From ba50c8fc286f6071c688a12606216cfb871a7ee0 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 12:44:53 -0500 Subject: analysis: add apical representation oracle audit --- experiments/analyze_oral_a_representation.py | 239 +++++++++++++++++++++++++++ 1 file changed, 239 insertions(+) create mode 100644 experiments/analyze_oral_a_representation.py (limited to 'experiments/analyze_oral_a_representation.py') diff --git a/experiments/analyze_oral_a_representation.py b/experiments/analyze_oral_a_representation.py new file mode 100644 index 0000000..da88ec3 --- /dev/null +++ b/experiments/analyze_oral_a_representation.py @@ -0,0 +1,239 @@ +#!/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() -- cgit v1.2.3