#!/usr/bin/env python3 """Validation-only runner for the four state-specific ResNet baselines.""" import argparse import hashlib import json import math import os import subprocess import sys import time import torch import torch.nn.functional as F sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from sdil.conv_crossover import ( CIFARDualPropResNet, CIFAREquilibriumPropResNet, CIFARForwardForwardResNet, CIFARPEPITAResNet, ) from sdil.data import DATA_DIR, get_cifar_image_splits ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) METHODS = ("pepita", "ff", "ep", "dualprop") def git_output(*args): return subprocess.run( ["git", *args], cwd=ROOT, check=True, capture_output=True, text=True).stdout.strip() def 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 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)] names.append("test_batch") return [ { "path": os.path.abspath(os.path.join(root, name)), "bytes": os.path.getsize(os.path.join(root, name)), "sha256": sha256(os.path.join(root, name)), } for name in names ] def sync(device): if str(device).startswith("cuda"): torch.cuda.synchronize(torch.device(device)) 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"): target = torch.device(device) properties = torch.cuda.get_device_properties(target) report.update({ "cuda_device_name": properties.name, "device_total_memory_bytes": properties.total_memory, "peak_memory_allocated_bytes": torch.cuda.max_memory_allocated(target), "peak_memory_reserved_bytes": torch.cuda.max_memory_reserved(target), }) 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 scheduled_rate(base, epoch, args): if args.lr_schedule == "constant": return base if args.lr_schedule == "pepita": return base * (0.1 if epoch >= 60 else 1.0) * ( 0.1 if epoch >= 90 else 1.0) milestones = [ int(value) for value in args.lr_milestones.split(",") if value] return base * args.lr_gamma ** sum( epoch >= milestone for milestone in milestones) def build(args): common = dict( depth=args.depth, base_width=args.width, n_classes=10, device=args.device, dtype=torch.float32, seed=args.seed, residual_scale=1.0, normalization="batchnorm", bn_momentum=0.1, bn_eps=1e-5, ) if args.method == "pepita": return CIFARPEPITAResNet( **common, projection_scale=args.pepita_projection_scale, projection_seed=args.feedback_seed) if args.method == "ff": return CIFARForwardForwardResNet( **common, threshold=args.ff_threshold, learning_rate=args.lr, score_from_layer=args.ff_score_from_layer) if args.method == "ep": return CIFAREquilibriumPropResNet( **common, ep_beta=args.ep_beta, dt=args.ep_dt, free_steps=args.ep_free_steps, nudge_steps=args.ep_nudge_steps, random_beta_sign=True) return CIFARDualPropResNet( **common, alpha=args.dp_alpha, dp_beta=args.dp_beta, inference_passes=args.dp_inference_passes) @torch.no_grad() def evaluate(net, method, loader): correct = 0 total = 0 loss_sum = 0.0 for image, labels in loader: if method == "ff": scores = net.ff_candidate_scores(image) prediction = scores.argmax(dim=1) loss = image.new_tensor(float("nan")) elif method == "ep": one_hot = F.one_hot(labels, net.n_classes).to(image.dtype) states = net.ep_settle( image, one_hot, beta=0.0, steps=net.ep_free_steps) scores = states[-1] prediction = scores.argmax(dim=1) loss = ( F.mse_loss(scores, one_hot, reduction="sum") / net.n_classes) else: scores = net.forward( image, training=False, update_stats=False)["logits"] prediction = scores.argmax(dim=1) loss = F.cross_entropy(scores, labels, reduction="sum") correct += int((prediction == labels).sum()) if torch.isfinite(loss): loss_sum += float(loss) total += labels.numel() return { "accuracy": correct / total, "loss": None if method == "ff" else loss_sum / total, "examples": total, } def work_report(net, args, ordinary_examples, validation_examples, completed_epochs): layers = net.n_hidden + 1 if args.method == "ff": presentations = 2 * ordinary_examples feedforward_examples = 2 * ordinary_examples candidate_examples = 10 * validation_examples relaxation_examples = 0 else: presentations = ( 2 * ordinary_examples if args.method == "pepita" else ordinary_examples) feedforward_multiplier = { "pepita": 3, "ep": 1, "dualprop": 1, }[args.method] feedforward_examples = ( feedforward_multiplier * ordinary_examples + validation_examples) candidate_examples = 0 relaxation_multiplier = { "pepita": 0, "ep": args.ep_free_steps + args.ep_nudge_steps, "dualprop": args.dp_inference_passes, }[args.method] relaxation_examples = relaxation_multiplier * ordinary_examples if args.method == "ep": relaxation_examples += ( args.ep_free_steps * validation_examples) local_vjp_examples = ( relaxation_examples if args.method in ("ep", "dualprop") else 0 ) local_target_backward_examples = ( 2 * ordinary_examples if args.method == "ff" else 0 ) return { "forward_parameter_count": net.n_forward_parameters, "forward_macs_per_example": net.forward_macs_per_example, "num_trainable_layers": layers, "ordinary_training_examples": ordinary_examples, "ordinary_validation_examples": validation_examples, "training_example_presentations": presentations, "feedforward_example_passes": feedforward_examples, "relaxation_example_passes": relaxation_examples, "candidate_label_evaluation_presentations": candidate_examples, "logical_task_loss_queries": 0, "local_vjp_example_evaluations": local_vjp_examples, "local_target_backward_example_evaluations": local_target_backward_examples, "completed_global_epochs": completed_epochs, } def run(args): if args.eval_split != "validation": raise ValueError("formal ResNet crossover never evaluates test") if os.path.exists(args.out): raise FileExistsError(f"refusing to overwrite {args.out}") torch.manual_seed(args.seed) if str(args.device).startswith("cuda"): if not torch.cuda.is_available(): raise RuntimeError("CUDA requested but unavailable") torch.cuda.manual_seed_all(args.seed) train, validation, _, input_shape, classes, 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))) if validation is None or input_shape != (3, 32, 32) or classes != 10: raise AssertionError("invalid CIFAR validation setup") split["cifar_source_files"] = cifar_source_records(args.data_dir) net = build(args) if str(args.device).startswith("cuda"): torch.cuda.reset_peak_memory_stats(torch.device(args.device)) provenance = { "git_commit": git_output("rev-parse", "HEAD"), "git_tracked_dirty": bool(git_output( "status", "--porcelain", "--untracked-files=no")), } started = time.time() ordinary_examples = 0 validation_examples = 0 nonfinite_step = None step = 0 epochs = [] layers = [] ep_generator = torch.Generator(device=torch.device(args.device)) ep_generator.manual_seed(args.feedback_seed) if args.method == "ff": for layer_index in range(net.ff_num_layers): layer_record = {"layer": layer_index, "epochs": []} for epoch in range(args.epochs): rate = scheduled_rate(args.lr, epoch, args) loss_sum = 0.0 examples = 0 metrics_sum = { "positive_goodness": 0.0, "negative_goodness": 0.0, "pair_accuracy": 0.0, } sync(args.device) epoch_started = time.time() for image, labels in train: metrics = net.ff_train_layer( layer_index, image, labels, learning_rate=rate) batch = labels.numel() loss_sum += metrics["loss"] * batch for key in metrics_sum: metrics_sum[key] += metrics[key] * batch examples += batch ordinary_examples += batch step += 1 if not math.isfinite(metrics["loss"]): nonfinite_step = step break if args.max_steps and step >= args.max_steps: break sync(args.device) row = { "epoch": epoch + 1, "lr": rate, "loss": loss_sum / examples, "examples": examples, "runtime_seconds": time.time() - epoch_started, **{ key: value / examples for key, value in metrics_sum.items() }, } layer_record["epochs"].append(row) print( f"layer={layer_index + 1}/{net.ff_num_layers} " f"epoch={epoch + 1}/{args.epochs} " f"loss={row['loss']:.6g}", flush=True) if nonfinite_step or ( args.max_steps and step >= args.max_steps): break layers.append(layer_record) if nonfinite_step or (args.max_steps and step >= args.max_steps): break else: for epoch in range(args.epochs): rate = scheduled_rate(args.lr, epoch, args) output_rate = scheduled_rate(args.output_lr, epoch, args) loss_sum = 0.0 examples = 0 sync(args.device) epoch_started = time.time() for image, labels in train: if args.method == "pepita": loss = net.pepita_step( image, labels, rate, eta_output=output_rate, momentum=args.momentum, weight_decay=args.weight_decay) elif args.method == "ep": loss, _ = net.ep_step( image, labels, rate, eta_output=output_rate, momentum=args.momentum, weight_decay=args.weight_decay, generator=ep_generator) else: loss = net.dualprop_step( image, labels, rate, eta_output=output_rate, momentum=args.momentum, weight_decay=args.weight_decay) batch = labels.numel() loss_sum += loss * batch examples += batch ordinary_examples += batch step += 1 if not math.isfinite(loss): nonfinite_step = step break if args.max_steps and step >= args.max_steps: break sync(args.device) evaluation = None eval_seconds = 0.0 if args.eval_every and (epoch + 1) % args.eval_every == 0: eval_started = time.time() evaluation = evaluate(net, args.method, validation) sync(args.device) eval_seconds = time.time() - eval_started validation_examples += evaluation["examples"] row = { "epoch": epoch + 1, "lr": rate, "output_lr": output_rate, "train_loss": loss_sum / examples, "train_examples": examples, "train_seconds": time.time() - epoch_started - eval_seconds, "validation": evaluation, "validation_seconds": eval_seconds, } epochs.append(row) accuracy = ( "" if evaluation is None else f" val={evaluation['accuracy']:.4f}") print( f"epoch={epoch + 1}/{args.epochs} " f"loss={row['train_loss']:.6g}{accuracy}", flush=True) if nonfinite_step or ( args.max_steps and step >= args.max_steps): break sync(args.device) final_started = time.time() final = evaluate(net, args.method, validation) sync(args.device) final_seconds = time.time() - final_started validation_examples += final["examples"] final["finite"] = ( final["loss"] is None or math.isfinite(final["loss"])) completed_epochs = ( sum(len(layer["epochs"]) for layer in layers) if args.method == "ff" else len(epochs)) record = { "schema_version": 1, "protocol_family": "resnet_local_learning_crossover", "args": vars(args), "provenance": provenance, "split": split, "architecture": { "family": "CIFAR 6n+2 ResNet, option-A shortcuts", "depth": net.depth, "base_width": net.base_width, "normalization": net.normalization, "residual_scale": net.residual_scale, "forward_parameter_count": net.n_forward_parameters, }, "epochs": epochs, "layers": layers, "first_nonfinite_step": nonfinite_step, "evaluation_protocol": { "split": "validation", "test_evaluations": 0, "test_used_for_selection": False, }, "final": {**final, "evaluation_seconds": final_seconds}, "work": work_report( net, args, ordinary_examples, validation_examples, completed_epochs), "hardware": hardware_report(args.device), "total_wall_seconds": time.time() - started, } os.makedirs(os.path.dirname(os.path.abspath(args.out)), exist_ok=True) with open(args.out, "w", encoding="utf-8") as handle: json.dump(record, handle, indent=2, sort_keys=True) handle.write("\n") print(json.dumps({ "out": args.out, "final": record["final"], "work": record["work"], }, indent=2, sort_keys=True)) def parse_args(): parser = argparse.ArgumentParser() parser.add_argument("--method", choices=METHODS, 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("--feedback_seed", type=int, default=1729) parser.add_argument("--batch_size", type=int, default=128) parser.add_argument("--epochs", type=int, default=10) 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",), default="validation") parser.add_argument("--eval_every", type=int, default=1) parser.add_argument("--augment_train", type=int, choices=(0, 1), default=1) parser.add_argument("--lr", type=float, required=True) parser.add_argument("--output_lr", type=float) parser.add_argument( "--lr_schedule", choices=("constant", "step", "pepita"), default="constant") parser.add_argument("--lr_milestones", default="100,150") parser.add_argument("--lr_gamma", type=float, default=0.1) parser.add_argument("--momentum", type=float, default=0.9) parser.add_argument("--weight_decay", type=float, default=1e-4) parser.add_argument("--pepita_projection_scale", type=float, default=0.05) parser.add_argument("--ff_threshold", type=float, default=2.0) parser.add_argument("--ff_score_from_layer", type=int, default=1) parser.add_argument("--ep_beta", type=float, default=0.5) parser.add_argument("--ep_dt", type=float, default=0.5) parser.add_argument("--ep_free_steps", type=int, default=20) parser.add_argument("--ep_nudge_steps", type=int, default=4) parser.add_argument("--dp_alpha", type=float, default=0.0) parser.add_argument("--dp_beta", type=float, default=0.1) parser.add_argument("--dp_inference_passes", type=int, default=16) args = parser.parse_args() if args.output_lr is None: args.output_lr = args.lr return args if __name__ == "__main__": run(parse_args())