summaryrefslogtreecommitdiff
path: root/experiments/rain_ep_bias_train.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 16:56:11 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-08-06 16:56:11 -0500
commitddb034680ddc46d848278ca89d9bb97f59cbe3d3 (patch)
tree9e53bbac894717f172948607a63afd26327b2262 /experiments/rain_ep_bias_train.py
parentd686f83aa35c2f1d87adfb5705ff02abe1901e9d (diff)
feat: add BP-free Rain layer-state adapter
Diffstat (limited to 'experiments/rain_ep_bias_train.py')
-rw-r--r--experiments/rain_ep_bias_train.py36
1 files changed, 29 insertions, 7 deletions
diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py
index defd973..fa8a243 100644
--- a/experiments/rain_ep_bias_train.py
+++ b/experiments/rain_ep_bias_train.py
@@ -19,6 +19,8 @@ sys.path.insert(0, str(ROOT))
from sdil.rain_ep_adapter import ( # noqa: E402
RainGradientCorrector,
+ RainLayerStateCorrector,
+ attach_layer_to_rain_estimator,
attach_to_rain_estimator,
observe_rain_neutral,
)
@@ -32,6 +34,8 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--author-root", type=Path, required=True)
parser.add_argument("--device", default="cuda")
parser.add_argument(
+ "--adapter", choices=("parameter", "layer"), default="parameter")
+ parser.add_argument(
"--mode", choices=sorted(RainGradientCorrector.MODES), required=True)
parser.add_argument("--bias-ratio", type=float, default=0.5)
parser.add_argument("--predictor-rate", type=float, default=0.1)
@@ -39,6 +43,7 @@ def parse_args() -> argparse.Namespace:
"--neutral-cadence", type=int, default=1,
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("--epochs", type=int, default=2)
parser.add_argument("--train-limit", type=int, default=2048)
parser.add_argument("--test-limit", type=int, default=1024)
@@ -139,13 +144,26 @@ def main() -> None:
training_minimizer)
estimator.variant = "positive"
estimator.nudging = 0.25
- corrector = RainGradientCorrector(
- mode=args.mode,
- bias_ratio=args.bias_ratio,
- predictor_rate=args.predictor_rate,
- neutral_cadence=args.neutral_cadence,
- seed=args.seed + 1729)
- attach_to_rain_estimator(estimator, corrector)
+ if args.adapter == "parameter":
+ corrector = RainGradientCorrector(
+ mode=args.mode,
+ bias_ratio=args.bias_ratio,
+ predictor_rate=args.predictor_rate,
+ neutral_cadence=args.neutral_cadence,
+ seed=args.seed + 1729)
+ attach_to_rain_estimator(estimator, corrector)
+ else:
+ if args.calibration_batches:
+ raise ValueError(
+ "layer adapter calibrates inside existing free phases; "
+ "external calibration batches must be zero")
+ corrector = RainLayerStateCorrector(
+ mode=args.mode,
+ bias_ratio=args.bias_ratio,
+ predictor_rate=args.predictor_rate,
+ calibration_steps=args.layer_calibration_steps,
+ seed=args.seed + 1729)
+ attach_layer_to_rain_estimator(estimator, corrector)
inference_minimizer = FixedPointMinimizer(
energy, network.free_layers())
@@ -164,6 +182,8 @@ def main() -> None:
calibration_start = time.time()
calibration_observations = 0
if args.calibration_batches:
+ if args.adapter != "parameter":
+ raise AssertionError("layer calibration was not rejected above")
if args.mode not in {"constant", "innovation"}:
raise ValueError(
"precalibration is defined only for constant or innovation mode")
@@ -227,10 +247,12 @@ def main() -> None:
"dataset": "FashionMNIST",
"network": "author ConvHopfieldEnergy28 32-64-10",
"algorithm": "positive equilibrium propagation",
+ "adapter": args.adapter,
"mode": args.mode,
"bias_ratio": args.bias_ratio,
"predictor_rate": args.predictor_rate,
"neutral_cadence": args.neutral_cadence,
+ "layer_calibration_steps": args.layer_calibration_steps,
"calibration_batches": args.calibration_batches,
"calibration_observations": calibration_observations,
"calibration_seconds": calibration_seconds,