diff options
Diffstat (limited to 'experiments/conv_run.py')
| -rw-r--r-- | experiments/conv_run.py | 454 |
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()) |
