summaryrefslogtreecommitdiff
path: root/sdil
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-10 04:31:50 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-10 04:31:50 -0500
commit927241c496d2508344f463c9cf96b280b5cc2830 (patch)
treef2a8bb017f305d57581103c71345a47ba09ca55f /sdil
parent29808ed4f8d550eb0ddcbc300b33f2dc596721e3 (diff)
feat: add local online least-squares residual predictor
Diffstat (limited to 'sdil')
-rw-r--r--sdil/rain_ep_adapter.py26
-rw-r--r--sdil/two_state_debias.py137
2 files changed, 156 insertions, 7 deletions
diff --git a/sdil/rain_ep_adapter.py b/sdil/rain_ep_adapter.py
index b16e96d..7dd47eb 100644
--- a/sdil/rain_ep_adapter.py
+++ b/sdil/rain_ep_adapter.py
@@ -18,6 +18,7 @@ import torch
from sdil.two_state_debias import (
BatchedLocalAffineDebiaser,
LocalAffineDebiaser,
+ LocalLeastSquaresDebiaser,
)
@@ -306,6 +307,7 @@ class DillavouUpdateCorrector:
neutral_cadence: int = 1,
drift_ratio: float = 0.0,
empirical_profile: DillavouBiasProfile | None = None,
+ predictor_kind: str = "nlms",
seed: int = 1729,
) -> None:
if mode not in self.MODES:
@@ -321,6 +323,8 @@ class DillavouUpdateCorrector:
raise ValueError("neutral cadence must be nonnegative")
if calibration_steps < 0:
raise ValueError("calibration steps must be nonnegative")
+ if predictor_kind not in {"nlms", "ols"}:
+ raise ValueError("predictor kind must be nlms or ols")
self.mode = mode
self.bias_ratio = bias_ratio
self.predictor_rate = predictor_rate
@@ -328,9 +332,10 @@ class DillavouUpdateCorrector:
self.neutral_cadence = neutral_cadence
self.drift_ratio = drift_ratio
self.empirical_profile = empirical_profile
+ self.predictor_kind = predictor_kind
self.seed = seed
self.steps = 0
- self.debiaser: LocalAffineDebiaser | None = None
+ self.debiaser: LocalAffineDebiaser | LocalLeastSquaresDebiaser | None = None
self._offsets: list[Tensor] | None = None
self._slopes: list[Tensor] | None = None
self._centers: list[Tensor] | None = None
@@ -402,12 +407,19 @@ class DillavouUpdateCorrector:
self._centers.append(parameter.clone())
self._scales.append(parameter_scale.clone())
feature_scales.append(torch.ones_like(parameter_scale))
- self.debiaser = LocalAffineDebiaser(
- clean,
- feature_centers=[0.0] * len(clean),
- feature_scales=feature_scales,
- affine=self.mode == "innovation",
- )
+ if self.mode == "innovation" and self.predictor_kind == "ols":
+ self.debiaser = LocalLeastSquaresDebiaser(
+ clean,
+ feature_centers=[0.0] * len(clean),
+ feature_scales=feature_scales,
+ )
+ else:
+ self.debiaser = LocalAffineDebiaser(
+ clean,
+ feature_centers=[0.0] * len(clean),
+ feature_scales=feature_scales,
+ affine=self.mode == "innovation",
+ )
@torch.no_grad()
def _measure(
diff --git a/sdil/two_state_debias.py b/sdil/two_state_debias.py
index b1dfdde..6273b2f 100644
--- a/sdil/two_state_debias.py
+++ b/sdil/two_state_debias.py
@@ -180,6 +180,143 @@ class LocalAffineDebiaser:
return clone
+@dataclass
+class LocalLeastSquaresState:
+ intercept: Tensor
+ slope: Tensor
+ feature_center: Tensor
+ feature_scale: Tensor
+ mean_feature: Tensor
+ mean_measurement: Tensor
+ feature_sum_squares: Tensor
+ cross_sum: Tensor
+
+
+class LocalLeastSquaresDebiaser:
+ """Per-element online affine regression from neutral observations.
+
+ Every element keeps its own scalar sufficient statistics. No observation,
+ coefficient, or update is shared across elements, and no autograd graph is
+ created. This is the direct online analogue of fitting a local baseline
+ relation before taking its innovation.
+ """
+
+ def __init__(
+ self,
+ templates: Iterable[Tensor],
+ *,
+ feature_centers: Iterable[Tensor | float],
+ feature_scales: Iterable[Tensor | float],
+ variance_floor: float = 1e-12,
+ ) -> None:
+ templates = _detached(templates)
+ centers = list(feature_centers)
+ scales = list(feature_scales)
+ if not (len(templates) == len(centers) == len(scales)):
+ raise ValueError("template and feature metadata lengths disagree")
+ if variance_floor <= 0.0:
+ raise ValueError("variance floor must be positive")
+ self.variance_floor = variance_floor
+ self.states = []
+ with torch.no_grad():
+ for template, center, scale in zip(templates, centers, scales):
+ center_tensor = torch.as_tensor(
+ center, dtype=template.dtype, device=template.device)
+ scale_tensor = torch.as_tensor(
+ scale, dtype=template.dtype, device=template.device)
+ if torch.any(scale_tensor <= 0):
+ raise ValueError("feature scales must be positive")
+ self.states.append(LocalLeastSquaresState(
+ intercept=torch.zeros_like(template),
+ slope=torch.zeros_like(template),
+ feature_center=center_tensor.clone(),
+ feature_scale=scale_tensor.clone(),
+ mean_feature=torch.zeros_like(template),
+ mean_measurement=torch.zeros_like(template),
+ feature_sum_squares=torch.zeros_like(template),
+ cross_sum=torch.zeros_like(template),
+ ))
+ self.neutral_observations = 0
+
+ @staticmethod
+ def _feature(
+ state: LocalLeastSquaresState, local_feature: Tensor
+ ) -> Tensor:
+ return (local_feature - state.feature_center) / state.feature_scale
+
+ @torch.no_grad()
+ def predict(self, local_features: Iterable[Tensor]) -> list[Tensor]:
+ features = _detached(local_features)
+ if len(features) != len(self.states):
+ raise ValueError("feature collection length changed")
+ return [
+ (
+ state.intercept
+ + state.slope * self._feature(state, feature)
+ ).clone()
+ for state, feature in zip(self.states, features)
+ ]
+
+ @torch.no_grad()
+ def update_neutral(
+ self,
+ local_features: Iterable[Tensor],
+ neutral_measurements: Iterable[Tensor],
+ learning_rate: float,
+ ) -> list[Tensor]:
+ if learning_rate != 1.0:
+ raise ValueError(
+ "online least squares requires predictor_rate=1")
+ features = _detached(local_features)
+ measurements = _detached(neutral_measurements)
+ if not (len(features) == len(measurements) == len(self.states)):
+ raise ValueError("neutral tuple lengths disagree")
+ count = self.neutral_observations + 1
+ residuals = []
+ for state, feature, measurement in zip(
+ self.states, features, measurements
+ ):
+ normalized = self._feature(state, feature)
+ residuals.append(
+ measurement
+ - state.intercept
+ - state.slope * normalized)
+ delta_feature = normalized - state.mean_feature
+ delta_measurement = measurement - state.mean_measurement
+ state.mean_feature.add_(delta_feature / count)
+ state.mean_measurement.add_(delta_measurement / count)
+ state.feature_sum_squares.add_(
+ delta_feature * (normalized - state.mean_feature))
+ state.cross_sum.add_(
+ delta_feature * (measurement - state.mean_measurement))
+ identifiable = state.feature_sum_squares > self.variance_floor
+ state.slope.copy_(torch.where(
+ identifiable,
+ state.cross_sum
+ / state.feature_sum_squares.clamp_min(self.variance_floor),
+ torch.zeros_like(state.slope),
+ ))
+ state.intercept.copy_(
+ state.mean_measurement - state.slope * state.mean_feature)
+ self.neutral_observations = count
+ return residuals
+
+ @torch.no_grad()
+ def residual(
+ self,
+ local_features: Iterable[Tensor],
+ teaching_measurements: Iterable[Tensor],
+ ) -> list[Tensor]:
+ measurements = _detached(teaching_measurements)
+ predictions = self.predict(local_features)
+ if len(measurements) != len(predictions):
+ raise ValueError("teaching tuple lengths disagree")
+ return [
+ measurement - prediction
+ for measurement, prediction in zip(measurements, predictions)
+ ]
+
+
class BatchedLocalAffineDebiaser:
"""Per-cell LMS whose coefficients are shared across observations.