From 2e2d166f5ee63526fa99bec3a82320c365f76f62 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Thu, 6 Aug 2026 16:53:27 -0500 Subject: feat: add per-cell batched local predictor --- sdil/two_state_debias.py | 84 ++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 84 insertions(+) (limited to 'sdil/two_state_debias.py') 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) + ] -- cgit v1.2.3