summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--experiments/two_state_debias_smoke.py22
-rw-r--r--sdil/two_state_debias.py84
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)
+ ]