#!/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): residual_scale = args.residual_scale if residual_scale is None and args.normalization == "batchnorm": residual_scale = 1.0 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=residual_scale, normalization=args.normalization, bn_momentum=args.bn_momentum, bn_eps=args.bn_eps) if args.mode == "bp": return CIFARLocalResNet(**common), None net = CIFARSDILResNet( **common, a_scale=args.a_scale, apical_seed=args.apical_seed, vectorizer_mode=args.vectorizer_mode) 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, apical_calibration_mode=args.apical_calibration_mode, 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"], "causal_scalar_observations": counters["causal_scalar_observations"], "per_example_cross_entropy_terms": counters["per_example_loss_terms"], "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, "causal_scalar_observations": 0, "per_example_loss_terms": 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": "CIFAR 6n+2 ResNet, option-A shortcuts", "normalization": net.normalization, "bn_momentum": net.bn_momentum if net.normalization == "batchnorm" else None, "bn_eps": net.bn_eps if net.normalization == "batchnorm" else None, "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), "vectorizer_parameters": getattr(net, "n_vectorizer_parameters", 0), "predictor_parameters": getattr(net, "n_predictor_parameters", 0), "vectorizer_mode": getattr(net, "vectorizer_mode", None), "fixed_traffic_coefficients": getattr( net, "n_fixed_traffic_coefficients", 0), }, "epochs": [], } sync(args.device) total_start = time.time() predictor_warmup_wall = 0.0 apical_warmup_wall = 0.0 loader_state = train.g.get_state().clone() if config is not None and config.learn_P and args.predictor_warmup_steps: sync(args.device) warmup_start = time.time() 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, training=True, update_stats=False) net.predictor_step( forward["hiddens"], config.eta_P, config.nuisance_scale) counters["predictor_warmup_examples"] += x.shape[0] train.g.set_state(loader_state) sync(args.device) predictor_warmup_wall = time.time() - warmup_start if config is not None and args.a_warmup_steps: sync(args.device) warmup_start = time.time() 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["causal_scalar_observations"] += 2 * config.pert_directions * ( 1 if net.normalization == "batchnorm" else batch) counters["per_example_loss_terms"] += 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], } sync(args.device) apical_warmup_wall = time.time() - warmup_start step = 0 train_wall = 0.0 eval_wall = 0.0 validation_evaluations = 0 test_evaluations = 0 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["causal_scalar_observations"] += ( 2 * config.pert_directions * (1 if net.normalization == "batchnorm" else batch)) counters["per_example_loss_terms"] += ( 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": { "predictor_warmup_wall_s": predictor_warmup_wall, "apical_warmup_wall_s": apical_warmup_wall, "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("--normalization", choices=("batchnorm", "none"), default="batchnorm") parser.add_argument("--bn_momentum", type=float, default=0.1) parser.add_argument("--bn_eps", type=float, default=1e-5) parser.add_argument("--a_scale", type=float, default=1.0) parser.add_argument("--vectorizer_mode", choices=("spatial_template", "channel_gated"), default="spatial_template") 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( "--apical_calibration_mode", choices=("unit_targets", "channel_subspace"), default="unit_targets") 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())