summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 13:03:46 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 13:03:46 -0500
commit45d2b57310d556c4d1e78ee185279eda14189585 (patch)
tree62c89f2b476b4e38fa0d89bec3f6a6b977ad495d /experiments
parent6fa80efc1c53905b6328c5b44df3db4e4da34dfe (diff)
analysis: add hierarchical feedback oracle audit
Diffstat (limited to 'experiments')
-rw-r--r--experiments/analyze_oral_a_hierarchical_oracle.py288
1 files changed, 288 insertions, 0 deletions
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()
+