From db4fa5aaae2857270f7fc512e65ce4be686fde35 Mon Sep 17 00:00:00 2001 From: xmarre Date: Tue, 31 Mar 2026 03:10:23 +0200 Subject: [PATCH] Move forecaster coefficient recompute off the predict path --- comfyui_spectrum/forecast.py | 24 +++++++++++++++++++----- tests/smoke_runtime.py | 29 +++++++++++++++++++++++++++++ 2 files changed, 48 insertions(+), 5 deletions(-) diff --git a/comfyui_spectrum/forecast.py b/comfyui_spectrum/forecast.py index 0a866ef..6c3efbb 100644 --- a/comfyui_spectrum/forecast.py +++ b/comfyui_spectrum/forecast.py @@ -46,6 +46,7 @@ class ChebyshevSpectrumForecaster: self._coeff = None self._cached_degree = None self._cache_dirty = True + self._recompute_coeff() @property def feature_shape(self) -> Optional[torch.Size]: @@ -74,6 +75,7 @@ class ChebyshevSpectrumForecaster: self._coeff = None self._cached_degree = None self._cache_dirty = True + self._recompute_coeff() def _build_design(self, coords: torch.Tensor, degree: int) -> torch.Tensor: coords = coords.reshape(-1, 1).to(torch.float32) @@ -98,18 +100,30 @@ class ChebyshevSpectrumForecaster: chol = torch.linalg.cholesky(lhs + jitter * torch.eye(p, device=lhs.device, dtype=lhs.dtype)) return torch.cholesky_solve(rhs, chol) - def _ensure_coeff(self) -> tuple[int, torch.Tensor]: - degree = min(self.degree, len(self._history) - 1) - if not self._cache_dirty and self._coeff is not None and self._cached_degree == degree: - return degree, self._coeff + def _recompute_coeff(self) -> None: + if not self.ready(): + self._coeff = None + self._cached_degree = None + self._cache_dirty = True + return + degree = min(self.degree, len(self._history) - 1) coords = torch.tensor([entry.time_coord for entry in self._history], device="cpu", dtype=torch.float32) features = torch.stack([entry.feature_flat for entry in self._history], dim=0) design = self._build_design(coords, degree) self._coeff = self._solve(design, features) self._cached_degree = degree self._cache_dirty = False - return degree, self._coeff + + def _ensure_coeff(self) -> tuple[int, torch.Tensor]: + degree = min(self.degree, len(self._history) - 1) + if not self._cache_dirty and self._coeff is not None and self._cached_degree == degree: + return degree, self._coeff + + self._recompute_coeff() + if self._coeff is None or self._cached_degree is None: + raise RuntimeError("Spectrum forecaster coefficients are not ready yet.") + return self._cached_degree, self._coeff def _linear_prediction(self, time_coord: float) -> torch.Tensor: last = self._history[-1] diff --git a/tests/smoke_runtime.py b/tests/smoke_runtime.py index 3d6d2ed..8d7fe38 100644 --- a/tests/smoke_runtime.py +++ b/tests/smoke_runtime.py @@ -36,6 +36,34 @@ def make_runtime(**overrides) -> SpectrumRuntime: return SpectrumRuntime(cfg) +def test_forecaster_recomputes_coeff_on_update_not_predict() -> None: + class CountingForecaster(ChebyshevSpectrumForecaster): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.solve_calls = 0 + + def _solve(self, design: torch.Tensor, features: torch.Tensor) -> torch.Tensor: + self.solve_calls += 1 + return super()._solve(design, features) + + forecaster = CountingForecaster(degree=4, ridge_lambda=0.1, max_history=8) + for idx in range(5): + forecaster.update(float(idx), torch.randn(1, 8, 4)) + + assert forecaster._history[0].feature_flat.device.type == "cpu" + + solve_calls_before_predict = forecaster.solve_calls + first = forecaster.predict(5.0, blend_weight=0.5) + second = forecaster.predict(5.5, blend_weight=0.5) + assert first.shape == (1, 8, 4) + assert second.shape == (1, 8, 4) + assert forecaster.solve_calls == solve_calls_before_predict + + solve_calls_before_update = forecaster.solve_calls + forecaster.update(6.0, torch.randn(1, 8, 4)) + assert forecaster.solve_calls == solve_calls_before_update + 1 + + def test_solver_step_scheduler() -> None: runtime = make_runtime() sample_sigmas = torch.linspace(1.0, 0.0, 51) @@ -642,6 +670,7 @@ def test_forecast_feature_sanitization_stats_only_report_real_violations() -> No def main() -> None: + test_forecaster_recomputes_coeff_on_update_not_predict() test_solver_step_scheduler() test_forecast_fallback_reconciles_bookkeeping() test_observe_actual_feature_clears_forecast_latch()