diff --git a/comfyui_spectrum/forecast.py b/comfyui_spectrum/forecast.py index 6c3efbb..07e7547 100644 --- a/comfyui_spectrum/forecast.py +++ b/comfyui_spectrum/forecast.py @@ -10,6 +10,7 @@ import torch class _HistoryEntry: time_coord: float feature_flat: torch.Tensor + basis_row: torch.Tensor class ChebyshevSpectrumForecaster: @@ -36,16 +37,23 @@ class ChebyshevSpectrumForecaster: self._coeff: Optional[torch.Tensor] = None self._cached_degree: Optional[int] = None self._cache_dirty = True + self._gram: Optional[torch.Tensor] = None + self._rhs: Optional[torch.Tensor] = None def configure(self, degree: int, ridge_lambda: float, max_history: int) -> None: self.degree = int(degree) self.ridge_lambda = float(ridge_lambda) self.max_history = int(max_history) - if len(self._history) > self.max_history: + if self.max_history < 0: + raise ValueError("max_history must be non-negative.") + if self.max_history == 0: + self._history = [] + elif len(self._history) > self.max_history: self._history = self._history[-self.max_history :] self._coeff = None self._cached_degree = None self._cache_dirty = True + self._rebuild_stats() self._recompute_coeff() @property @@ -69,14 +77,51 @@ class ChebyshevSpectrumForecaster: ) feature_flat = feat.reshape(-1).to(device="cpu", dtype=torch.float32, copy=True) - self._history.append(_HistoryEntry(float(time_coord), feature_flat)) + basis_row = self._build_design( + torch.tensor([float(time_coord)], device="cpu", dtype=torch.float32), + self.degree, + ).reshape(-1) + entry = _HistoryEntry(float(time_coord), feature_flat, basis_row) + self._history.append(entry) + self._ensure_stats_initialized(feature_flat.numel()) + self._accumulate_entry(entry, sign=1.0) if len(self._history) > self.max_history: - self._history.pop(0) + oldest = self._history.pop(0) + self._accumulate_entry(oldest, sign=-1.0) self._coeff = None self._cached_degree = None self._cache_dirty = True self._recompute_coeff() + def _ensure_stats_initialized(self, feature_dim: int) -> None: + p = self.degree + 1 + if self._gram is None or self._gram.shape != (p, p): + self._gram = torch.zeros((p, p), device="cpu", dtype=torch.float32) + if self._rhs is None or self._rhs.shape != (p, feature_dim): + self._rhs = torch.zeros((p, feature_dim), device="cpu", dtype=torch.float32) + + def _accumulate_entry(self, entry: _HistoryEntry, *, sign: float) -> None: + if self._gram is None or self._rhs is None: + self._ensure_stats_initialized(entry.feature_flat.numel()) + self._gram.add_(float(sign) * torch.outer(entry.basis_row, entry.basis_row)) + self._rhs.add_(float(sign) * (entry.basis_row.unsqueeze(1) * entry.feature_flat.unsqueeze(0))) + + def _rebuild_stats(self) -> None: + feature_dim = self._history[0].feature_flat.numel() if self._history else 0 + p = self.degree + 1 + self._gram = torch.zeros((p, p), device="cpu", dtype=torch.float32) + self._rhs = torch.zeros((p, feature_dim), device="cpu", dtype=torch.float32) if feature_dim > 0 else None + rebuilt: List[_HistoryEntry] = [] + for entry in self._history: + basis_row = self._build_design( + torch.tensor([entry.time_coord], device="cpu", dtype=torch.float32), + self.degree, + ).reshape(-1) + rebuilt_entry = _HistoryEntry(entry.time_coord, entry.feature_flat, basis_row) + rebuilt.append(rebuilt_entry) + self._accumulate_entry(rebuilt_entry, sign=1.0) + self._history = rebuilt + def _build_design(self, coords: torch.Tensor, degree: int) -> torch.Tensor: coords = coords.reshape(-1, 1).to(torch.float32) cols = [torch.ones((coords.shape[0], 1), device=coords.device, dtype=torch.float32)] @@ -101,22 +146,34 @@ class ChebyshevSpectrumForecaster: return torch.cholesky_solve(rhs, chol) def _recompute_coeff(self) -> None: - if not self.ready(): + if not self.ready() or self._gram is None or self._rhs is None: 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) + degree = self.degree + lhs = self._gram + rhs = self._rhs + if lhs.numel() == 0 or rhs.numel() == 0: + self._coeff = None + self._cached_degree = None + self._cache_dirty = True + return + if self.ridge_lambda > 0.0: + lhs = lhs + self.ridge_lambda * torch.eye(degree + 1, device=lhs.device, dtype=lhs.dtype) + try: + chol = torch.linalg.cholesky(lhs) + except RuntimeError: + diag_mean = lhs.diag().mean() if lhs.numel() else torch.tensor(1.0, device=lhs.device) + jitter = max(float(diag_mean.item()) * 1e-6, 1e-8) + chol = torch.linalg.cholesky(lhs + jitter * torch.eye(degree + 1, device=lhs.device, dtype=lhs.dtype)) + self._coeff = torch.cholesky_solve(rhs, chol) self._cached_degree = degree self._cache_dirty = False def _ensure_coeff(self) -> tuple[int, torch.Tensor]: - degree = min(self.degree, len(self._history) - 1) + degree = self.degree if not self._cache_dirty and self._coeff is not None and self._cached_degree == degree: return degree, self._coeff diff --git a/tests/smoke_runtime.py b/tests/smoke_runtime.py index 8d7fe38..cd534cd 100644 --- a/tests/smoke_runtime.py +++ b/tests/smoke_runtime.py @@ -40,11 +40,11 @@ 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 + self.recompute_calls = 0 - def _solve(self, design: torch.Tensor, features: torch.Tensor) -> torch.Tensor: - self.solve_calls += 1 - return super()._solve(design, features) + def _recompute_coeff(self) -> None: + self.recompute_calls += 1 + return super()._recompute_coeff() forecaster = CountingForecaster(degree=4, ridge_lambda=0.1, max_history=8) for idx in range(5): @@ -52,16 +52,16 @@ def test_forecaster_recomputes_coeff_on_update_not_predict() -> None: assert forecaster._history[0].feature_flat.device.type == "cpu" - solve_calls_before_predict = forecaster.solve_calls + recompute_calls_before_predict = forecaster.recompute_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 + assert forecaster.recompute_calls == recompute_calls_before_predict - solve_calls_before_update = forecaster.solve_calls + recompute_calls_before_update = forecaster.recompute_calls forecaster.update(6.0, torch.randn(1, 8, 4)) - assert forecaster.solve_calls == solve_calls_before_update + 1 + assert forecaster.recompute_calls == recompute_calls_before_update + 1 def test_solver_step_scheduler() -> None: