diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-06 14:32:26 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-08-06 14:32:26 -0500 |
| commit | 640522076c770d0746dc58dde868586ace568752 (patch) | |
| tree | a10d4f6364aeaf8caa287a2459b82608b42efce1 /experiments | |
| parent | 042c43f36d7ebac22f4b3ff078b1de42445de0b3 (diff) | |
experiment: freeze no-KP causal capture runner
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/analyze_oral_a_v5_calibration.py | 102 | ||||
| -rw-r--r-- | experiments/oral_a_v5_calibration_screen.py | 207 |
2 files changed, 309 insertions, 0 deletions
diff --git a/experiments/analyze_oral_a_v5_calibration.py b/experiments/analyze_oral_a_v5_calibration.py new file mode 100644 index 0000000..2844743 --- /dev/null +++ b/experiments/analyze_oral_a_v5_calibration.py @@ -0,0 +1,102 @@ +#!/usr/bin/env python3 +"""Validate and gate the frozen no-KP layerwise causal-bootstrap screen.""" +import argparse +import json +import math +import os + + +SPLIT_HASH = "8328b206a97c420e49e54e3eca4abe3274c4756b084355784ea3fb8059e4515b" + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument( + "--input", default="results/oral_a_v5_calibration/result.json") + parser.add_argument( + "--out", default="results/oral_a_v5_calibration_gate.json") + args = parser.parse_args() + with open(args.input) as handle: + record = json.load(handle) + if record.get("protocol") != "oral_a_v5_layerwise_causal_bootstrap_capture_v1": + raise ValueError("unexpected V5 protocol") + expected = { + "depth": 20, "width": 16, "seed": 0, "loader_seed": 0, + "batch_size": 128, "train_limit": 10000, + "val_examples": 5000, "split_seed": 2027, + "normalization": "batchnorm", "residual_scale": 1.0, + "feedback_scale": 1.0, "sigma": 0.01, "eta_A": 0.1, + "perturb_seed": 5000, "sweeps": 20, "alignment_probe": 64, + "calibration_augmentation": False, + } + if record.get("settings") != expected: + raise ValueError("V5 settings drift") + if record["provenance"]["git_tracked_dirty"]: + raise ValueError("V5 result came from a tracked-dirty tree") + if record["split"]["validation_index_sha256"] != SPLIT_HASH: + raise ValueError("V5 split drift") + if record["test_examples_touched"] or record["validation_endpoints_observed"]: + raise ValueError("V5 touched a held-out endpoint") + work = record["work"] + audit = record["method_audit"] + fixed = record["fixed_hfa"] + learned = record["learned_lcb"] + finite_metrics = [ + fixed["early_third_alignment"], fixed["all_layer_alignment"], + learned["early_third_alignment"], learned["all_layer_alignment"], + learned["min_feedback_forward_norm_ratio"], + learned["max_feedback_forward_norm_ratio"], + ] + checks = { + "finite": bool(record["finite"]) + and all(math.isfinite(value) for value in finite_metrics), + "exactly_380_edge_events": work["edge_events"] == 380, + "exactly_760_batch_loss_queries": ( + work["logical_batch_loss_queries"] == 760), + "exactly_48640_per_example_observations": ( + work["per_example_causal_observations"] == 48640), + "forward_state_bitwise_fixed": ( + audit["forward_state_max_absolute_difference"] == 0.0), + "zero_forward_weight_reads_in_update": ( + audit["forward_weight_reads_in_feedback_update"] == 0), + "zero_reverse_mode_learning_operations": ( + audit["reverse_mode_learning_operations"] == 0), + "early_third_at_least_0.10": ( + learned["early_third_alignment"] >= 0.10), + "all_layer_at_least_0.20": ( + learned["all_layer_alignment"] >= 0.20), + "early_gain_over_fixed_hfa_at_least_0.08": ( + learned["early_third_alignment"] + - fixed["early_third_alignment"] >= 0.08), + "feedback_norm_ratios_in_0.1_to_3": ( + learned["min_feedback_forward_norm_ratio"] >= 0.1 + and learned["max_feedback_forward_norm_ratio"] <= 3.0), + } + output = { + "protocol": "oral_a_v5_layerwise_causal_bootstrap_gate_v1", + "status": "passed" if all(checks.values()) else "failed", + "checks": checks, + "fixed_hfa": fixed, + "learned_lcb": learned, + "work": work, + "source_commit": record["provenance"]["git_commit"], + "source_result": args.input, + "conditional_short_task_gate_open": all(checks.values()), + "confirmation_test_seeds_touched": False, + "review_score_before": 5, + "review_score_after": 5, + "score_change_rule": "causal capture alone cannot raise score", + } + 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({ + "status": output["status"], "checks": checks, + "fixed_hfa": fixed, "learned_lcb": learned, + }, indent=2)) + + +if __name__ == "__main__": + main() + diff --git a/experiments/oral_a_v5_calibration_screen.py b/experiments/oral_a_v5_calibration_screen.py new file mode 100644 index 0000000..133c18d --- /dev/null +++ b/experiments/oral_a_v5_calibration_screen.py @@ -0,0 +1,207 @@ +#!/usr/bin/env python3 +"""Run the frozen no-KP layerwise causal-bootstrap capture screen.""" +import argparse +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, + conv_hierarchical_alignment_report, + layerwise_causal_bootstrap_sweep) +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 summarize_alignment(report): + values = report["teaching_negative_gradient_cosine"] + early = max(1, len(values) // 3) + ratios = report["feedback_forward_norm_ratio"] + cosines = report["feedback_forward_cosine"] + return { + "per_layer": values, + "early_third_alignment": sum(values[:early]) / early, + "all_layer_alignment": sum(values) / len(values), + "mean_feedback_forward_cosine": sum(cosines) / len(cosines), + "min_feedback_forward_norm_ratio": min(ratios), + "max_feedback_forward_norm_ratio": max(ratios), + "feedback_forward_cosine": cosines, + "feedback_forward_norm_ratio": ratios, + } + + +def forward_state(net): + return [value.clone() for value in ( + net.W + net.gamma + net.beta + net.running_mean + net.running_var + + net.mW + net.mgamma + net.mbeta + + [net.W_out, net.b_out, net.mW_out, net.mb_out])] + + +def main(): + parser = argparse.ArgumentParser() + parser.add_argument("--device", default="cuda") + parser.add_argument("--data_dir", default=DATA_DIR) + parser.add_argument("--out", default="results/oral_a_v5_calibration/result.json") + args = parser.parse_args() + settings = { + "depth": 20, "width": 16, "seed": 0, "loader_seed": 0, + "batch_size": 128, "train_limit": 10000, + "val_examples": 5000, "split_seed": 2027, + "normalization": "batchnorm", "residual_scale": 1.0, + "feedback_scale": 1.0, "sigma": 0.01, "eta_A": 0.1, + "perturb_seed": 5000, "sweeps": 20, "alignment_probe": 64, + "calibration_augmentation": False, + } + torch.manual_seed(settings["seed"]) + if str(args.device).startswith("cuda"): + if not torch.cuda.is_available(): + raise RuntimeError("CUDA requested but unavailable") + torch.cuda.manual_seed_all(settings["seed"]) + torch.cuda.reset_peak_memory_stats(torch.device(args.device)) + train, _, _, input_shape, n_out, split = get_cifar_image_splits( + batch_size=settings["batch_size"], data_dir=args.data_dir, + device=args.device, train_limit=settings["train_limit"], + val_examples=settings["val_examples"], split_seed=settings["split_seed"], + loader_seed=settings["loader_seed"], augment_train=False) + if input_shape != (3, 32, 32) or n_out != 10: + raise AssertionError("unexpected CIFAR dimensions") + net = CIFARHierarchicalFAResNet( + depth=settings["depth"], base_width=settings["width"], + n_classes=10, device=args.device, seed=settings["seed"], + residual_scale=settings["residual_scale"], + normalization=settings["normalization"], + feedback_scale=settings["feedback_scale"]) + audit_x = train.x[:settings["alignment_probe"]] + audit_y = train.y[:settings["alignment_probe"]] + fixed = summarize_alignment( + conv_hierarchical_alignment_report(net, audit_x, audit_y)) + state_before = forward_state(net) + generator = torch.Generator(device=torch.device(args.device)).manual_seed( + settings["perturb_seed"]) + + if str(args.device).startswith("cuda"): + torch.cuda.synchronize(torch.device(args.device)) + start = time.time() + sweeps = [] + for sweep_index in range(settings["sweeps"]): + start_index = sweep_index * settings["batch_size"] + stop_index = start_index + settings["batch_size"] + metric = layerwise_causal_bootstrap_sweep( + net, train.x[start_index:stop_index], train.y[start_index:stop_index], + sigma=settings["sigma"], eta=settings["eta_A"], + generator=generator) + sweeps.append({ + key: value for key, value in metric.items() if key != "edges"}) + print(json.dumps({"sweep": sweep_index + 1, **sweeps[-1]}), flush=True) + if str(args.device).startswith("cuda"): + torch.cuda.synchronize(torch.device(args.device)) + wall_seconds = time.time() - start + state_after = forward_state(net) + forward_state_max_difference = max( + float((before - after).abs().max()) + for before, after in zip(state_before, state_after)) + learned = summarize_alignment( + conv_hierarchical_alignment_report(net, audit_x, audit_y)) + + total_events = sum(value["events"] for value in sweeps) + total_queries = sum(value["logical_batch_loss_queries"] for value in sweeps) + total_observations = sum( + value["per_example_causal_observations"] for value in sweeps) + batch = settings["batch_size"] + clean_forward_examples = settings["sweeps"] * batch + perturbation_forward_examples = 2 * total_events * batch + teaching_macs = total_events * batch * net.apical_macs_per_example + work = { + "edge_events": total_events, + "logical_batch_loss_queries": total_queries, + "per_example_causal_observations": total_observations, + "per_example_cross_entropy_terms": 2 * total_observations, + "clean_forward_examples": clean_forward_examples, + "perturbation_forward_examples": perturbation_forward_examples, + "forward_macs": ((clean_forward_examples + perturbation_forward_examples) + * net.forward_macs_per_example), + "hierarchical_teaching_and_local_correlation_macs_estimate": teaching_macs, + } + work["total_macs_estimate"] = ( + work["forward_macs"] + + work["hierarchical_teaching_and_local_correlation_macs_estimate"]) + finite_values = [ + fixed["early_third_alignment"], fixed["all_layer_alignment"], + learned["early_third_alignment"], learned["all_layer_alignment"], + learned["min_feedback_forward_norm_ratio"], + learned["max_feedback_forward_norm_ratio"], + ] + [value[key] for value in sweeps for key in ( + "mean_field_prediction_target_cosine", + "mean_parameter_update_rms", "max_parameter_update_rms")] + output = { + "schema_version": 1, + "protocol": "oral_a_v5_layerwise_causal_bootstrap_capture_v1", + "settings": settings, + "provenance": provenance(), + "split": split, + "architecture": { + "family": "CIFAR 6n+2 ResNet, option-A shortcuts", + "forward_parameters": net.n_forward_parameters, + "adaptive_feedback_parameters": net.n_fixed_feedback_parameters, + "forward_macs_per_example": net.forward_macs_per_example, + "feedback_macs_per_example": net.apical_macs_per_example, + }, + "method_audit": { + "forward_weight_reads_in_feedback_update": 0, + "reverse_mode_learning_operations": 0, + "causal_query_normalization_state": "evaluation_running_statistics", + "ordinary_task_normalization_state": "not_run_forward_frozen", + "forward_state_max_absolute_difference": ( + forward_state_max_difference), + }, + "fixed_hfa": fixed, + "learned_lcb": learned, + "sweeps": sweeps, + "work": work, + "wall_seconds": wall_seconds, + "finite": all(math.isfinite(value) for value in finite_values), + "test_examples_touched": 0, + "validation_endpoints_observed": 0, + "hardware": { + "device": str(args.device), "torch_version": torch.__version__, + "cuda_visible_devices": os.environ.get("CUDA_VISIBLE_DEVICES"), + "cuda_device_name": (torch.cuda.get_device_name(torch.device(args.device)) + if str(args.device).startswith("cuda") else None), + "peak_memory_allocated_bytes": ( + torch.cuda.max_memory_allocated(torch.device(args.device)) + if str(args.device).startswith("cuda") else None), + }, + } + 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({ + "fixed_hfa": fixed, "learned_lcb": learned, "work": work, + "finite": output["finite"], "wall_seconds": wall_seconds, + "out": args.out, + }, indent=2), flush=True) + + +if __name__ == "__main__": + main() + |
