Merge pull request #10 from xmarre/codex/speed-up-forecast-steps

Fix forecast-step coefficient recomputation timing
This commit is contained in:
xmarre
2026-03-31 03:14:40 +02:00
committed by GitHub
2 changed files with 48 additions and 5 deletions
+19 -5
View File
@@ -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]
+29
View File
@@ -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()