Merge pull request #10 from xmarre/codex/speed-up-forecast-steps
Fix forecast-step coefficient recomputation timing
This commit is contained in:
@@ -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]
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user