From 45d2b57310d556c4d1e78ee185279eda14189585 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 22 Jul 2026 13:03:46 -0500 Subject: analysis: add hierarchical feedback oracle audit --- experiments/analyze_oral_a_hierarchical_oracle.py | 288 ++++++++++++++++++++++ 1 file changed, 288 insertions(+) create mode 100644 experiments/analyze_oral_a_hierarchical_oracle.py diff --git a/experiments/analyze_oral_a_hierarchical_oracle.py b/experiments/analyze_oral_a_hierarchical_oracle.py new file mode 100644 index 0000000..71488b1 --- /dev/null +++ b/experiments/analyze_oral_a_hierarchical_oracle.py @@ -0,0 +1,288 @@ +#!/usr/bin/env python3 +"""Post-failure oracle audit for spatially hierarchical feedback. + +The audit uses exact gradients only as held-out regression targets. It asks +whether the first ResNet stage can be represented by direct output feedback, +ordinary downstream activation context, or a learned local map from the true +error fields at the actual child nodes in the residual DAG. No parameter of the +network or learning algorithm is updated. +""" +import argparse +import json +import os +import subprocess +import sys + +import torch +import torch.nn.functional as F + +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__))) +EARLY_CHILDREN = ((1, 2), (2,), (3, 4), (4,), (5, 6), (6,)) + + +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 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.detach(), forward["features"].detach()) + + +def cosine(prediction, target): + return float(F.cosine_similarity( + prediction.flatten(1), target.flatten(1), dim=1).mean()) + + +def energy_ratio(prediction, target): + return float(prediction.square().sum() + / target.square().sum().clamp_min(1e-30)) + + +def solve_ridge(features, target, relative_ridge=1e-6): + gram = features.t() @ features + ridge = relative_ridge * gram.diag().mean().clamp_min(1e-30) + return torch.linalg.solve( + gram + ridge * torch.eye( + gram.shape[0], dtype=gram.dtype, device=gram.device), + features.t() @ target) + + +def direct_context_oracle(output_train, output_eval, hidden_train, hidden_eval, + target_train): + """Best held-out A c + tanh(h) G c map at fixed initialization.""" + batch, channels, height, width = target_train.shape + spatial = height * width + context = output_train.shape[1] + train_fields = (torch.ones_like(hidden_train), torch.tanh(hidden_train)) + eval_fields = (torch.ones_like(hidden_eval), torch.tanh(hidden_eval)) + gram = target_train.new_zeros(channels, 2 * context, 2 * context) + rhs = target_train.new_zeros(channels, 2 * context) + for left in range(2): + projected = (train_fields[left] * target_train).flatten(2).sum(2) + rhs[:, left * context:(left + 1) * context] = torch.einsum( + "bd,bc->cd", output_train, projected) + for right in range(2): + weights = (train_fields[left] * train_fields[right]).flatten( + 2).sum(2) + gram[:, left * context:(left + 1) * context, + right * context:(right + 1) * context] = torch.einsum( + "bd,be,bc->cde", output_train, output_train, weights) + coefficients = (torch.linalg.pinv(gram) @ rhs[:, :, None]).squeeze(-1) + coefficients = coefficients.reshape(channels, 2, context) + prediction = target_train.new_zeros( + output_eval.shape[0], channels, height, width) + for field in range(2): + prediction.add_( + (output_eval @ coefficients[:, field].t())[:, :, None, None] + * eval_fields[field]) + return prediction + + +def activation_context_oracle(output_train, output_eval, hidden_train, + hidden_eval, child_train, child_eval, + target_train): + """Add one spatial basis from ordinary downstream activation context.""" + context_train = torch.tanh(child_train).mean(dim=1, keepdim=True).expand_as( + hidden_train) + context_eval = torch.tanh(child_eval).mean(dim=1, keepdim=True).expand_as( + hidden_eval) + fields_train = (torch.ones_like(hidden_train), torch.tanh(hidden_train), + context_train) + fields_eval = (torch.ones_like(hidden_eval), torch.tanh(hidden_eval), + context_eval) + batch, channels, height, width = target_train.shape + context = output_train.shape[1] + n_fields = len(fields_train) + gram = target_train.new_zeros( + channels, n_fields * context, n_fields * context) + rhs = target_train.new_zeros(channels, n_fields * context) + for left in range(n_fields): + projected = (fields_train[left] * target_train).flatten(2).sum(2) + rhs[:, left * context:(left + 1) * context] = torch.einsum( + "bd,bc->cd", output_train, projected) + for right in range(n_fields): + weights = (fields_train[left] * fields_train[right]).flatten( + 2).sum(2) + gram[:, left * context:(left + 1) * context, + right * context:(right + 1) * context] = torch.einsum( + "bd,be,bc->cde", output_train, output_train, weights) + coefficients = (torch.linalg.pinv(gram) @ rhs[:, :, None]).squeeze(-1) + coefficients = coefficients.reshape(channels, n_fields, context) + prediction = target_train.new_zeros( + output_eval.shape[0], channels, height, width) + for field in range(n_fields): + prediction.add_( + (output_eval @ coefficients[:, field].t())[:, :, None, None] + * fields_eval[field]) + return prediction + + +def hierarchical_features(hiddens, targets, child_indices, kernel, gated): + features = [] + for child in child_indices: + signal = targets[child] + if gated: + signal = signal * (hiddens[child] > 0).to(signal.dtype) + if kernel == 1: + features.append(signal.permute(0, 2, 3, 1).reshape( + -1, signal.shape[1])) + else: + features.append(F.unfold( + signal, kernel_size=3, padding=1).transpose(1, 2).reshape( + -1, signal.shape[1] * 9)) + return torch.cat(features, dim=1).to(torch.float64) + + +def flat_target(target): + return target.permute(0, 2, 3, 1).reshape( + -1, target.shape[1]).to(torch.float64) + + +def summarize(values): + return {"per_layer": values, "early_third_mean": sum(values) / len(values)} + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--device", default="cpu") + parser.add_argument("--data_dir", default=DATA_DIR) + parser.add_argument("--fit_examples", type=int, default=128) + parser.add_argument("--eval_examples", type=int, default=128) + parser.add_argument( + "--out", default="results/oral_a_hierarchical_oracle.json") + args = parser.parse_args() + if (args.fit_examples, args.eval_examples) != (128, 128): + raise ValueError("the audited hierarchical oracle is fixed to 128/128") + torch.manual_seed(0) + train, _, _, _, _, split = get_cifar_image_splits( + batch_size=256, 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) + fit_end = args.fit_examples + eval_end = fit_end + args.eval_examples + h_train, q_train, c_train, _ = exact_batch( + net, train.x[:fit_end], train.y[:fit_end]) + h_eval, q_eval, c_eval, _ = exact_batch( + net, train.x[fit_end:eval_end], train.y[fit_end:eval_end]) + for parameter in parameters: + parameter.requires_grad_(False) + h_train = [value.to(torch.float64) for value in h_train] + h_eval = [value.to(torch.float64) for value in h_eval] + q_train = [value.to(torch.float64) for value in q_train] + q_eval = [value.to(torch.float64) for value in q_eval] + c_train = c_train.to(torch.float64) + c_eval = c_eval.to(torch.float64) + + reports = { + "direct_output_channel_gated": [], + "downstream_activation_context": [], + "hierarchical_1x1": [], + "hierarchical_1x1_gated": [], + "hierarchical_3x3": [], + "hierarchical_3x3_gated": [], + } + energies = {key: [] for key in reports} + for layer, children in enumerate(EARLY_CHILDREN): + prediction = direct_context_oracle( + c_train, c_eval, h_train[layer], h_eval[layer], q_train[layer]) + reports["direct_output_channel_gated"].append( + cosine(prediction, q_eval[layer])) + energies["direct_output_channel_gated"].append( + energy_ratio(prediction, q_eval[layer])) + + prediction = activation_context_oracle( + c_train, c_eval, h_train[layer], h_eval[layer], + h_train[children[0]], h_eval[children[0]], q_train[layer]) + reports["downstream_activation_context"].append( + cosine(prediction, q_eval[layer])) + energies["downstream_activation_context"].append( + energy_ratio(prediction, q_eval[layer])) + + for kernel, gated, key in ( + (1, False, "hierarchical_1x1"), + (1, True, "hierarchical_1x1_gated"), + (3, False, "hierarchical_3x3"), + (3, True, "hierarchical_3x3_gated")): + features_train = hierarchical_features( + h_train, q_train, children, kernel, gated) + features_eval = hierarchical_features( + h_eval, q_eval, children, kernel, gated) + weights = solve_ridge( + features_train, flat_target(q_train[layer])) + flat_prediction = features_eval @ weights + prediction = flat_prediction.reshape( + args.eval_examples, 32, 32, q_eval[layer].shape[1]).permute( + 0, 3, 1, 2) + reports[key].append(cosine(prediction, q_eval[layer])) + energies[key].append(energy_ratio(prediction, q_eval[layer])) + + output = { + "protocol": "oral_a_post_failure_hierarchical_oracle_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": { + "fit_examples": args.fit_examples, + "evaluation_examples": args.eval_examples, + "source": "unaugmented first training-prefix examples", + "separate_batchnorm_graphs": True, "test_examples_touched": 0, + "layers": list(range(6)), + "residual_dag_children": [list(value) for value in EARLY_CHILDREN], + }, + "cosine": {key: summarize(value) for key, value in reports.items()}, + "prediction_energy_over_target_energy": { + key: summarize(value) for key, value in energies.items()}, + "fit": { + "hierarchical_relative_ridge": 1e-6, + "ridge_role": "fixed numerical conditioning, not endpoint selection", + }, + "interpretation": { + "status": "post_failure_oracle_not_a_trainable_method_or_gate", + "direct_output_channel_gated": "current output-error-only family", + "downstream_activation_context": ( + "current family plus ordinary next-node activation map"), + "hierarchical": ( + "local linear feedback from exact child error fields following " + "the true residual DAG; gated variants multiply by child ReLU state"), + }, + } + 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