#!/usr/bin/env python3 """Post-failure oracle audit of convolutional apical representability. This script never updates the network. Exact hidden gradients are used only to distinguish estimator failure from feedback-family misspecification at the same deterministic, untrained ResNet-20 state used by the v2 calibration screen. """ import argparse import json import os import subprocess import torch import torch.nn.functional as F import sys sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from sdil.conv import CIFARSDILResNet 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 mean_sample_cosine(prediction, target): return float(F.cosine_similarity( prediction.flatten(1), target.flatten(1), dim=1).mean()) def energy_fraction(prediction, target): return float(prediction.square().sum() / target.square().sum().clamp_min(1e-30)) def per_example_spatial_oracle(hidden, target): """Project each example/channel independently onto [1, tanh(h)].""" gate = torch.tanh(hidden) mean_gate = gate.mean(dim=(2, 3), keepdim=True) mean_target = target.mean(dim=(2, 3), keepdim=True) centered_gate = gate - mean_gate centered_target = target - mean_target slope = ((centered_gate * centered_target).mean(dim=(2, 3), keepdim=True) / centered_gate.square().mean( dim=(2, 3), keepdim=True).clamp_min(1e-12)) return mean_target + slope * centered_gate def spatial_template_cv(output_train, output_eval, target_train): """Fit the best output-error linear map independently at every unit.""" flat = target_train.flatten(1) coefficients = torch.linalg.lstsq(output_train, flat).solution return output_eval @ coefficients def basis_fields(hidden): gate = torch.tanh(hidden) local = F.avg_pool2d( gate, kernel_size=3, stride=1, padding=1, count_include_pad=False) channel_context = gate.mean(dim=1, keepdim=True).expand_as(gate) return { "channel_gated": [torch.ones_like(gate), gate], "local_context": [torch.ones_like(gate), gate, local, channel_context], } def state_basis_cv(output_train, output_eval, fields, target_train): """Fit output-error-conditioned coefficients for fixed local state bases.""" batch_train, channels, height, width = target_train.shape batch_eval = output_eval.shape[0] spatial = height * width train_features = torch.cat([ field["train"].reshape(batch_train, channels, spatial, 1) * output_train[:, None, None, :] for field in fields], dim=3) train_target = target_train.reshape(batch_train, channels, spatial) gram = torch.einsum("bcsp,bcsq->cpq", train_features, train_features) rhs = torch.einsum("bcsp,bcs->cp", train_features, train_target) coefficients = torch.linalg.pinv(gram) @ rhs[:, :, None] eval_features = torch.cat([ field["eval"].reshape(batch_eval, channels, spatial, 1) * output_eval[:, None, None, :] for field in fields], dim=3) prediction = torch.einsum( "bcsp,cpk->bcsk", eval_features, coefficients).squeeze(-1) return prediction.reshape(batch_eval, channels, height, width) def summarize(values): early = max(1, len(values) // 3) return { "per_layer": values, "early_third_mean": sum(values[:early]) / early, "all_layer_mean": sum(values) / len(values), } def exact_batch(net, x, y): forward = net.forward(x, training=True, update_stats=False) gradients = torch.autograd.grad( F.cross_entropy(forward["logits"], y), forward["hiddens"]) targets = [-x.shape[0] * value.detach() for value in gradients] output_signal = (torch.softmax(forward["logits"].detach(), dim=1) - F.one_hot(y, 10).to(forward["logits"].dtype)) return [value.detach() for value in forward["hiddens"]], targets, output_signal def main(): parser = argparse.ArgumentParser() parser.add_argument("--device", default="cpu") parser.add_argument("--data_dir", default=DATA_DIR) parser.add_argument("--batch", type=int, default=64) parser.add_argument("--fit_examples", type=int, default=32) parser.add_argument( "--out", default="results/oral_a_representation_diagnosis.json") args = parser.parse_args() if args.batch != 64 or args.fit_examples != 32: raise ValueError("the audited diagnosis is fixed to a 32/32 split") torch.manual_seed(0) train, _, _, _, _, split = get_cifar_image_splits( batch_size=64, data_dir=args.data_dir, device=args.device, train_limit=10000, val_examples=5000, split_seed=2027, loader_seed=0, augment_train=False) net = CIFARSDILResNet( depth=20, base_width=16, n_classes=10, device=args.device, seed=0, residual_scale=1.0, normalization="batchnorm", vectorizer_mode="channel_gated", a_scale=0.0) parameters = net.W + net.gamma + net.beta + [net.W_out, net.b_out] for parameter in parameters: parameter.requires_grad_(True) x = train.x[:args.batch] y = train.y[:args.batch] hidden_train, target_train, output_train = exact_batch( net, x[:args.fit_examples], y[:args.fit_examples]) hidden_eval, target_eval, output_eval = exact_batch( net, x[args.fit_examples:], y[args.fit_examples:]) for parameter in parameters: parameter.requires_grad_(False) output_train = output_train.to(torch.float64) output_eval = output_eval.to(torch.float64) reports = { "per_example_spatial_oracle": [], "channel_gated_cv": [], "local_context_cv": [], "spatial_template_cv": [], } energies = {key: [] for key in reports} for h_train, q_train, h_eval, q_eval in zip( hidden_train, target_train, hidden_eval, target_eval): h_train = h_train.to(torch.float64) q_train = q_train.to(torch.float64) h_eval = h_eval.to(torch.float64) q_eval = q_eval.to(torch.float64) prediction = per_example_spatial_oracle(h_eval, q_eval) reports["per_example_spatial_oracle"].append( mean_sample_cosine(prediction, q_eval)) energies["per_example_spatial_oracle"].append( energy_fraction(prediction, q_eval)) train_fields = basis_fields(h_train) eval_fields = basis_fields(h_eval) for name in ("channel_gated", "local_context"): paired = [{"train": fit_field, "eval": eval_field} for fit_field, eval_field in zip( train_fields[name], eval_fields[name])] prediction = state_basis_cv( output_train, output_eval, paired, q_train) key = f"{name}_cv" reports[key].append(mean_sample_cosine(prediction, q_eval)) energies[key].append(energy_fraction(prediction, q_eval)) flat_prediction = spatial_template_cv( output_train, output_eval, q_train).reshape_as(q_eval) reports["spatial_template_cv"].append( mean_sample_cosine(flat_prediction, q_eval)) energies["spatial_template_cv"].append( energy_fraction(flat_prediction, q_eval)) output = { "protocol": "oral_a_post_failure_representation_diagnosis_v1", "provenance": provenance(), "split": split, "network": { "depth": 20, "width": 16, "seed": 0, "normalization": "batchnorm", "residual_scale": 1.0, "forward_state": "deterministic_initialization_no_updates", }, "probe": { "examples": args.batch, "fit_examples": args.fit_examples, "evaluation_examples": args.batch - args.fit_examples, "source": "unaugmented first training-prefix examples", "batchnorm_cv": ( "fit and evaluation halves use separate 32-example batch " "statistics and exact-gradient graphs"), "test_examples_touched": 0, }, "cosine": {key: summarize(value) for key, value in reports.items()}, "prediction_energy_over_target_energy": { key: summarize(value) for key, value in energies.items()}, "interpretation": { "per_example_spatial_oracle": ( "upper bound for the two [1,tanh(h)] fields with unconstrained " "per-example/channel coefficients"), "channel_gated_cv": ( "actual v2 vectorizer family fit on 32 examples and evaluated " "on 32 disjoint examples"), "local_context_cv": ( "diagnostic four-field family adding local spatial average and " "cross-channel somatic context"), "spatial_template_cv": ( "output-error linear map with an independent coefficient at " "every hidden unit"), "status": "post_failure_diagnosis_not_a_learning_result_or_gate", }, } os.makedirs(os.path.dirname(os.path.abspath(args.out)), exist_ok=True) with open(args.out, "w") as handle: json.dump(output, handle, indent=2, sort_keys=True) handle.write("\n") print(json.dumps({"cosine": output["cosine"]}, indent=2)) if __name__ == "__main__": main()