summaryrefslogtreecommitdiff
path: root/experiments/audit_native_baselines.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 05:36:17 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 05:36:17 -0500
commit46517d604a8caa23242f209f809f159044d9878b (patch)
treee6f76e2b042490bf6e72544eb6776db415eb48ba /experiments/audit_native_baselines.py
parentb329856d15b91a6f31fb75f010f0a550cddd0ee8 (diff)
baselines: freeze strict C4 audit gate
Diffstat (limited to 'experiments/audit_native_baselines.py')
-rw-r--r--experiments/audit_native_baselines.py236
1 files changed, 236 insertions, 0 deletions
diff --git a/experiments/audit_native_baselines.py b/experiments/audit_native_baselines.py
new file mode 100644
index 0000000..8bae59b
--- /dev/null
+++ b/experiments/audit_native_baselines.py
@@ -0,0 +1,236 @@
+#!/usr/bin/env python3
+"""Strict C4 audit and deterministic table for native author-code baselines."""
+import argparse
+import hashlib
+import json
+import math
+import os
+import re
+
+
+ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
+DEFAULT_PROTOCOLS = os.path.join(
+ ROOT, "experiments", "native_baseline_protocols.json")
+PROTOCOL_IDS = (
+ "burstccn_cifar10_symplastic_paper_seed0",
+ "dualprop_cifar10_vgg16_a0_b0p1_seed1988",
+)
+SHA256_RE = re.compile(r"^[0-9a-f]{64}$")
+GIT_REVISION_RE = re.compile(r"^[0-9a-f]{40,64}$")
+
+
+def file_hash(path):
+ digest = hashlib.sha256()
+ with open(path, "rb") as handle:
+ for chunk in iter(lambda: handle.read(1024 * 1024), b""):
+ digest.update(chunk)
+ return digest.hexdigest()
+
+
+def require(condition, message):
+ if not condition:
+ raise ValueError(message)
+
+
+def number(value, label, low=None, high=None):
+ require(isinstance(value, (int, float)) and math.isfinite(value),
+ f"{label} must be finite, found {value!r}")
+ if low is not None:
+ require(value >= low, f"{label} must be >= {low}, found {value}")
+ if high is not None:
+ require(value <= high, f"{label} must be <= {high}, found {value}")
+ return float(value)
+
+
+def accuracy(metrics, key, label):
+ return number(metrics.get(key), f"{label}.{key}", 0.0, 100.0)
+
+
+def verify_artifacts(records, label):
+ require(isinstance(records, list) and records, f"{label} must be nonempty")
+ for index, record in enumerate(records):
+ prefix = f"{label}[{index}]"
+ require(isinstance(record.get("path"), str) and record["path"],
+ f"{prefix}.path is missing")
+ require(isinstance(record.get("bytes"), int) and record["bytes"] > 0,
+ f"{prefix}.bytes must be positive")
+ require(bool(SHA256_RE.fullmatch(str(record.get("sha256", "")))),
+ f"{prefix}.sha256 is invalid")
+
+
+def validate_common(row, protocol_id, protocol):
+ label = protocol_id
+ require(row.get("schema_version") == 1, f"{label}: unsupported schema")
+ require(row.get("protocol_id") == protocol_id, f"{label}: protocol ID mismatch")
+ require(row.get("protocol") == protocol, f"{label}: frozen protocol drift")
+ require(row.get("method") == protocol["method"], f"{label}: method mismatch")
+ require(row.get("dataset") == protocol["dataset"], f"{label}: dataset mismatch")
+ require(row.get("seed") == protocol["seed"], f"{label}: seed mismatch")
+ audit = row.get("audit", {})
+ for key in ("complete", "finite", "source_revision_match",
+ "source_tracked_clean", "selection_semantics_frozen"):
+ require(audit.get(key) is True, f"{label}: audit.{key} is not true")
+ require(audit.get("expected_epochs") == protocol["expected_epochs"],
+ f"{label}: expected epoch audit mismatch")
+ source = row.get("source_provenance", {})
+ require(source.get("revision") == protocol["source_revision"],
+ f"{label}: author source revision mismatch")
+ require(source.get("tracked_dirty") is False,
+ f"{label}: author source is not tracked-clean")
+ importer = row.get("importer_provenance", {})
+ require(bool(GIT_REVISION_RE.fullmatch(str(importer.get("revision", "")))),
+ f"{label}: importer revision is invalid")
+ require(importer.get("tracked_dirty") is False,
+ f"{label}: importer source was not tracked-clean")
+ environment = row.get("runtime_environment", {})
+ freeze = environment.get("pip_freeze")
+ require(isinstance(freeze, list) and freeze, f"{label}: package freeze missing")
+ freeze_text = ("\n".join(freeze) + "\n").encode()
+ require(hashlib.sha256(freeze_text).hexdigest()
+ == environment.get("pip_freeze_sha256"),
+ f"{label}: package freeze hash mismatch")
+ gpu = row.get("gpu", {})
+ for key, expected in protocol["expected_gpu"].items():
+ require(gpu.get(key) == expected,
+ f"{label}: GPU {key} expected {expected!r}, found {gpu.get(key)!r}")
+ verify_artifacts(row.get("dataset_artifacts"), f"{label}.dataset_artifacts")
+ dataset_hashes = {item["sha256"] for item in row["dataset_artifacts"]}
+ require(set(protocol["required_dataset_sha256"]).issubset(dataset_hashes),
+ f"{label}: required dataset hashes missing")
+ verify_artifacts(row.get("run_artifacts"), f"{label}.run_artifacts")
+ metrics = row.get("metrics", {})
+ require(metrics.get("epochs_observed") == protocol["expected_epochs"],
+ f"{label}: incomplete metric history")
+ number(metrics.get("wall_s"), f"{label}.wall_s", 0.0)
+ require(metrics["wall_s"] > 0, f"{label}: wall time must be positive")
+ require(isinstance(metrics.get("wall_definition"), str)
+ and metrics["wall_definition"], f"{label}: wall definition missing")
+
+
+def validate_burst(row, protocol):
+ label = row["protocol_id"]
+ validate_common(row, label, protocol)
+ metrics = row["metrics"]
+ require(metrics.get("optimization_epochs") == 399,
+ f"{label}: author-loop optimization epoch count mismatch")
+ require(metrics.get("initial_evaluation_epoch") == 0,
+ f"{label}: epoch-0 evaluation missing")
+ require(metrics.get("primary_final_epoch") == 399,
+ f"{label}: final epoch mismatch")
+ require(metrics.get("test_evaluations") == 400,
+ f"{label}: native per-epoch test count mismatch")
+ for key in (
+ "primary_final_test_accuracy_percent",
+ "final_validation_accuracy_percent",
+ "validation_selected_validation_accuracy_percent",
+ "validation_selected_test_accuracy_percent",
+ "test_selected_test_accuracy_percent"):
+ accuracy(metrics, key, label)
+ require(0 <= metrics.get("validation_selected_epoch", -1) <= 399,
+ f"{label}: validation-selected epoch invalid")
+ require(0 <= metrics.get("test_selected_epoch", -1) <= 399,
+ f"{label}: test-selected epoch invalid")
+ require(metrics["test_selected_test_accuracy_percent"]
+ >= metrics["primary_final_test_accuracy_percent"],
+ f"{label}: test-selected result is below final result")
+
+
+def validate_dualprop(row, protocol):
+ label = row["protocol_id"]
+ validate_common(row, label, protocol)
+ metrics = row["metrics"]
+ require(metrics.get("final_epoch") == 130, f"{label}: final epoch mismatch")
+ require(metrics.get("test_evaluations") == 1,
+ f"{label}: expected a single test evaluation")
+ require(metrics.get("test_selected_test_accuracy_percent") is None,
+ f"{label}: test-selected metric must be unavailable")
+ for key in (
+ "final_train_accuracy_percent",
+ "final_validation_accuracy_percent",
+ "validation_selected_validation_accuracy_percent",
+ "validation_selected_test_accuracy_percent",
+ "primary_test_accuracy_percent",
+ "primary_test_top5_accuracy_percent"):
+ accuracy(metrics, key, label)
+ require(1 <= metrics.get("validation_selected_epoch", 0) <= 130,
+ f"{label}: validation-selected epoch invalid")
+ require(metrics["primary_test_accuracy_percent"]
+ == metrics["validation_selected_test_accuracy_percent"],
+ f"{label}: primary is not the validation-selected checkpoint")
+ component_sum = (metrics.get("training_wall_s", 0)
+ + metrics.get("validation_wall_s", 0)
+ + metrics.get("test_wall_s", 0))
+ require(math.isclose(component_sum, metrics["wall_s"], rel_tol=1e-12,
+ abs_tol=1e-9),
+ f"{label}: component wall times do not sum to total")
+
+
+def render(rows, protocols):
+ burst = rows[PROTOCOL_IDS[0]]
+ dual = rows[PROTOCOL_IDS[1]]
+ bm, dm = burst["metrics"], dual["metrics"]
+ bp, dp = protocols[PROTOCOL_IDS[0]], protocols[PROTOCOL_IDS[1]]
+ return f"""# Native author-code baseline audit
+
+Both records passed the strict C4 artifact, source, protocol, environment,
+dataset, completeness, finiteness, selection, and cost-definition checks.
+These are method-native reproductions, not equal-compute comparisons with SDIL.
+
+| method | reproduced seed | primary test (%) | validation-selected test (%) | test-selected test (%) | published (%) | audited wall (s) |
+|:--|--:|--:|--:|--:|--:|--:|
+| BurstCCN | {bp['seed']} | {bm['primary_final_test_accuracy_percent']:.3f} | {bm['validation_selected_test_accuracy_percent']:.3f} | {bm['test_selected_test_accuracy_percent']:.3f} | {bp['published_accuracy_percent_mean']:.2f} ± {bp['published_accuracy_percent_sd']:.2f} | {bm['wall_s']:.1f} |
+| Dual Prop | {dp['seed']} | {dm['primary_test_accuracy_percent']:.3f} | {dm['validation_selected_test_accuracy_percent']:.3f} | — | {dp['published_accuracy_percent_mean']:.2f} ± {dp['published_accuracy_percent_sd']:.2f} | {dm['wall_s']:.1f} |
+
+BurstCCN's primary metric is its final epoch. Its validation-selected metric is
+the test accuracy from the minimum-validation-error epoch; its best test value
+is explicitly test-selected because the native trainer evaluates test every
+epoch. The `n_epochs=400` author configuration performs an epoch-0 evaluation
+and 399 optimization epochs. Dual Prop evaluates test once after restoring the
+best-validation checkpoint, so its primary and validation-selected values are
+identical and no test-selected value exists.
+
+The BurstCCN wall value is W&B run time through the final epoch log. The Dual
+Prop wall value sums author-measured train, validation, and final test time and
+excludes setup, checkpoint I/O, and diagnostic plotting. One reproduced seed
+establishes executable fidelity; published multi-seed uncertainty is retained
+for context and is not replaced by a one-seed error bar.
+"""
+
+
+def main():
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--protocols", default=DEFAULT_PROTOCOLS)
+ parser.add_argument("--results", default=os.path.join(ROOT, "results", "native"))
+ parser.add_argument("--output", default=os.path.join(
+ ROOT, "results", "native_baseline_audit.md"))
+ args = parser.parse_args()
+ with open(args.protocols) as handle:
+ ledger = json.load(handle)
+ require(ledger.get("schema_version") == 1, "unsupported protocol ledger schema")
+ protocols = ledger["protocols"]
+ rows = {}
+ for protocol_id in PROTOCOL_IDS:
+ path = os.path.join(args.results, protocol_id + ".json")
+ require(os.path.isfile(path), f"missing native result: {path}")
+ with open(path) as handle:
+ rows[protocol_id] = json.load(handle)
+ validate_burst(rows[PROTOCOL_IDS[0]], protocols[PROTOCOL_IDS[0]])
+ validate_dualprop(rows[PROTOCOL_IDS[1]], protocols[PROTOCOL_IDS[1]])
+ text = render(rows, protocols)
+ os.makedirs(os.path.dirname(os.path.abspath(args.output)), exist_ok=True)
+ with open(args.output, "w") as handle:
+ handle.write(text)
+ manifest = {
+ "strict": True,
+ "c4_pass": True,
+ "protocol_ledger_sha256": file_hash(args.protocols),
+ "result_sha256": {
+ protocol_id: file_hash(os.path.join(args.results, protocol_id + ".json"))
+ for protocol_id in PROTOCOL_IDS
+ },
+ }
+ print(json.dumps(manifest, indent=2, sort_keys=True))
+
+
+if __name__ == "__main__":
+ main()