summaryrefslogtreecommitdiff
path: root/experiments/rain_ep_layer_adapter_smoke.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 17:19:02 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 17:19:02 -0500
commit25f9b58f2fe478b7a4c728f404364ef8b8a92155 (patch)
treefca17dc57dac2513cbbcd5ddb4cdacb35255f26e /experiments/rain_ep_layer_adapter_smoke.py
parent7ea36581a0469532630e891a032622c1f21a914b (diff)
feat: make Rain hardware bias beta-independent
Diffstat (limited to 'experiments/rain_ep_layer_adapter_smoke.py')
-rw-r--r--experiments/rain_ep_layer_adapter_smoke.py20
1 files changed, 20 insertions, 0 deletions
diff --git a/experiments/rain_ep_layer_adapter_smoke.py b/experiments/rain_ep_layer_adapter_smoke.py
index 8959217..24c9e0a 100644
--- a/experiments/rain_ep_layer_adapter_smoke.py
+++ b/experiments/rain_ep_layer_adapter_smoke.py
@@ -125,12 +125,32 @@ def main() -> None:
assert innovation.debiaser.neutral_observations == 64
assert constant.debiaser.neutral_observations == 64
assert all(not value.requires_grad for value in innovation_used.values())
+
+ positive_field = RainLayerStateCorrector(
+ mode="raw", bias_ratio=2e-4, bias_normalization="first_state", seed=47)
+ negative_field = RainLayerStateCorrector(
+ mode="raw", bias_ratio=2e-4, bias_normalization="first_state", seed=47)
+ positive_field.apply(first, second, layer_names)
+ negative_second = {
+ name: value - clean_difference[name] for name, value in first.items()
+ }
+ negative_field.apply(first, negative_second, layer_names)
+ positive_bias = positive_field._measure(
+ [first[name] for name in layer_names])[1]
+ negative_bias = negative_field._measure(
+ [first[name] for name in layer_names])[1]
+ beta_independent_bias = all(
+ torch.equal(positive, negative)
+ for positive, negative in zip(positive_bias, negative_bias)
+ )
+ assert beta_independent_bias
print({
"oracle_parameter_gradient_relative_error": oracle_relative_error,
"innovation_heldout_state_residual_rms": innovation_error,
"constant_heldout_state_residual_rms": constant_error,
"matched_neutral_observations": 64,
"extra_equilibrium_phases": 0,
+ "first_state_bias_bitwise_independent_of_beta_sign": True,
"requires_grad": False,
})