summaryrefslogtreecommitdiff
path: root/experiments/rain_ep_bias_train.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/rain_ep_bias_train.py')
-rw-r--r--experiments/rain_ep_bias_train.py60
1 files changed, 49 insertions, 11 deletions
diff --git a/experiments/rain_ep_bias_train.py b/experiments/rain_ep_bias_train.py
index db6f668..79e9d6b 100644
--- a/experiments/rain_ep_bias_train.py
+++ b/experiments/rain_ep_bias_train.py
@@ -19,8 +19,10 @@ ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from sdil.rain_ep_adapter import ( # noqa: E402
+ DillavouUpdateCorrector,
RainGradientCorrector,
RainLayerStateCorrector,
+ attach_dillavou_to_rain_estimator,
attach_layer_to_rain_estimator,
attach_to_rain_estimator,
observe_rain_neutral,
@@ -35,18 +37,23 @@ 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")
+ "--adapter", choices=("parameter", "layer", "dillavou"),
+ default="parameter")
parser.add_argument(
"--network-protocol", choices=("conv28_screen", "comparative32"),
default="conv28_screen")
parser.add_argument(
"--beta-policy",
- choices=("fixed_positive", "fixed_negative", "random_sign"),
+ choices=("fixed_positive", "fixed_negative", "random_sign", "centered"),
default="fixed_positive")
parser.add_argument("--beta-seed", type=int, default=7100)
+ parser.add_argument("--beta-value", type=float, default=0.25)
parser.add_argument(
"--mode", choices=sorted(RainGradientCorrector.MODES), required=True)
parser.add_argument("--bias-ratio", type=float, default=0.5)
+ parser.add_argument(
+ "--dillavou-drift-ratio", type=float, default=0.0,
+ help="zero is the exact fixed update-offset model from Dillavou et al.")
parser.add_argument("--predictor-rate", type=float, default=0.1)
parser.add_argument(
"--neutral-cadence", type=int, default=1,
@@ -103,11 +110,15 @@ def main() -> None:
args = parse_args()
if args.calibration_batches < 0:
raise ValueError("calibration batches must be nonnegative")
+ if args.beta_value <= 0.0:
+ raise ValueError("beta value must be positive")
if args.beta_policy != "fixed_positive" and not (
- args.adapter == "layer" and args.mode == "raw"
+ (args.adapter == "layer" and args.mode == "raw")
+ or args.adapter == "dillavou"
):
raise ValueError(
- "non-positive beta policies are raw layer-measurement baselines")
+ "non-positive beta policies require a raw layer baseline or "
+ "the post-estimator Dillavou adapter")
author_root = args.author_root.resolve()
author_revision = revision(author_root)
if author_revision != PINNED_REVISION:
@@ -185,8 +196,12 @@ def main() -> None:
estimator = EquilibriumProp(
energy.params(), energy.layers(), augmented, cost,
training_minimizer)
- estimator.variant = "positive"
- estimator.nudging = 0.25
+ estimator.variant = (
+ "negative" if args.beta_policy == "fixed_negative"
+ else "centered" if args.beta_policy == "centered"
+ else "positive"
+ )
+ estimator.nudging = args.beta_value
if args.adapter == "parameter":
corrector = RainGradientCorrector(
mode=args.mode,
@@ -195,7 +210,7 @@ def main() -> None:
neutral_cadence=args.neutral_cadence,
seed=args.seed + 1729)
attach_to_rain_estimator(estimator, corrector)
- else:
+ elif args.adapter == "layer":
if args.calibration_batches:
raise ValueError(
"layer adapter calibrates inside existing free phases; "
@@ -208,6 +223,20 @@ def main() -> None:
bias_normalization=args.layer_bias_normalization,
seed=args.seed + 1729)
attach_layer_to_rain_estimator(estimator, corrector)
+ else:
+ if args.calibration_batches:
+ raise ValueError(
+ "Dillavou calibration probes the local update circuit and "
+ "does not require equilibrium batches")
+ 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,
+ seed=args.seed + 1729,
+ )
+ attach_dillavou_to_rain_estimator(estimator, corrector)
inference_minimizer = FixedPointMinimizer(
energy, network.free_layers())
@@ -261,8 +290,9 @@ def main() -> None:
if args.beta_policy == "fixed_positive":
beta_sign_counts["positive"] += 1
elif args.beta_policy == "fixed_negative":
- estimator._first_nudging = 0.0
- estimator._second_nudging = -estimator.nudging
+ beta_sign_counts["negative"] += 1
+ elif args.beta_policy == "centered":
+ beta_sign_counts["positive"] += 1
beta_sign_counts["negative"] += 1
else:
sign = 1 if int(torch.randint(
@@ -316,9 +346,11 @@ def main() -> None:
"algorithm": "equilibrium propagation",
"beta_policy": args.beta_policy,
"beta_seed": args.beta_seed,
+ "beta_value": args.beta_value,
"adapter": args.adapter,
"mode": args.mode,
"bias_ratio": args.bias_ratio,
+ "dillavou_drift_ratio": args.dillavou_drift_ratio,
"predictor_rate": args.predictor_rate,
"neutral_cadence": args.neutral_cadence,
"layer_calibration_steps": args.layer_calibration_steps,
@@ -327,16 +359,22 @@ def main() -> None:
"calibration_observations": calibration_observations,
"calibration_seconds": calibration_seconds,
"extra_equilibrium_phases_for_predictor": (
- 0 if args.adapter == "layer" else args.calibration_batches),
+ 0 if args.adapter in {"layer", "dillavou"}
+ else args.calibration_batches),
"predictor_neutral_source": (
"existing_first_EP_phase"
- if args.adapter == "layer" else "separate_free_equilibrium"),
+ if args.adapter == "layer"
+ else "instruction_off_local_update_probe"
+ if args.adapter == "dillavou"
+ else "separate_free_equilibrium"),
"bias_ratio_normalization": (
(
"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_clean_local_update_rms_for_simulation_only"
+ if args.adapter == "dillavou"
else "initial_local_parameter_state_rms"),
"bias_ratio_normalization_visible_to_predictor": False,
"epochs": args.epochs,