From 5d7adb2604f9f1e083523e71880533ac750c1cc4 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Thu, 6 Aug 2026 16:34:53 -0500 Subject: feat: add shared no-grad two-state debiaser --- experiments/two_state_debias_smoke.py | 119 ++++++++++++++++++++++++++++++++++ 1 file changed, 119 insertions(+) create mode 100644 experiments/two_state_debias_smoke.py (limited to 'experiments') diff --git a/experiments/two_state_debias_smoke.py b/experiments/two_state_debias_smoke.py new file mode 100644 index 0000000..9b57819 --- /dev/null +++ b/experiments/two_state_debias_smoke.py @@ -0,0 +1,119 @@ +#!/usr/bin/env python3 +"""Strict-locality smoke test for the shared two-state SDIL filter.""" + +from __future__ import annotations + +from pathlib import Path +import sys + +import torch + +ROOT = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(ROOT)) + +from sdil.two_state_debias import ( # noqa: E402 + LocalAffineDebiaser, + two_state_difference, +) + + +def main() -> None: + torch.manual_seed(20260806) + templates = [torch.zeros(7), torch.zeros(3, 4)] + centers = [0.2, -0.1] + scales = [1.3, 0.8] + affine = LocalAffineDebiaser( + templates, + feature_centers=centers, + feature_scales=scales, + affine=True, + ) + constant = LocalAffineDebiaser( + templates, + feature_centers=centers, + feature_scales=scales, + affine=False, + ) + intercepts = [torch.linspace(-0.4, 0.5, 7), torch.randn(3, 4) * 0.2] + slopes = [torch.linspace(0.3, 1.0, 7), torch.randn(3, 4) * 0.4] + + train_features = [] + for value in torch.linspace(-1.5, 1.1, 50): + train_features.append([ + torch.full_like(templates[0], value), + torch.full_like(templates[1], -0.6 * value + 0.2), + ]) + for _ in range(20): + for features in train_features: + measurements = [ + intercept + slope * ((feature - center) / scale) + for intercept, slope, feature, center, scale in zip( + intercepts, slopes, features, centers, scales) + ] + affine.update_neutral(features, measurements, 0.2) + constant.update_neutral(features, measurements, 0.2) + + held_features = [ + torch.full_like(templates[0], 1.35), + torch.full_like(templates[1], -0.75), + ] + held_bias = [ + intercept + slope * ((feature - center) / scale) + for intercept, slope, feature, center, scale in zip( + intercepts, slopes, held_features, centers, scales) + ] + affine_error = torch.mean(torch.stack([ + (prediction - target).square().mean() + for prediction, target in zip(affine.predict(held_features), held_bias) + ])).sqrt() + constant_error = torch.mean(torch.stack([ + (prediction - target).square().mean() + for prediction, target in zip(constant.predict(held_features), held_bias) + ])).sqrt() + assert affine_error < 1e-5 + assert constant_error > 0.1 + + first = [torch.randn_like(template) for template in templates] + common = [torch.randn_like(template) for template in templates] + second = [value + 0.03 * torch.ones_like(value) for value in first] + difference_a = two_state_difference(first, second, 0.2) + difference_b = two_state_difference( + [value + offset for value, offset in zip(first, common)], + [value + offset for value, offset in zip(second, common)], + 0.2, + ) + common_mode_error = max( + float((a - b).abs().max()) for a, b in zip(difference_a, difference_b) + ) + assert common_mode_error < 2e-6 + + teaching = [torch.randn_like(template) for template in templates] + eligibility = [torch.randn_like(template) for template in templates] + update_a = affine.replay_updates( + held_features, teaching, eligibility, 0.07) + downstream = torch.randn(1024, 1024) + downstream.normal_() + update_b = affine.replay_updates( + held_features, teaching, eligibility, 0.07) + assert all(torch.equal(a, b) for a, b in zip(update_a, update_b)) + assert all(not update.requires_grad for update in update_a) + + try: + affine.predict([held_features[0].requires_grad_(), held_features[1]]) + except ValueError: + pass + else: + raise AssertionError("requires-grad input was not rejected") + + print({ + "affine_heldout_rmse": float(affine_error), + "constant_heldout_rmse": float(constant_error), + "neutral_observations_each": affine.neutral_observations, + "common_mode_max_float_error": common_mode_error, + "downstream_independence_exact": True, + "requires_grad_rejected": True, + }) + + +if __name__ == "__main__": + main() -- cgit v1.2.3