summaryrefslogtreecommitdiff
path: root/sdil
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 16:53:27 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 16:53:27 -0500
commit2e2d166f5ee63526fa99bec3a82320c365f76f62 (patch)
tree3ea9f6aa3acdfcff9db942bebed6ffb3d4132742 /sdil
parenteb2416ee22f916f4e1d29405d634a0ee613a525c (diff)
feat: add per-cell batched local predictor
Diffstat (limited to 'sdil')
-rw-r--r--sdil/two_state_debias.py84
1 files changed, 84 insertions, 0 deletions
diff --git a/sdil/two_state_debias.py b/sdil/two_state_debias.py
index fad980a..b1dfdde 100644
--- a/sdil/two_state_debias.py
+++ b/sdil/two_state_debias.py
@@ -179,3 +179,87 @@ class LocalAffineDebiaser:
clone.neutral_observations = self.neutral_observations
return clone
+
+class BatchedLocalAffineDebiaser:
+ """Per-cell LMS whose coefficients are shared across observations.
+
+ Inputs have shape ``[observations, *local_shape]``. The leading axis is
+ reduced only inside each cell's predictor update; cells remain independent.
+ """
+
+ def __init__(
+ self,
+ templates: Iterable[Tensor],
+ *,
+ feature_centers: Iterable[Tensor | float],
+ feature_scales: Iterable[Tensor | float],
+ affine: bool = True,
+ ) -> None:
+ templates = _detached(templates)
+ if any(template.ndim < 1 for template in templates):
+ raise ValueError("batched templates require an observation axis")
+ self.filter = LocalAffineDebiaser(
+ [template[0] for template in templates],
+ feature_centers=feature_centers,
+ feature_scales=feature_scales,
+ affine=affine,
+ )
+ self.affine = affine
+ self.neutral_observations = 0
+
+ @torch.no_grad()
+ def predict(self, local_features: Iterable[Tensor]) -> list[Tensor]:
+ features = _detached(local_features)
+ if len(features) != len(self.filter.states):
+ raise ValueError("feature collection length changed")
+ predictions = []
+ for state, feature in zip(self.filter.states, features):
+ normalized = self.filter._feature(state, feature)
+ prediction = state.intercept.unsqueeze(0).expand_as(feature)
+ if self.affine:
+ prediction = prediction + state.slope.unsqueeze(0) * normalized
+ predictions.append(prediction.clone())
+ return predictions
+
+ @torch.no_grad()
+ def update_neutral(
+ self,
+ local_features: Iterable[Tensor],
+ neutral_measurements: Iterable[Tensor],
+ learning_rate: float,
+ ) -> None:
+ features = _detached(local_features)
+ measurements = _detached(neutral_measurements)
+ if not (len(features) == len(measurements) == len(self.filter.states)):
+ raise ValueError("neutral tuple lengths disagree")
+ batch_sizes = {feature.shape[0] for feature in features}
+ if len(batch_sizes) != 1:
+ raise ValueError("local populations have different observation counts")
+ if any(
+ feature.shape != measurement.shape
+ for feature, measurement in zip(features, measurements)
+ ):
+ raise ValueError("feature and measurement shapes disagree")
+ observations = next(iter(batch_sizes))
+ for index in range(observations):
+ self.filter.update_neutral(
+ [feature[index] for feature in features],
+ [measurement[index] for measurement in measurements],
+ learning_rate,
+ )
+ self.neutral_observations += observations
+
+ @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)
+ ]