summaryrefslogtreecommitdiff
path: root/experiments/conv_run.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:07:38 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 06:07:38 -0500
commitbcd845b60dc85e4f46bf1ab343405c1e40ed0860 (patch)
tree37cb160e8e89e393c13d91a47ab2f9b35c2c24f5 /experiments/conv_run.py
parentb323ddb1bd0c11e313c452a63707c01b4254b71a (diff)
oral-a: add audited convolutional runner
Diffstat (limited to 'experiments/conv_run.py')
-rw-r--r--experiments/conv_run.py454
1 files changed, 454 insertions, 0 deletions
diff --git a/experiments/conv_run.py b/experiments/conv_run.py
new file mode 100644
index 0000000..48eebd1
--- /dev/null
+++ b/experiments/conv_run.py
@@ -0,0 +1,454 @@
+#!/usr/bin/env python3
+"""Audited BP/DFA/SDIL/direct-NP runner for CIFAR residual networks."""
+import argparse
+import hashlib
+import json
+import math
+import os
+import subprocess
+import sys
+import time
+
+import torch
+
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+from sdil.conv import (CIFARLocalResNet, CIFARSDILResNet, ConvSDILConfig,
+ conv_alignment_report, conv_apical_calibration_step,
+ conv_local_step, evaluate_conv)
+from sdil.data import DATA_DIR, get_cifar_image_splits
+
+
+ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+
+
+def provenance():
+ def run(command):
+ return subprocess.run(
+ command, cwd=ROOT, check=True, capture_output=True, text=True).stdout.strip()
+ return {
+ "git_commit": run(["git", "rev-parse", "HEAD"]),
+ "git_tracked_dirty": bool(run(
+ ["git", "status", "--porcelain", "--untracked-files=no"])),
+ }
+
+
+def file_sha256(path):
+ digest = hashlib.sha256()
+ with open(path, "rb") as handle:
+ for chunk in iter(lambda: handle.read(1024 * 1024), b""):
+ digest.update(chunk)
+ return {"path": os.path.abspath(path), "bytes": os.path.getsize(path),
+ "sha256": digest.hexdigest()}
+
+
+def cifar_source_records(data_dir):
+ root = os.path.join(data_dir, "cifar-10-batches-py")
+ names = [f"data_batch_{index}" for index in range(1, 6)] + ["test_batch"]
+ return [file_sha256(os.path.join(root, name)) for name in names]
+
+
+def hardware_report(device):
+ report = {
+ "device": str(device),
+ "torch_version": torch.__version__,
+ "cuda_visible_devices": os.environ.get("CUDA_VISIBLE_DEVICES"),
+ }
+ if str(device).startswith("cuda"):
+ index = torch.device(device)
+ props = torch.cuda.get_device_properties(index)
+ report.update({
+ "cuda_device_name": props.name,
+ "device_total_memory_bytes": props.total_memory,
+ "peak_memory_allocated_bytes": torch.cuda.max_memory_allocated(index),
+ "peak_memory_reserved_bytes": torch.cuda.max_memory_reserved(index),
+ })
+ else:
+ report.update({
+ "cuda_device_name": None,
+ "device_total_memory_bytes": None,
+ "peak_memory_allocated_bytes": None,
+ "peak_memory_reserved_bytes": None,
+ })
+ return report
+
+
+def sync(device):
+ if str(device).startswith("cuda"):
+ torch.cuda.synchronize(torch.device(device))
+
+
+def scheduled_lr(base, epoch, args):
+ if args.warmup_epochs and epoch < args.warmup_epochs:
+ return base * (epoch + 1) / args.warmup_epochs
+ if args.lr_schedule == "constant":
+ return base
+ if args.lr_schedule == "step":
+ milestones = [int(value) for value in args.lr_milestones.split(",") if value]
+ return base * args.lr_gamma ** sum(epoch >= milestone for milestone in milestones)
+ progress = ((epoch - args.warmup_epochs)
+ / max(1, args.epochs - args.warmup_epochs))
+ return base * 0.5 * (1.0 + math.cos(math.pi * progress))
+
+
+def build(args):
+ common = dict(
+ depth=args.depth, base_width=args.width, n_classes=10,
+ device=args.device, seed=args.seed, weight_scale=args.weight_scale,
+ residual_scale=args.residual_scale)
+ if args.mode == "bp":
+ return CIFARLocalResNet(**common), None
+ net = CIFARSDILResNet(
+ **common, a_scale=args.a_scale, apical_seed=args.apical_seed)
+ config = ConvSDILConfig(
+ eta=args.lr, eta_output=args.output_lr, eta_A=args.eta_A,
+ eta_P=args.eta_P, momentum=args.momentum, weight_decay=args.weight_decay,
+ learn_A=args.mode == "sdil", learn_P=bool(args.learn_P),
+ use_residual=bool(args.use_residual), nuisance_scale=args.nuisance_scale,
+ pert_sigma=args.pert_sigma, pert_every=args.pert_every,
+ pert_directions=args.pert_directions,
+ direct_node_perturbation=args.mode == "nodepert")
+ config.validate()
+ return net, config
+
+
+def work_report(net, mode, counters):
+ forward_macs = net.forward_macs_per_example
+ apical_macs = getattr(net, "apical_macs_per_example", 0)
+ normal_forward = counters["ordinary_examples"] * forward_macs
+ warmup_forward = ((counters["predictor_warmup_examples"]
+ + counters["apical_warmup_examples"]) * forward_macs)
+ calibration_forward = counters["perturbation_forward_examples"] * forward_macs
+ if mode == "bp":
+ bp_reverse = 2 * counters["ordinary_examples"] * forward_macs
+ local_correlation = 0
+ apical_inference = 0
+ apical_regression = 0
+ else:
+ bp_reverse = 0
+ local_correlation = counters["ordinary_examples"] * forward_macs
+ apical_inference = ((counters["ordinary_examples"]
+ + counters["apical_warmup_examples"]) * apical_macs)
+ apical_regression = counters["calibration_event_examples"] * apical_macs
+ components = {
+ "ordinary_forward_macs": normal_forward,
+ "warmup_clean_forward_macs": warmup_forward,
+ "perturbation_forward_macs": calibration_forward,
+ "bp_reverse_macs_estimate": bp_reverse,
+ "local_weight_correlation_macs": local_correlation,
+ "apical_projection_macs": apical_inference,
+ "apical_regression_macs": apical_regression,
+ }
+ return {
+ "forward_macs_per_example": forward_macs,
+ "apical_macs_per_example": apical_macs,
+ "components": components,
+ "total_macs_estimate": sum(components.values()),
+ "total_clean_forward_examples": (
+ counters["ordinary_examples"] + counters["predictor_warmup_examples"]
+ + counters["apical_warmup_examples"]),
+ "total_forward_equivalent_examples": (
+ counters["ordinary_examples"] + counters["predictor_warmup_examples"]
+ + counters["apical_warmup_examples"]
+ + counters["perturbation_forward_examples"]),
+ "logical_batch_loss_queries": counters["logical_batch_loss_queries"],
+ "scalar_loss_evaluations": counters["scalar_loss_evaluations"],
+ "definition": (
+ "multiply-accumulates in conv/linear maps; one local weight correlation "
+ "equals one forward-weight MAC count; BP reverse is estimated as one "
+ "weight-gradient plus one activation-gradient convolution per forward "
+ "convolution; elementwise nonlinearities and optimizer arithmetic excluded"),
+ }
+
+
+def run(args):
+ if args.eval_split == "test" and args.eval_every:
+ raise ValueError("test protocols must use --eval_every 0 (one final evaluation)")
+ if args.mode != "sdil" and (args.a_warmup_steps or args.learn_P):
+ raise ValueError("apical/predictor warmup is restricted to SDIL")
+ torch.manual_seed(args.seed)
+ if str(args.device).startswith("cuda"):
+ if not torch.cuda.is_available():
+ raise RuntimeError("CUDA device requested but CUDA is unavailable")
+ torch.cuda.manual_seed_all(args.seed)
+
+ train, validation, test, input_shape, n_out, split = get_cifar_image_splits(
+ batch_size=args.batch_size, data_dir=args.data_dir, device=args.device,
+ train_limit=args.train_limit or None, val_examples=args.val_examples,
+ split_seed=args.split_seed, loader_seed=args.loader_seed,
+ augment_train=bool(args.augment_train))
+ evaluation = validation if args.eval_split == "validation" else test
+ if evaluation is None:
+ raise ValueError("validation evaluation requires a nonzero validation split")
+ split["evaluation_split"] = args.eval_split
+ split["cifar_source_files"] = cifar_source_records(args.data_dir)
+ net, config = build(args)
+ if input_shape != (3, 32, 32) or n_out != 10:
+ raise AssertionError("unexpected CIFAR dimensions")
+ if str(args.device).startswith("cuda"):
+ torch.cuda.reset_peak_memory_stats(torch.device(args.device))
+ perturb_generator = torch.Generator(device=torch.device(args.device)).manual_seed(
+ args.perturb_seed)
+ warmup_generator = torch.Generator(device=torch.device(args.device)).manual_seed(
+ args.perturb_seed + 1)
+
+ counters = {
+ "ordinary_examples": 0,
+ "predictor_warmup_examples": 0,
+ "apical_warmup_examples": 0,
+ "perturbation_forward_examples": 0,
+ "calibration_event_examples": 0,
+ "logical_batch_loss_queries": 0,
+ "scalar_loss_evaluations": 0,
+ "perturbation_events": 0,
+ }
+ log = {
+ "schema_version": 1,
+ "protocol_family": "oral_a_cifar_local_resnet_development",
+ "args": vars(args),
+ "provenance": provenance(),
+ "split": split,
+ "architecture": {
+ "family": "normalization-free CIFAR 6n+2 ResNet, option-A shortcuts",
+ "depth": net.depth,
+ "blocks_per_stage": net.blocks_per_stage,
+ "base_width": net.base_width,
+ "residual_scale": net.residual_scale,
+ "hidden_shapes": net.hidden_shapes,
+ "forward_parameters": net.n_forward_parameters,
+ "adaptive_apical_parameters": getattr(net, "n_apical_parameters", 0),
+ "fixed_traffic_coefficients": getattr(
+ net, "n_fixed_traffic_coefficients", 0),
+ },
+ "epochs": [],
+ }
+
+ loader_state = train.g.get_state().clone()
+ if config is not None and config.learn_P and args.predictor_warmup_steps:
+ iterator = iter(train)
+ for _ in range(args.predictor_warmup_steps):
+ try:
+ x, _ = next(iterator)
+ except StopIteration:
+ iterator = iter(train)
+ x, _ = next(iterator)
+ forward = net.forward(x)
+ net.predictor_step(
+ forward["hiddens"], config.eta_P, config.nuisance_scale)
+ counters["predictor_warmup_examples"] += x.shape[0]
+ train.g.set_state(loader_state)
+
+ if config is not None and args.a_warmup_steps:
+ iterator = iter(train)
+ warmup_metrics = []
+ for _ in range(args.a_warmup_steps):
+ try:
+ x, y = next(iterator)
+ except StopIteration:
+ iterator = iter(train)
+ x, y = next(iterator)
+ _, metric = conv_apical_calibration_step(
+ net, x, y, config, generator=warmup_generator)
+ warmup_metrics.append(metric)
+ batch = x.shape[0]
+ counters["apical_warmup_examples"] += batch
+ counters["perturbation_forward_examples"] += (
+ 2 * config.pert_directions * batch)
+ counters["calibration_event_examples"] += batch
+ counters["logical_batch_loss_queries"] += 2 * config.pert_directions
+ counters["scalar_loss_evaluations"] += 2 * config.pert_directions * batch
+ counters["perturbation_events"] += 1
+ train.g.set_state(loader_state)
+ log["apical_warmup"] = {
+ "steps": args.a_warmup_steps,
+ "last": warmup_metrics[-1],
+ }
+
+ step = 0
+ train_wall = 0.0
+ eval_wall = 0.0
+ validation_evaluations = 0
+ test_evaluations = 0
+ sync(args.device)
+ total_start = time.time()
+ for epoch in range(args.epochs):
+ lr = scheduled_lr(args.lr, epoch, args)
+ output_lr = (scheduled_lr(args.output_lr, epoch, args)
+ if args.output_lr is not None else lr)
+ if config is not None:
+ config.eta = lr
+ config.eta_output = output_lr
+ sync(args.device)
+ epoch_start = time.time()
+ loss_sum = 0.0
+ examples = 0
+ calibration_metrics = []
+ for x, y in train:
+ batch = x.shape[0]
+ if args.mode == "bp":
+ loss = net.bp_step(
+ x, y, lr, momentum=args.momentum,
+ weight_decay=args.weight_decay)
+ did_perturb = False
+ else:
+ result = conv_local_step(
+ net, x, y, config, step, generator=perturb_generator)
+ loss = result["loss"]
+ did_perturb = result["did_perturb"]
+ if result["calibration"] is not None:
+ calibration_metrics.append(result["calibration"])
+ loss_sum += loss * batch
+ examples += batch
+ counters["ordinary_examples"] += batch
+ if did_perturb:
+ counters["perturbation_forward_examples"] += (
+ 2 * config.pert_directions * batch)
+ counters["calibration_event_examples"] += batch
+ counters["logical_batch_loss_queries"] += 2 * config.pert_directions
+ counters["scalar_loss_evaluations"] += (
+ 2 * config.pert_directions * batch)
+ counters["perturbation_events"] += 1
+ step += 1
+ if args.max_steps and step >= args.max_steps:
+ break
+ sync(args.device)
+ train_wall += time.time() - epoch_start
+ record = {
+ "epoch": epoch + 1,
+ "step": step,
+ "lr": lr,
+ "output_lr": output_lr,
+ "train_loss": loss_sum / examples,
+ "train_examples": examples,
+ }
+ if calibration_metrics:
+ record["calibration"] = {
+ key: sum(value[key] for value in calibration_metrics)
+ / len(calibration_metrics)
+ for key in calibration_metrics[0]
+ }
+ if args.eval_every and (epoch + 1) % args.eval_every == 0:
+ sync(args.device)
+ eval_start = time.time()
+ accuracy, eval_loss = evaluate_conv(net, evaluation)
+ sync(args.device)
+ eval_wall += time.time() - eval_start
+ record.update({"eval_split": args.eval_split,
+ "eval_accuracy": accuracy, "eval_loss": eval_loss})
+ validation_evaluations += int(args.eval_split == "validation")
+ test_evaluations += int(args.eval_split == "test")
+ log["epochs"].append(record)
+ metric = (f" eval={record['eval_accuracy']:.4f}"
+ if "eval_accuracy" in record else "")
+ print(f"epoch={epoch + 1} step={step} lr={lr:.6g} "
+ f"loss={record['train_loss']:.5f}{metric}", flush=True)
+ if args.max_steps and step >= args.max_steps:
+ break
+
+ sync(args.device)
+ final_eval_start = time.time()
+ final_accuracy, final_loss = evaluate_conv(net, evaluation)
+ sync(args.device)
+ eval_wall += time.time() - final_eval_start
+ validation_evaluations += int(args.eval_split == "validation")
+ test_evaluations += int(args.eval_split == "test")
+ total_wall = time.time() - total_start
+ # Snapshot the training/evaluation peak before optional autograd-only
+ # alignment diagnostics, which can otherwise make local methods look more
+ # memory hungry merely because BP does not need that post-hoc probe.
+ training_hardware = hardware_report(args.device)
+
+ diagnostics = None
+ if args.alignment_probe and args.mode != "bp":
+ probe = min(args.alignment_probe, train.x.shape[0])
+ sync(args.device)
+ diagnostic_start = time.time()
+ diagnostics = conv_alignment_report(
+ net, train.x[:probe], train.y[:probe], config)
+ sync(args.device)
+ diagnostics["wall_s"] = time.time() - diagnostic_start
+ values = diagnostics["teaching_negative_gradient_cosine"]
+ early = max(1, len(values) // 3)
+ diagnostics["early_third_mean"] = sum(values[:early]) / early
+
+ log.update({
+ "counters": counters,
+ "work": work_report(net, args.mode, counters),
+ "hardware": training_hardware,
+ "timing": {
+ "train_wall_s": train_wall,
+ "evaluation_wall_s": eval_wall,
+ "total_timed_wall_s": total_wall,
+ "timing_excludes_data_loading_hashing_and_model_construction": True,
+ },
+ "evaluation_protocol": {
+ "validation_evaluations": validation_evaluations,
+ "test_evaluations": test_evaluations,
+ "test_used_for_selection": False,
+ },
+ "diagnostics": diagnostics,
+ "final": {
+ "evaluation_split": args.eval_split,
+ "accuracy": final_accuracy,
+ "loss": final_loss,
+ "epoch": len(log["epochs"]),
+ "step": step,
+ "finite": math.isfinite(final_loss),
+ },
+ })
+ os.makedirs(os.path.dirname(os.path.abspath(args.out)), exist_ok=True)
+ with open(args.out, "w") as handle:
+ json.dump(log, handle, indent=2, sort_keys=True)
+ handle.write("\n")
+ print(json.dumps({"out": args.out, "final": log["final"],
+ "work": log["work"]}, indent=2))
+
+
+def parse_args():
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--mode", choices=("bp", "dfa", "sdil", "nodepert"), required=True)
+ parser.add_argument("--out", required=True)
+ parser.add_argument("--device", default="cpu")
+ parser.add_argument("--data_dir", default=DATA_DIR)
+ parser.add_argument("--depth", type=int, default=20)
+ parser.add_argument("--width", type=int, default=16)
+ parser.add_argument("--seed", type=int, default=0)
+ parser.add_argument("--loader_seed", type=int, default=0)
+ parser.add_argument("--split_seed", type=int, default=2027)
+ parser.add_argument("--perturb_seed", type=int, default=1000)
+ parser.add_argument("--apical_seed", type=int)
+ parser.add_argument("--batch_size", type=int, default=128)
+ parser.add_argument("--epochs", type=int, default=200)
+ parser.add_argument("--max_steps", type=int, default=0)
+ parser.add_argument("--train_limit", type=int, default=0)
+ parser.add_argument("--val_examples", type=int, default=5000)
+ parser.add_argument("--eval_split", choices=("validation", "test"), default="validation")
+ parser.add_argument("--eval_every", type=int, default=10)
+ parser.add_argument("--augment_train", type=int, choices=(0, 1), default=1)
+ parser.add_argument("--lr", type=float, default=0.1)
+ parser.add_argument("--output_lr", type=float)
+ parser.add_argument("--lr_schedule", choices=("constant", "step", "cosine"),
+ default="cosine")
+ parser.add_argument("--lr_milestones", default="100,150")
+ parser.add_argument("--lr_gamma", type=float, default=0.1)
+ parser.add_argument("--warmup_epochs", type=int, default=5)
+ parser.add_argument("--momentum", type=float, default=0.9)
+ parser.add_argument("--weight_decay", type=float, default=5e-4)
+ parser.add_argument("--weight_scale", type=float, default=1.0)
+ parser.add_argument("--residual_scale", type=float)
+ parser.add_argument("--a_scale", type=float, default=1.0)
+ parser.add_argument("--eta_A", type=float, default=0.01)
+ parser.add_argument("--eta_P", type=float, default=0.01)
+ parser.add_argument("--learn_P", type=int, choices=(0, 1), default=0)
+ parser.add_argument("--use_residual", type=int, choices=(0, 1), default=1)
+ parser.add_argument("--nuisance_scale", type=float, default=0.0)
+ parser.add_argument("--pert_sigma", type=float, default=0.01)
+ parser.add_argument("--pert_every", type=int, default=4)
+ parser.add_argument("--pert_directions", type=int, default=1)
+ parser.add_argument("--predictor_warmup_steps", type=int, default=0)
+ parser.add_argument("--a_warmup_steps", type=int, default=0)
+ parser.add_argument("--alignment_probe", type=int, default=0)
+ return parser.parse_args()
+
+
+if __name__ == "__main__":
+ run(parse_args())