diff options
| -rw-r--r-- | experiments/two_state_debias_smoke.py | 22 | ||||
| -rw-r--r-- | sdil/two_state_debias.py | 84 |
2 files changed, 106 insertions, 0 deletions
diff --git a/experiments/two_state_debias_smoke.py b/experiments/two_state_debias_smoke.py index 9b57819..fc659e6 100644 --- a/experiments/two_state_debias_smoke.py +++ b/experiments/two_state_debias_smoke.py @@ -12,6 +12,7 @@ ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) from sdil.two_state_debias import ( # noqa: E402 + BatchedLocalAffineDebiaser, LocalAffineDebiaser, two_state_difference, ) @@ -105,6 +106,24 @@ def main() -> None: else: raise AssertionError("requires-grad input was not rejected") + feature = torch.linspace(-1.0, 1.0, 64).reshape(64, 1, 1) + feature = torch.cat((feature, feature.square()), dim=2) + target = 0.2 + torch.tensor([[[0.7, -0.4]]]) * feature + batched_affine = BatchedLocalAffineDebiaser( + [feature], feature_centers=[0.0], feature_scales=[1.0], affine=True) + batched_constant = BatchedLocalAffineDebiaser( + [feature], feature_centers=[0.0], feature_scales=[1.0], affine=False) + for _ in range(20): + batched_affine.update_neutral([feature], [target], 0.2) + batched_constant.update_neutral([feature], [target], 0.2) + held_feature = torch.tensor([[[-1.3, 1.4]], [[1.3, 1.4]]]) + held_target = 0.2 + torch.tensor([[[0.7, -0.4]]]) * held_feature + held_affine = batched_affine.residual([held_feature], [held_target])[0] + held_constant = batched_constant.residual([held_feature], [held_target])[0] + batched_affine_rmse = float(held_affine.square().mean().sqrt()) + batched_constant_rmse = float(held_constant.square().mean().sqrt()) + assert batched_affine_rmse < 0.05 * batched_constant_rmse + print({ "affine_heldout_rmse": float(affine_error), "constant_heldout_rmse": float(constant_error), @@ -112,6 +131,9 @@ def main() -> None: "common_mode_max_float_error": common_mode_error, "downstream_independence_exact": True, "requires_grad_rejected": True, + "batched_affine_heldout_rmse": batched_affine_rmse, + "batched_constant_heldout_rmse": batched_constant_rmse, + "batched_neutral_observations": batched_affine.neutral_observations, }) 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) + ] |
