#!/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 (CIFARHierarchicalFAResNet, CIFARKPMixedTrafficResNet, CIFARKPResNet, CIFARLocalResNet, CIFARSDILResNet, ConvSDILConfig, conv_alignment_report, conv_apical_calibration_step, conv_hierarchical_alignment_report, conv_hierarchical_step, conv_kolen_pollack_step, conv_kp_mixed_traffic_alignment_report, conv_kp_mixed_traffic_step, conv_learned_hierarchical_step, conv_local_step, evaluate_conv, hierarchical_feedback_tracking_report, hierarchical_parameter_subspace_calibration) from sdil.conv import (normalized_residual_mirror_step, normalized_response_mirror_step) 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 if args.mode in ("hfa", "lhfa", "wm", "rrm", "kp", "kp_traffic"): if args.mode == "kp_traffic": net = CIFARKPMixedTrafficResNet( **common, feedback_seed=args.apical_seed, feedback_scale=args.a_scale, traffic_seed=args.traffic_seed) net.traffic_rule = args.traffic_rule else: network_class = ( CIFARKPResNet if args.mode == "kp" else CIFARHierarchicalFAResNet) net = network_class( **common, feedback_seed=args.apical_seed, feedback_scale=args.a_scale) config = ConvSDILConfig( eta=args.lr, eta_output=args.output_lr, eta_A=(args.eta_A if args.mode == "lhfa" else 0.0), eta_P=args.eta_P, momentum=args.momentum, weight_decay=args.weight_decay, learn_A=args.mode == "lhfa", learn_P=args.mode == "kp_traffic", pert_sigma=args.pert_sigma, pert_every=args.pert_every, pert_directions=args.pert_directions, apical_calibration_mode="hierarchical_parameter_subspace") config.validate() return net, config 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"] + counters["traffic_calibration_examples"] + counters["traffic_audit_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"] + counters["traffic_calibration_examples"]) * apical_macs) regression_multiplier = 2 if mode == "lhfa" else 1 apical_regression = (regression_multiplier * counters["calibration_event_examples"] * apical_macs) mirror_conv_macs = max(0, apical_macs - getattr(net, "R_out", torch.empty(0)).numel()) mirror_readout_macs = getattr(net, "R_out", torch.empty(0)).numel() mirror_forward = (counters["mirror_conv_examples"] * mirror_conv_macs + counters["mirror_readout_examples"] * mirror_readout_macs) mirror_prediction = mirror_forward if mode == "rrm" else 0 mirror_correlation = mirror_forward kp_feedback_correlation = ( counters["ordinary_examples"] * apical_macs if mode in ("kp", "kp_traffic") else 0) if mode == "kp_traffic": elementwise_operations = ( counters["ordinary_examples"] * net.mixed_elementwise_ops_per_example() + (counters["predictor_warmup_examples"] + counters["predictor_update_examples"]) * net.predictor_elementwise_ops_per_example + counters["traffic_calibration_examples"] * net.traffic_calibration_elementwise_ops_per_example + counters["traffic_audit_examples"] * net.predictor_audit_elementwise_ops_per_example + counters["neutral_projection_examples"] * net.neutral_projection_elementwise_ops_per_example) else: elementwise_operations = 0 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, "mirror_response_macs": mirror_forward, "mirror_feedback_prediction_macs": mirror_prediction, "mirror_local_correlation_macs": mirror_correlation, "kp_reciprocal_correlation_macs": kp_feedback_correlation, } return { "forward_macs_per_example": forward_macs, "apical_macs_per_example": apical_macs, "components": components, "total_macs_estimate": sum(components.values()), "elementwise_operations_estimate": elementwise_operations, "total_clean_forward_examples": ( counters["ordinary_examples"] + counters["predictor_warmup_examples"] + counters["apical_warmup_examples"] + counters["traffic_calibration_examples"] + counters["traffic_audit_examples"]), "total_forward_equivalent_examples": ( counters["ordinary_examples"] + counters["predictor_warmup_examples"] + counters["apical_warmup_examples"] + counters["traffic_calibration_examples"] + counters["traffic_audit_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"], "mirror_probe_examples": counters["mirror_conv_examples"], "mirror_readout_probe_examples": counters["mirror_readout_examples"], "neutral_projection_observations": counters[ "neutral_projection_examples"], "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; mixed-traffic/predictor elementwise arithmetic is reported " "as a conservative operation estimate separately and is not folded " "into MACs"), } 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 not in ("sdil", "lhfa", "kp_traffic") and ( args.a_warmup_steps or args.learn_P): raise ValueError("apical/predictor warmup is restricted to SDIL") if args.mode == "lhfa" and args.learn_P: raise ValueError("predictor learning is not defined for learned HFA") if args.mode == "kp_traffic": if not args.learn_P or args.predictor_warmup_steps < 1: raise ValueError("mixed-traffic KP requires predictor warmup") if args.traffic_calibration_examples < 1 or args.traffic_ratio <= 0: raise ValueError("invalid mixed-traffic calibration") if args.neutral_projection: if (args.predictor_mode != "closed_form" or args.predictor_warmup_steps != 1 or args.predictor_every != 0): raise ValueError( "neutral projection requires frozen one-step closed-form " "prediction; raw and matched rules compute it as a sham " "control without applying the subtractive direction") elif args.predictor_every < 1: raise ValueError("mixed-traffic KP requires a predictor cadence") if (args.predictor_mode == "closed_form" and args.predictor_warmup_steps != 1): raise ValueError("closed-form predictor requires one warmup fit") elif args.predictor_every: raise ValueError("predictor cadence is restricted to mixed-traffic KP") if args.mode != "kp_traffic" and args.neutral_projection: raise ValueError("neutral projection is restricted to mixed-traffic KP") if args.mode != "kp_traffic" and args.predictor_mode != "nlms": raise ValueError("predictor mode is restricted to mixed-traffic KP") if args.mode not in ("wm", "rrm") and args.mirror_warmup_steps: raise ValueError("mirror warmup is restricted to weight mirror mode") if args.mirror_every < 1 or args.mirror_batch_size < 1: raise ValueError("invalid mirror cadence or batch size") if not 0.0 < args.mirror_eta <= 1.0 or args.mirror_noise_std <= 0: raise ValueError("invalid mirror learning hyperparameters") 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) mirror_generator = torch.Generator(device=torch.device(args.device)).manual_seed( args.mirror_seed) 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, "mirror_conv_examples": 0, "mirror_readout_examples": 0, "mirror_events": 0, "traffic_calibration_examples": 0, "traffic_audit_examples": 0, "predictor_update_examples": 0, "neutral_projection_examples": 0, } log = { "schema_version": 1, "protocol_family": "oral_a_cifar_local_resnet_development", "calibration_metric_space": ( None if config is None or args.mode == "hfa" else "reciprocal_local_activity_products_with_mixed_apical_traffic" if args.mode == "kp_traffic" else "hierarchical_feedback_parameters" if args.mode == "lhfa" else "reciprocal_local_activity_products" if args.mode == "kp" else { "wm": "local_parent_child_response", "rrm": "local_parent_child_response_residual", "unit_targets": "full_hidden_field", "channel_subspace": "channel_basis_moments", "vectorizer_subspace": "vectorizer_parameter_gradients", }[args.mode if args.mode in ("wm", "rrm") else config.apical_calibration_mode]), "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), "fixed_feedback_parameters": ( getattr(net, "n_fixed_feedback_parameters", 0) if args.mode == "hfa" else 0), "adaptive_feedback_parameters": ( getattr(net, "n_fixed_feedback_parameters", 0) if args.mode in ("lhfa", "wm", "rrm", "kp", "kp_traffic") else 0), }, "epochs": [], } sync(args.device) total_start = time.time() predictor_warmup_wall = 0.0 apical_warmup_wall = 0.0 mirror_warmup_wall = 0.0 loader_state = train.g.get_state().clone() traffic_calibration_batch = None closed_form_predictor_ready = False if args.mode == "kp_traffic": count = min(args.traffic_calibration_examples, train.x.shape[0]) if count != args.traffic_calibration_examples: raise ValueError("traffic calibration prefix is unavailable") calibration_x = train.x[:count] calibration_y = train.y[:count] traffic_calibration_batch = (calibration_x, calibration_y) calibration_forward = net.forward( calibration_x, return_cache=True, training=True, update_stats=False) calibration_output_error = ( torch.softmax(calibration_forward["logits"], dim=1) - torch.nn.functional.one_hot( calibration_y, net.n_classes).to( calibration_forward["logits"].dtype)) calibration_instruction = net.hierarchical_teaching( calibration_output_error, calibration_forward) log["traffic_calibration"] = net.calibrate_traffic_gain( calibration_instruction, calibration_forward["hiddens"], args.traffic_ratio) log["traffic_calibration"].update({ "examples": count, "data_source": "first unaugmented development-training examples", "uses_validation_endpoint": False, }) counters["traffic_calibration_examples"] += count if args.predictor_mode == "closed_form": sync(args.device) closed_form_start = time.time() fit = net.predictor_closed_form_fit( calibration_forward["hiddens"], stability_margin=0.0) sync(args.device) predictor_warmup_wall += time.time() - closed_form_start counters["predictor_update_examples"] += count post_fit_ratio = net.predictor_traffic_residual_rms_ratio( calibration_forward["hiddens"]) log["predictor_warmup"] = { "mode": "closed_form", "steps": 1, "examples": count, "closed_form_fit": fit, "post_warmup_traffic_residual_rms_ratio": post_fit_ratio, "instruction_present": False, "task_loader_state_restored": True, "reuses_traffic_calibration_forward": True, } closed_form_predictor_ready = True del (calibration_forward, calibration_output_error, calibration_instruction) if args.mode in ("wm", "rrm") and args.mirror_warmup_steps: sync(args.device) mirror_start = time.time() mirror_metrics = [] for _ in range(args.mirror_warmup_steps): mirror_function = (normalized_residual_mirror_step if args.mode == "rrm" else normalized_response_mirror_step) metric = mirror_function( net, batch_size=args.mirror_batch_size, noise_std=args.mirror_noise_std, eta=args.mirror_eta, generator=mirror_generator) mirror_metrics.append(metric) counters["mirror_conv_examples"] += args.mirror_batch_size counters["mirror_readout_examples"] += metric["readout_batch_size"] counters["mirror_events"] += 1 log["mirror_warmup"] = { "steps": args.mirror_warmup_steps, "first": mirror_metrics[0], "mean": {key: sum(value[key] for value in mirror_metrics) / len(mirror_metrics) for key in mirror_metrics[0]}, "last": mirror_metrics[-1], } sync(args.device) mirror_warmup_wall = time.time() - mirror_start if (config is not None and config.learn_P and args.predictor_warmup_steps and not closed_form_predictor_ready): sync(args.device) warmup_start = time.time() iterator = iter(train) predictor_metrics = [] 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) predictor_metric = (net.predictor_step( forward["hiddens"], config.eta_P) if args.mode == "kp_traffic" else net.predictor_step( forward["hiddens"], config.eta_P, config.nuisance_scale)) predictor_metrics.append(predictor_metric) counters["predictor_warmup_examples"] += x.shape[0] del forward train.g.set_state(loader_state) loader_state_restored = torch.equal(train.g.get_state(), loader_state) if not loader_state_restored: raise AssertionError("predictor warmup did not restore loader state") log["predictor_warmup"] = { "mode": "nlms", "steps": args.predictor_warmup_steps, "first_mse": predictor_metrics[0], "mean_mse": sum(predictor_metrics) / len(predictor_metrics), "last_mse": predictor_metrics[-1], "instruction_present": False, "task_loader_state_restored": loader_state_restored, } if args.mode == "kp_traffic": audit_x, _ = traffic_calibration_batch audit_forward = net.forward( audit_x, training=True, update_stats=False) log["predictor_warmup"][ "post_warmup_traffic_residual_rms_ratio"] = ( net.predictor_traffic_residual_rms_ratio( audit_forward["hiddens"])) counters["traffic_audit_examples"] += audit_x.shape[0] del audit_forward 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) if args.mode == "lhfa": forward = net.forward( x, return_cache=True, training=True, update_stats=False) output_signal = ( torch.softmax(forward["logits"], dim=1) - torch.nn.functional.one_hot( y, net.n_classes).to(forward["logits"].dtype)) metric = hierarchical_parameter_subspace_calibration( net, x, y, forward, output_signal, sigma=config.pert_sigma, n_directions=config.pert_directions, eta=config.eta_A, generator=warmup_generator) else: _, 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) warmup_mean = { key: sum(metric[key] for metric in warmup_metrics) / len(warmup_metrics) for key in warmup_metrics[0] } log["apical_warmup"] = { "steps": args.a_warmup_steps, "first": warmup_metrics[0], "mean": warmup_mean, "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 = [] mirror_metrics = [] signal_metrics = [] neutral_projection_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 elif args.mode == "hfa": result = conv_hierarchical_step(net, x, y, config) loss = result["loss"] did_perturb = False elif args.mode == "lhfa": result = conv_learned_hierarchical_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"]) elif args.mode == "wm": if step % args.mirror_every == 0: mirror_metric = normalized_response_mirror_step( net, batch_size=args.mirror_batch_size, noise_std=args.mirror_noise_std, eta=args.mirror_eta, generator=mirror_generator) mirror_metrics.append(mirror_metric) counters["mirror_conv_examples"] += args.mirror_batch_size counters["mirror_readout_examples"] += mirror_metric[ "readout_batch_size"] counters["mirror_events"] += 1 result = conv_hierarchical_step(net, x, y, config) loss = result["loss"] did_perturb = False elif args.mode == "rrm": if step % args.mirror_every == 0: mirror_metric = normalized_residual_mirror_step( net, batch_size=args.mirror_batch_size, noise_std=args.mirror_noise_std, eta=args.mirror_eta, generator=mirror_generator) mirror_metrics.append(mirror_metric) counters["mirror_conv_examples"] += args.mirror_batch_size counters["mirror_readout_examples"] += mirror_metric[ "readout_batch_size"] counters["mirror_events"] += 1 result = conv_hierarchical_step(net, x, y, config) loss = result["loss"] did_perturb = False elif args.mode == "kp": result = conv_kolen_pollack_step(net, x, y, config) loss = result["loss"] did_perturb = False elif args.mode == "kp_traffic": result = conv_kp_mixed_traffic_step( net, x, y, config, step, args.traffic_rule, args.predictor_every, neutral_projection=bool(args.neutral_projection)) loss = result["loss"] did_perturb = False signal_metrics.append({key: result[key] for key in ( "teaching_rms", "instruction_rms", "raw_apical_rms", "innovation_rms", "traffic_rms")}) if result["did_predictor_update"]: counters["predictor_update_examples"] += batch if result["neutral_projection"] is not None: neutral_projection_metrics.append( result["neutral_projection"]) counters["neutral_projection_examples"] += batch 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 mirror_metrics: record["mirror"] = { key: sum(value[key] for value in mirror_metrics) / len(mirror_metrics) for key in mirror_metrics[0] } if signal_metrics: record["mixed_apical"] = { key: sum(value[key] for value in signal_metrics) / len(signal_metrics) for key in signal_metrics[0] } if neutral_projection_metrics: record["neutral_projection"] = { "maximum_pre_projection_traffic_rms_ratio": max( value["pre_projection_traffic_rms_ratio"] for value in neutral_projection_metrics), "maximum_post_projection_traffic_rms_ratio": max( value["post_projection_traffic_rms_ratio"] for value in neutral_projection_metrics), "maximum_absolute_post_projection_soma_slope": max( value["max_absolute_post_projection_soma_slope"] for value in neutral_projection_metrics), "maximum_absolute_correction_slope": max( value["max_absolute_correction_slope"] for value in neutral_projection_metrics), "minimum_observations": min( value["observations"] for value in neutral_projection_metrics), "maximum_observations": max( value["observations"] for value in neutral_projection_metrics), "instruction_observations": sum( value["instruction_observations"] for value in neutral_projection_metrics), } if args.mode in ("kp", "kp_traffic"): tracking = hierarchical_feedback_tracking_report(net) record["feedback_tracking"] = { "mean_feedback_forward_cosine": tracking[ "mean_feedback_forward_cosine"], "mean_feedback_forward_relative_error": tracking[ "mean_feedback_forward_relative_error"], "min_feedback_forward_cosine": min( tracking["feedback_forward_cosine"]), "max_feedback_forward_relative_error": max( tracking["feedback_forward_relative_error"]), } 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() if args.mode == "kp_traffic": diagnostics = conv_kp_mixed_traffic_alignment_report( net, train.x[:probe], train.y[:probe], args.traffic_rule, neutral_projection=bool(args.neutral_projection)) elif args.mode in ("hfa", "lhfa", "wm", "rrm", "kp"): diagnostics = conv_hierarchical_alignment_report( net, train.x[:probe], train.y[:probe]) else: 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, "mirror_warmup_wall_s": mirror_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", "hfa", "lhfa", "wm", "rrm", "kp", "kp_traffic", "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", "vectorizer_subspace"), default="unit_targets") parser.add_argument("--predictor_warmup_steps", type=int, default=0) parser.add_argument("--predictor_every", type=int, default=0) parser.add_argument("--predictor_mode", choices=("nlms", "closed_form"), default="nlms") parser.add_argument("--neutral_projection", type=int, choices=(0, 1), default=0) parser.add_argument("--traffic_rule", choices=("raw", "matched", "innovation"), default="innovation") parser.add_argument("--traffic_seed", type=int, default=4000) parser.add_argument("--traffic_ratio", type=float, default=4.0) parser.add_argument("--traffic_calibration_examples", type=int, default=64) parser.add_argument("--a_warmup_steps", type=int, default=0) parser.add_argument("--mirror_warmup_steps", type=int, default=0) parser.add_argument("--mirror_every", type=int, default=16) parser.add_argument("--mirror_batch_size", type=int, default=1) parser.add_argument("--mirror_eta", type=float, default=0.1) parser.add_argument("--mirror_noise_std", type=float, default=1.0) parser.add_argument("--mirror_seed", type=int, default=3000) parser.add_argument("--alignment_probe", type=int, default=0) return parser.parse_args() if __name__ == "__main__": run(parse_args())