summaryrefslogtreecommitdiff
path: root/experiments/rain_ep_bias_train.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_bias_train.py
parent7ea36581a0469532630e891a032622c1f21a914b (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.py14
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,