Cache incremental forecaster stats on history updates
This commit is contained in:
@@ -10,6 +10,7 @@ import torch
|
||||
class _HistoryEntry:
|
||||
time_coord: float
|
||||
feature_flat: torch.Tensor
|
||||
basis_row: torch.Tensor
|
||||
|
||||
|
||||
class ChebyshevSpectrumForecaster:
|
||||
@@ -36,6 +37,8 @@ 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)
|
||||
@@ -46,6 +49,7 @@ class ChebyshevSpectrumForecaster:
|
||||
self._coeff = None
|
||||
self._cached_degree = None
|
||||
self._cache_dirty = True
|
||||
self._rebuild_stats()
|
||||
self._recompute_coeff()
|
||||
|
||||
@property
|
||||
@@ -69,14 +73,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 +142,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
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user