diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-06 17:19:02 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-06 17:19:02 -0500 |
| commit | 25f9b58f2fe478b7a4c728f404364ef8b8a92155 (patch) | |
| tree | fca17dc57dac2513cbbcd5ddb4cdacb35255f26e /experiments/rain_ep_bias_train.py | |
| parent | 7ea36581a0469532630e891a032622c1f21a914b (diff) | |
feat: make Rain hardware bias beta-independent
Diffstat (limited to 'experiments/rain_ep_bias_train.py')
| -rw-r--r-- | experiments/rain_ep_bias_train.py | 14 |
1 files changed, 12 insertions, 2 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, |
