summaryrefslogtreecommitdiff
path: root/experiments/rain_ep_bias_train.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-07 13:18:15 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-07 13:18:15 -0500
commit65f617eaf07a87c84475ad0010dea844422163e1 (patch)
treed5ddf759838081cf6a9e5167f3ed8ed083cdcef6 /experiments/rain_ep_bias_train.py
parent0be7f2d8b71343084da3cd4a97c714b7f74ffc3c (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.py14
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,