summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--experiments/rain_ep_bias_train.py14
-rw-r--r--experiments/rain_ep_layer_adapter_smoke.py20
-rw-r--r--sdil/rain_ep_adapter.py12
3 files changed, 42 insertions, 4 deletions
diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py
index 6528be6..eb181a3 100644
--- a/experiments/rain_ep_bias_train.py
+++ b/experiments/rain_ep_bias_train.py
@@ -49,6 +49,10 @@ def parse_args() -> argparse.Namespace:
help="training steps per neutral update; zero freezes after calibration")
parser.add_argument("--calibration-batches", type=int, default=0)
parser.add_argument("--layer-calibration-steps", type=int, default=1)
+ parser.add_argument(
+ "--layer-bias-normalization",
+ choices=("clean_difference", "first_state"),
+ default="clean_difference")
parser.add_argument("--epochs", type=int, default=2)
parser.add_argument("--train-limit", type=int, default=2048)
parser.add_argument("--test-limit", type=int, default=1024)
@@ -187,6 +191,7 @@ def main() -> None:
bias_ratio=args.bias_ratio,
predictor_rate=args.predictor_rate,
calibration_steps=args.layer_calibration_steps,
+ bias_normalization=args.layer_bias_normalization,
seed=args.seed + 1729)
attach_layer_to_rain_estimator(estimator, corrector)
@@ -290,6 +295,7 @@ def main() -> None:
"predictor_rate": args.predictor_rate,
"neutral_cadence": args.neutral_cadence,
"layer_calibration_steps": args.layer_calibration_steps,
+ "layer_bias_normalization": args.layer_bias_normalization,
"calibration_batches": args.calibration_batches,
"calibration_observations": calibration_observations,
"calibration_seconds": calibration_seconds,
@@ -299,8 +305,12 @@ def main() -> None:
"existing_first_EP_phase"
if args.adapter == "layer" else "separate_free_equilibrium"),
"bias_ratio_normalization": (
- "experimenter_initial_clean_layer_state_difference_rms"
- if args.adapter == "layer" else "initial_local_parameter_state_rms"),
+ (
+ "initial_free_layer_state_rms"
+ if args.layer_bias_normalization == "first_state"
+ else "experimenter_initial_clean_layer_state_difference_rms"
+ ) if args.adapter == "layer"
+ else "initial_local_parameter_state_rms"),
"bias_ratio_normalization_visible_to_predictor": False,
"epochs": args.epochs,
"train_limit": args.train_limit,
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,
})
diff --git a/sdil/rain_ep_adapter.py b/sdil/rain_ep_adapter.py
index 73cd70d..ecdb5e5 100644
--- a/sdil/rain_ep_adapter.py
+++ b/sdil/rain_ep_adapter.py
@@ -234,6 +234,7 @@ class RainLayerStateCorrector:
bias_ratio: float,
predictor_rate: float = 0.2,
calibration_steps: int = 1,
+ bias_normalization: str = "clean_difference",
seed: int = 1729,
) -> None:
if mode not in self.MODES:
@@ -242,10 +243,13 @@ class RainLayerStateCorrector:
raise ValueError("bias ratio must be nonnegative")
if calibration_steps < 0:
raise ValueError("calibration steps must be nonnegative")
+ if bias_normalization not in {"clean_difference", "first_state"}:
+ raise ValueError("unrecognized layer-bias normalization")
self.mode = mode
self.bias_ratio = bias_ratio
self.predictor_rate = predictor_rate
self.calibration_steps = calibration_steps
+ self.bias_normalization = bias_normalization
self.seed = seed
self.steps = 0
self.last_diagnostics: dict[str, float | int] = {}
@@ -275,8 +279,12 @@ class RainLayerStateCorrector:
basis = torch.tanh(first / scale)
source = offset.unsqueeze(0) + slope.unsqueeze(0) * basis
source_rms = source.square().mean().sqrt().clamp_min(1e-30)
- clean_rms = clean.square().mean().sqrt()
- gain = self.bias_ratio * clean_rms / source_rms
+ reference_rms = (
+ clean.square().mean().sqrt()
+ if self.bias_normalization == "clean_difference"
+ else first.square().mean().sqrt()
+ )
+ gain = self.bias_ratio * reference_rms / source_rms
self._metadata.append(
(scale, gain * offset, gain * slope)
)