diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-07 13:18:15 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-07 13:18:15 -0500 |
| commit | 65f617eaf07a87c84475ad0010dea844422163e1 (patch) | |
| tree | d5ddf759838081cf6a9e5167f3ed8ed083cdcef6 /experiments/rain_ep_bias_train.py | |
| parent | 0be7f2d8b71343084da3cd4a97c714b7f74ffc3c (diff) | |
feat: transfer released physical drift profile to Rain EP
Diffstat (limited to 'experiments/rain_ep_bias_train.py')
| -rw-r--r-- | experiments/rain_ep_bias_train.py | 14 |
1 files changed, 14 insertions, 0 deletions
diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py index fa62137..2ba3dd0 100644 --- a/experiments/rain_ep_bias_train.py +++ b/experiments/rain_ep_bias_train.py @@ -19,6 +19,7 @@ ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) from sdil.rain_ep_adapter import ( # noqa: E402 + DillavouBiasProfile, DillavouUpdateCorrector, RainGradientCorrector, RainLayerStateCorrector, @@ -55,6 +56,9 @@ def parse_args() -> argparse.Namespace: "--dillavou-drift-ratio", type=float, default=0.0, help="zero is the exact fixed update-offset model from Dillavou et al.") parser.add_argument("--dillavou-calibration-steps", type=int, default=1) + parser.add_argument( + "--dillavou-profile-json", type=Path, + help="released physical state-dependence report; omitted means constant B") parser.add_argument("--predictor-rate", type=float, default=0.1) parser.add_argument( "--neutral-cadence", type=int, default=1, @@ -230,12 +234,18 @@ def main() -> None: raise ValueError( "Dillavou calibration probes the local update circuit and " "does not require equilibrium batches") + empirical_profile = None + if args.dillavou_profile_json is not None: + profile_path = args.dillavou_profile_json.resolve() + empirical_profile = DillavouBiasProfile.from_state_dependence_report( + json.loads(profile_path.read_text()), source=str(profile_path)) corrector = DillavouUpdateCorrector( mode=args.mode, bias_ratio=args.bias_ratio, predictor_rate=args.predictor_rate, neutral_cadence=args.neutral_cadence, drift_ratio=args.dillavou_drift_ratio, + empirical_profile=empirical_profile, seed=args.seed + 1729, ) attach_dillavou_to_rain_estimator(estimator, corrector) @@ -354,6 +364,10 @@ def main() -> None: "bias_ratio": args.bias_ratio, "dillavou_drift_ratio": args.dillavou_drift_ratio, "dillavou_calibration_steps": args.dillavou_calibration_steps, + "dillavou_profile": ( + None if args.adapter != "dillavou" + or corrector.empirical_profile is None + else corrector.empirical_profile.as_dict()), "predictor_rate": args.predictor_rate, "neutral_cadence": args.neutral_cadence, "layer_calibration_steps": args.layer_calibration_steps, |
