Merge pull request #16 from xmarre/codex/fix-spectrum-update-overhead
Fix Spectrum forecaster update overhead on actual steps
This commit is contained in:
+135
-51
@@ -10,7 +10,6 @@ import torch
|
||||
class _HistoryEntry:
|
||||
time_coord: float
|
||||
feature_flat: torch.Tensor
|
||||
basis_row: torch.Tensor
|
||||
|
||||
|
||||
class ChebyshevSpectrumForecaster:
|
||||
@@ -34,6 +33,7 @@ class ChebyshevSpectrumForecaster:
|
||||
self._feature_dtype: Optional[torch.dtype] = None
|
||||
self._device: Optional[torch.device] = None
|
||||
self._predict_device: Optional[torch.device] = None
|
||||
self._predict_dtype: Optional[torch.dtype] = None
|
||||
self._output_device: Optional[torch.device] = None
|
||||
self._coeff: Optional[torch.Tensor] = None
|
||||
self._coeff_device: Optional[torch.Tensor] = None
|
||||
@@ -45,6 +45,7 @@ class ChebyshevSpectrumForecaster:
|
||||
self._previous_time_coord: Optional[float] = None
|
||||
self._latest_feature_flat_device: Optional[torch.Tensor] = None
|
||||
self._latest_time_coord: Optional[float] = None
|
||||
self._linear_mirrors_enabled = True
|
||||
|
||||
def configure(self, degree: int, ridge_lambda: float, max_history: int) -> None:
|
||||
self.degree = int(degree)
|
||||
@@ -79,73 +80,118 @@ class ChebyshevSpectrumForecaster:
|
||||
predict_device: Optional[torch.device] = None,
|
||||
output_device: Optional[torch.device] = None,
|
||||
output_dtype: Optional[torch.dtype] = None,
|
||||
blend_weight: Optional[float] = None,
|
||||
) -> None:
|
||||
feat = feature.detach()
|
||||
if self._feature_shape is None:
|
||||
self._feature_shape = feat.shape
|
||||
self._device = torch.device("cpu")
|
||||
elif feat.shape != self._feature_shape:
|
||||
raise ValueError(
|
||||
f"Spectrum feature shape changed from {tuple(self._feature_shape)} to {tuple(feat.shape)}."
|
||||
)
|
||||
self._feature_dtype = output_dtype if output_dtype is not None else feat.dtype
|
||||
if predict_device is not None:
|
||||
self._predict_device = predict_device
|
||||
else:
|
||||
self._predict_device = feat.device
|
||||
resolved_output_dtype = output_dtype if output_dtype is not None else feat.dtype
|
||||
resolved_predict_device = predict_device if predict_device is not None else feat.device
|
||||
resolved_predict_dtype = self._resolve_predict_dtype(resolved_output_dtype)
|
||||
resolved_stats_device = resolved_predict_device if resolved_predict_device is not None else feat.device
|
||||
previous_predict_device = self._predict_device
|
||||
previous_predict_dtype = self._predict_dtype
|
||||
previous_stats_device = self._device
|
||||
self._feature_dtype = resolved_output_dtype
|
||||
self._predict_device = resolved_predict_device
|
||||
if output_device is not None:
|
||||
self._output_device = output_device
|
||||
else:
|
||||
self._output_device = feat.device
|
||||
if self._predict_device is None:
|
||||
self._predict_device = self._output_device
|
||||
if resolved_stats_device is None:
|
||||
resolved_stats_device = self._predict_device
|
||||
self._device = resolved_stats_device
|
||||
self._predict_dtype = resolved_predict_dtype
|
||||
|
||||
feature_flat = feat.reshape(-1).to(device="cpu", dtype=torch.float32, copy=True)
|
||||
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)
|
||||
linear_mirrors_enabled = True if blend_weight is None else float(blend_weight) < (1.0 - 1e-12)
|
||||
predict_context_changed = (
|
||||
previous_predict_device != self._predict_device or previous_predict_dtype != self._predict_dtype
|
||||
)
|
||||
if previous_stats_device != self._device:
|
||||
self._coeff = None
|
||||
self._coeff_device = None
|
||||
self._cached_degree = None
|
||||
self._cache_dirty = True
|
||||
self._rebuild_stats()
|
||||
|
||||
feature_flat = feat.reshape(-1).to(device=self._device, dtype=torch.float32, copy=False)
|
||||
entry = _HistoryEntry(float(time_coord), self._archive_feature_for_history(feat))
|
||||
self._history.append(entry)
|
||||
self._ensure_stats_initialized(feature_flat.numel())
|
||||
self._accumulate_entry(entry, sign=1.0)
|
||||
self._accumulate_feature(time_coord=entry.time_coord, feature_flat=feature_flat, sign=1.0)
|
||||
if len(self._history) > self.max_history:
|
||||
oldest = self._history.pop(0)
|
||||
self._accumulate_entry(oldest, sign=-1.0)
|
||||
self._refresh_prediction_mirrors()
|
||||
self._sync_linear_prediction_mirrors(
|
||||
entry,
|
||||
feat=feat,
|
||||
linear_mirrors_enabled=linear_mirrors_enabled,
|
||||
force_rebuild=predict_context_changed or linear_mirrors_enabled != self._linear_mirrors_enabled,
|
||||
)
|
||||
self._coeff = None
|
||||
self._coeff_device = 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)
|
||||
if self._device is None:
|
||||
raise RuntimeError("Spectrum forecaster stats device is not configured.")
|
||||
if self._gram is None or self._gram.shape != (p, p) or self._gram.device != self._device:
|
||||
self._gram = torch.zeros((p, p), device=self._device, dtype=torch.float32)
|
||||
if self._rhs is None or self._rhs.shape != (p, feature_dim) or self._rhs.device != self._device:
|
||||
self._rhs = torch.zeros((p, feature_dim), device=self._device, dtype=torch.float32)
|
||||
|
||||
def _accumulate_feature(self, time_coord: float, feature_flat: torch.Tensor, *, sign: float) -> None:
|
||||
if self._gram is None or self._rhs is None:
|
||||
self._ensure_stats_initialized(feature_flat.numel())
|
||||
basis_row = self._build_design(
|
||||
torch.tensor([float(time_coord)], device=self._device, dtype=torch.float32),
|
||||
self.degree,
|
||||
).reshape(-1)
|
||||
self._gram.add_(float(sign) * torch.outer(basis_row, basis_row))
|
||||
self._rhs.addmm_(basis_row.unsqueeze(1), feature_flat.unsqueeze(0), beta=1.0, alpha=float(sign))
|
||||
|
||||
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)))
|
||||
feature_flat = entry.feature_flat.to(device=self._device, dtype=torch.float32)
|
||||
self._accumulate_feature(entry.time_coord, feature_flat, sign=sign)
|
||||
|
||||
def _archive_feature_for_history(self, feature: torch.Tensor) -> torch.Tensor:
|
||||
flat = feature.reshape(-1)
|
||||
if flat.device.type == "cpu":
|
||||
return flat.clone()
|
||||
use_pinned_copy = flat.device.type == "cuda"
|
||||
archived = torch.empty(flat.shape, device="cpu", dtype=flat.dtype, pin_memory=use_pinned_copy)
|
||||
archived.copy_(flat, non_blocking=use_pinned_copy)
|
||||
return archived
|
||||
|
||||
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] = []
|
||||
if self._device is None:
|
||||
self._gram = None
|
||||
self._rhs = None
|
||||
return
|
||||
self._gram = torch.zeros((p, p), device=self._device, dtype=torch.float32)
|
||||
self._rhs = torch.zeros((p, feature_dim), device=self._device, dtype=torch.float32) if feature_dim > 0 else None
|
||||
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
|
||||
self._accumulate_entry(entry, sign=1.0)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_predict_dtype(dtype: torch.dtype) -> torch.dtype:
|
||||
return dtype if torch.is_floating_point(torch.empty((), dtype=dtype)) else torch.float32
|
||||
|
||||
def _mirror_feature_for_prediction(self, feature: torch.Tensor) -> torch.Tensor:
|
||||
if self._predict_device is None or self._predict_dtype is None:
|
||||
raise RuntimeError("Spectrum forecaster prediction device is not configured.")
|
||||
return feature.reshape(-1).to(device=self._predict_device, dtype=self._predict_dtype, copy=True)
|
||||
|
||||
def _refresh_prediction_mirrors(self) -> None:
|
||||
self._previous_feature_flat_device = None
|
||||
@@ -153,19 +199,48 @@ class ChebyshevSpectrumForecaster:
|
||||
self._latest_feature_flat_device = None
|
||||
self._latest_time_coord = None
|
||||
self._coeff_device = None
|
||||
if self._predict_device is None or not self._history:
|
||||
if (
|
||||
self._predict_device is None
|
||||
or self._predict_dtype is None
|
||||
or not self._history
|
||||
or not self._linear_mirrors_enabled
|
||||
):
|
||||
return
|
||||
if len(self._history) >= 2:
|
||||
previous = self._history[-2]
|
||||
self._previous_feature_flat_device = previous.feature_flat.to(
|
||||
device=self._predict_device, dtype=torch.float32
|
||||
device=self._predict_device, dtype=self._predict_dtype
|
||||
)
|
||||
self._previous_time_coord = previous.time_coord
|
||||
latest = self._history[-1]
|
||||
self._latest_feature_flat_device = latest.feature_flat.to(device=self._predict_device, dtype=torch.float32)
|
||||
self._latest_feature_flat_device = latest.feature_flat.to(device=self._predict_device, dtype=self._predict_dtype)
|
||||
self._latest_time_coord = latest.time_coord
|
||||
if self._coeff is not None:
|
||||
self._coeff_device = self._coeff.to(device=self._predict_device, dtype=torch.float32)
|
||||
|
||||
def _sync_linear_prediction_mirrors(
|
||||
self,
|
||||
entry: _HistoryEntry,
|
||||
*,
|
||||
feat: torch.Tensor,
|
||||
linear_mirrors_enabled: bool,
|
||||
force_rebuild: bool,
|
||||
) -> None:
|
||||
self._linear_mirrors_enabled = linear_mirrors_enabled
|
||||
if not self._linear_mirrors_enabled or not self._history:
|
||||
self._previous_feature_flat_device = None
|
||||
self._previous_time_coord = None
|
||||
self._latest_feature_flat_device = None
|
||||
self._latest_time_coord = None
|
||||
return
|
||||
if force_rebuild or self._latest_feature_flat_device is None:
|
||||
self._refresh_prediction_mirrors()
|
||||
self._latest_feature_flat_device = self._mirror_feature_for_prediction(feat)
|
||||
self._latest_time_coord = entry.time_coord
|
||||
return
|
||||
|
||||
self._previous_feature_flat_device = self._latest_feature_flat_device
|
||||
self._previous_time_coord = self._latest_time_coord
|
||||
self._latest_feature_flat_device = self._mirror_feature_for_prediction(feat)
|
||||
self._latest_time_coord = entry.time_coord
|
||||
|
||||
def _build_design(self, coords: torch.Tensor, degree: int) -> torch.Tensor:
|
||||
coords = coords.reshape(-1, 1).to(torch.float32)
|
||||
@@ -216,18 +291,13 @@ class ChebyshevSpectrumForecaster:
|
||||
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)
|
||||
if self._predict_device is not None:
|
||||
self._coeff_device = self._coeff.to(device=self._predict_device, dtype=torch.float32)
|
||||
else:
|
||||
self._coeff_device = None
|
||||
self._coeff_device = None
|
||||
self._cached_degree = degree
|
||||
self._cache_dirty = False
|
||||
|
||||
def _ensure_coeff(self) -> tuple[int, torch.Tensor]:
|
||||
degree = self.degree
|
||||
if not self._cache_dirty and self._coeff is not None and self._cached_degree == degree:
|
||||
if self._coeff_device is None and self._predict_device is not None:
|
||||
self._coeff_device = self._coeff.to(device=self._predict_device, dtype=torch.float32)
|
||||
return degree, self._coeff
|
||||
|
||||
self._recompute_coeff()
|
||||
@@ -235,6 +305,16 @@ class ChebyshevSpectrumForecaster:
|
||||
raise RuntimeError("Spectrum forecaster coefficients are not ready yet.")
|
||||
return self._cached_degree, self._coeff
|
||||
|
||||
def _ensure_coeff_device(self) -> torch.Tensor:
|
||||
if self._coeff is None or self._predict_device is None or self._predict_dtype is None:
|
||||
raise RuntimeError("Spectrum forecaster prediction coefficients are not ready yet.")
|
||||
if self._coeff_device is None:
|
||||
if self._coeff.device == self._predict_device and self._coeff.dtype == self._predict_dtype:
|
||||
self._coeff_device = self._coeff
|
||||
else:
|
||||
self._coeff_device = self._coeff.to(device=self._predict_device, dtype=self._predict_dtype)
|
||||
return self._coeff_device
|
||||
|
||||
def _linear_prediction(self, time_coord: float) -> torch.Tensor:
|
||||
if self._latest_feature_flat_device is None or self._latest_time_coord is None:
|
||||
raise RuntimeError("Spectrum forecaster has no cached feature history.")
|
||||
@@ -256,6 +336,7 @@ class ChebyshevSpectrumForecaster:
|
||||
or self._feature_dtype is None
|
||||
or self._device is None
|
||||
or self._predict_device is None
|
||||
or self._predict_dtype is None
|
||||
or self._output_device is None
|
||||
):
|
||||
raise RuntimeError("Spectrum forecaster has no cached feature history.")
|
||||
@@ -263,13 +344,16 @@ class ChebyshevSpectrumForecaster:
|
||||
raise RuntimeError("Spectrum forecaster is not ready yet.")
|
||||
|
||||
degree, _ = self._ensure_coeff()
|
||||
if self._coeff_device is None:
|
||||
raise RuntimeError("Spectrum forecaster prediction coefficients are not mirrored to the model device.")
|
||||
coeff_device = self._ensure_coeff_device()
|
||||
|
||||
coord_star = torch.tensor([float(time_coord)], device=self._predict_device, dtype=torch.float32)
|
||||
design_star = self._build_design(coord_star, degree)
|
||||
spectral = (design_star @ self._coeff_device).reshape(self._feature_shape)
|
||||
design_star = self._build_design(coord_star, degree).to(dtype=coeff_device.dtype)
|
||||
spectral = (design_star @ coeff_device).reshape(self._feature_shape)
|
||||
|
||||
linear = self._linear_prediction(time_coord).reshape(self._feature_shape)
|
||||
out = float(blend_weight) * spectral + (1.0 - float(blend_weight)) * linear
|
||||
blend = float(blend_weight)
|
||||
if blend >= (1.0 - 1e-12):
|
||||
out = spectral
|
||||
else:
|
||||
linear = self._linear_prediction(time_coord).reshape(self._feature_shape)
|
||||
out = blend * spectral + (1.0 - blend) * linear
|
||||
return out.to(device=self._output_device, dtype=self._feature_dtype)
|
||||
|
||||
@@ -567,6 +567,7 @@ class SpectrumRuntime:
|
||||
predict_device=step.actual_feature_device,
|
||||
output_device=step.actual_feature_device,
|
||||
output_dtype=step.actual_feature_dtype,
|
||||
blend_weight=self.cfg.blend_weight,
|
||||
)
|
||||
except ValueError:
|
||||
self._disable_forecasting("combined actual feature shape changed across solver steps")
|
||||
|
||||
Reference in New Issue
Block a user