experiment: implement AttnRes forward training runner

This commit is contained in:
wuyang
2026-07-30 15:14:15 +08:00
parent d81aacfc68
commit 7ea91caabb
11 changed files with 17842 additions and 1 deletions
+684
View File
@@ -0,0 +1,684 @@
#!/usr/bin/env python3
"""Aggregate and gate preregistered Round 08 forward-training results."""
from __future__ import annotations
import argparse
import copy
import hashlib
import json
import math
import statistics
from pathlib import Path
from typing import Any, Iterable
PROTOCOL_ID = "llm-atlas-k3-attnres-forward-training-v1"
PARENT_PROTOCOL_ID = "llm-atlas-k3-attnres-gradient-scale-v1"
METRICS = ("spike_contrast", "peak_normalized")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--manifest", type=Path, required=True)
parser.add_argument("--formal", type=Path, action="append", required=True)
parser.add_argument("--replay", type=Path, required=True)
parser.add_argument("--reference-dir", type=Path, required=True)
parser.add_argument("--aggregate-output", type=Path, required=True)
parser.add_argument("--compact-output", type=Path, required=True)
parser.add_argument("--reproduction-output", type=Path, required=True)
return parser.parse_args()
def canonical_sha256(value: Any) -> str:
return hashlib.sha256(
json.dumps(
value, ensure_ascii=False, sort_keys=True, separators=(",", ":")
).encode()
).hexdigest()
def file_sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def mean(values: Iterable[float]) -> float:
return statistics.fmean(values)
def read_result(path: Path, expected_protocol: str) -> dict[str, Any]:
value = json.loads(path.read_text())
if value.get("protocol_id") != expected_protocol:
raise RuntimeError(f"protocol mismatch: {path}")
expected = value.get("canonical_sha256_without_self")
payload = {
key: item
for key, item in value.items()
if key != "canonical_sha256_without_self"
}
if not isinstance(expected, str) or canonical_sha256(payload) != expected:
raise RuntimeError(f"canonical self-hash mismatch: {path}")
return value
def exactly_one(values: list[dict[str, Any]], step: int) -> dict[str, Any]:
matches = [value for value in values if value["step"] == step]
if len(matches) != 1:
raise RuntimeError(f"step {step} missing or duplicated")
return matches[0]
def spectrum_metrics(
diagnostic: dict[str, Any],
spike_layers: tuple[int, ...],
epsilon: float,
) -> dict[str, Any]:
values = [
float(value)
for value in diagnostic["activation_grad_rms_by_block"]
]
if len(values) != 32:
raise RuntimeError("activation-gradient spectrum must have 32 layers")
if any(not math.isfinite(value) or value <= epsilon for value in values):
raise RuntimeError("activation-gradient spectrum is non-finite/non-positive")
spike_indices = {layer - 1 for layer in spike_layers}
spike_values = [
value for index, value in enumerate(values) if index in spike_indices
]
reference_values = [
value for index, value in enumerate(values) if index not in spike_indices
]
spike_mean = mean(spike_values)
reference_mean = mean(reference_values)
global_mean = mean(values)
contrast = spike_mean / reference_mean
peak = max(values) / global_mean
if any(
not math.isfinite(value) or value <= epsilon
for value in (spike_mean, reference_mean, contrast, peak)
):
raise RuntimeError("derived spike metric is non-finite/non-positive")
ordered = sorted(range(32), key=lambda index: (-values[index], index))
return {
"values": values,
"normalized": [value / global_mean for value in values],
"spike_mean": spike_mean,
"reference_mean": reference_mean,
"global_mean": global_mean,
"spike_contrast": contrast,
"peak_normalized": peak,
"peak_layer_1based": ordered[0] + 1,
"top_five_layers_1based": [index + 1 for index in ordered[:5]],
}
def final_bpc(value: dict[str, Any], step: int) -> float:
result = float(exactly_one(value["evaluations"], step)["bits_per_byte"])
if not math.isfinite(result):
raise RuntimeError("final validation BPC is non-finite")
return result
def stable_environment(value: dict[str, Any]) -> dict[str, Any]:
keys = (
"cublas_workspace_config",
"deterministic_algorithms",
"autocast",
"compile",
)
return {key: value["environment"][key] for key in keys}
def pairing_checks(
run: dict[str, Any], reference: dict[str, Any]
) -> dict[str, bool]:
manifest_fields = (
"formal_schedule_sha256",
"validation_tensor_sha256",
"diagnostic_tensor_sha256",
"input_gate_tensor_hashes",
)
checks = {
"seed": run["seed"] == reference["seed"],
"architecture": (
run["architecture"] == reference["architecture"] == "block"
),
"depth": run["depth"] == reference["depth"] == 32,
"steps": run["steps"] == reference["steps"] == 8000,
"batch_size": run["batch_size"] == reference["batch_size"] == 32,
"initial_public_parameters": (
run["hashes"]["initial_public_parameters"]
== reference["hashes"]["initial_public_parameters"]
),
"initial_mixer_parameters": (
run["hashes"]["initial_mixer_parameters"]
== reference["hashes"]["initial_mixer_parameters"]
),
"model_topology": run["model"] == reference["model"],
"optimizer_hyperparameters": (
run["optimizer"] == reference["optimizer"]
),
"scientific_environment": (
stable_environment(run) == stable_environment(reference)
),
}
for field in manifest_fields:
checks[f"manifest.{field}"] = (
run["manifest"][field] == reference["manifest"][field]
)
return checks
def scientific_replay_payload(value: dict[str, Any]) -> dict[str, Any]:
payload = copy.deepcopy(value)
for key in (
"run_kind",
"timing",
"canonical_sha256_without_self",
"parent_runner_canonical_sha256",
):
payload.pop(key, None)
payload["manifest"].pop("path", None)
payload["study_manifest"].pop("path", None)
payload["environment"] = stable_environment(value)
return payload
def quality_gate(
variant_runs: dict[int, dict[str, Any]],
references: dict[int, dict[str, Any]],
*,
step: int,
per_seed_maximum: float,
mean_maximum: float,
) -> dict[str, Any]:
per_seed = {}
for seed, run in sorted(variant_runs.items()):
variant_bpc = final_bpc(run, step)
reference_bpc = final_bpc(references[seed], step)
delta = variant_bpc - reference_bpc
per_seed[str(seed)] = {
"variant_bpc": variant_bpc,
"reference_bpc": reference_bpc,
"delta_bpc": delta,
"passed": delta <= per_seed_maximum,
}
mean_delta = mean(item["delta_bpc"] for item in per_seed.values())
per_seed_passed = all(item["passed"] for item in per_seed.values())
mean_passed = mean_delta <= mean_maximum
return {
"passed": per_seed_passed and mean_passed,
"passed_checks": (
sum(item["passed"] for item in per_seed.values())
+ int(mean_passed)
),
"required_checks": 4,
"per_seed_maximum": per_seed_maximum,
"mean_maximum": mean_maximum,
"mean_delta_bpc": mean_delta,
"mean_passed": mean_passed,
"per_seed": per_seed,
}
def variant_effect(
variant: str,
runs: dict[int, dict[str, Any]],
references: dict[int, dict[str, Any]],
metrics_by_cell: dict[tuple[str, int, int], dict[str, Any]],
*,
step: int,
threshold: float,
quality: dict[str, Any],
) -> dict[str, Any]:
cells = []
for seed in sorted(runs):
candidate = metrics_by_cell[(variant, seed, step)]
reference = metrics_by_cell[("learned_reference", seed, step)]
for metric in METRICS:
reference_value = reference[metric]
candidate_value = candidate[metric]
relative_drop = (
reference_value - candidate_value
) / reference_value
cells.append(
{
"seed": seed,
"metric": metric,
"reference": reference_value,
"variant": candidate_value,
"relative_drop": relative_drop,
"passed": relative_drop >= threshold,
}
)
attenuation_passed = all(cell["passed"] for cell in cells)
return {
"variant": variant,
"threshold": threshold,
"passed_cells": sum(cell["passed"] for cell in cells),
"required_cells": len(cells),
"attenuation_passed": attenuation_passed,
"quality": quality,
"material_response_passed": (
attenuation_passed and quality["passed"]
),
"cells": cells,
}
def interaction_map(
metrics_by_cell: dict[tuple[str, int, int], dict[str, Any]],
seeds: tuple[int, ...],
steps: tuple[int, ...],
) -> dict[str, Any]:
cells = []
for step in steps:
for seed in seeds:
reference = metrics_by_cell[
("learned_reference", seed, step)
]
group6 = metrics_by_cell[
("uniform_group_6_forward", seed, step)
]
group7 = metrics_by_cell[
("uniform_group_7_forward", seed, step)
]
joint = metrics_by_cell[
("uniform_groups_6_7_forward", seed, step)
]
for metric in METRICS:
ref = reference[metric]
effects = {
"group6": math.log(ref / group6[metric]),
"group7": math.log(ref / group7[metric]),
"groups6_7": math.log(ref / joint[metric]),
}
residual = (
effects["groups6_7"]
- effects["group6"]
- effects["group7"]
)
cells.append(
{
"step": step,
"seed": seed,
"metric": metric,
"log_effects": effects,
"interaction_residual": residual,
"relative_drops": {
"group6": (ref - group6[metric]) / ref,
"group7": (ref - group7[metric]) / ref,
"groups6_7": (ref - joint[metric]) / ref,
},
}
)
summaries = []
for step in steps:
for metric in METRICS:
selected = [
cell
for cell in cells
if cell["step"] == step and cell["metric"] == metric
]
residuals = [
cell["interaction_residual"] for cell in selected
]
summaries.append(
{
"step": step,
"metric": metric,
"mean_interaction_residual": mean(residuals),
"minimum": min(residuals),
"maximum": max(residuals),
}
)
return {
"definition": "I67=ln(Xref/X67)-ln(Xref/X6)-ln(Xref/X7)",
"interpretation": (
"descriptive cross-run log-attenuation residual from three "
"independently trained variants; not a causal interaction"
),
"cells": cells,
"summaries": summaries,
}
def environment_metadata(value: dict[str, Any]) -> dict[str, Any]:
return {
key: value["environment"].get(key)
for key in ("gpu", "torch", "cuda", "compute_capability")
}
def write_hashed(path: Path, value: dict[str, Any]) -> None:
value["canonical_sha256_without_self"] = canonical_sha256(value)
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(
json.dumps(value, ensure_ascii=False, indent=2, sort_keys=True) + "\n"
)
def main() -> None:
args = parse_args()
manifest = json.loads(args.manifest.read_text())
if (
manifest["protocol_id"] != PROTOCOL_ID
or manifest["status"] != "frozen-before-model-output"
):
raise RuntimeError("manifest is not the frozen Round 08 contract")
variants = tuple(manifest["variants"].keys())
seeds = tuple(manifest["formal_seeds"])
steps = tuple(manifest["diagnostic_steps"])
primary_step = manifest["primary_step"]
epsilon = manifest["thresholds"]["positive_denominator_epsilon"]
spike_layers = tuple(manifest["spike_layers_1based"])
expected_cells = {(variant, seed) for variant in variants for seed in seeds}
if len(args.formal) != len(expected_cells):
raise RuntimeError("formal path count does not match the 4×3 matrix")
runs: dict[tuple[str, int], dict[str, Any]] = {}
run_paths: dict[tuple[str, int], Path] = {}
pairing: dict[str, Any] = {}
references: dict[int, dict[str, Any]] = {}
reference_paths: dict[int, Path] = {}
for seed in seeds:
path = args.reference_dir / (
f"formal-depth-32-block-seed-{seed}.json"
)
references[seed] = read_result(path, PARENT_PROTOCOL_ID)
reference_paths[seed] = path
for path in args.formal:
value = read_result(path, PROTOCOL_ID)
identity = (value["variant"], value["seed"])
if identity in runs:
raise RuntimeError(f"duplicate formal cell: {identity}")
if (
value["run_kind"] != "formal"
or value["steps"] != manifest["formal_steps"]
or not value["forward_intervention"]["passed"]
):
raise RuntimeError(f"invalid formal cell: {path}")
runs[identity] = value
run_paths[identity] = path
if set(runs) != expected_cells:
raise RuntimeError("formal matrix identities do not match manifest")
for (variant, seed), value in sorted(runs.items()):
checks = pairing_checks(value, references[seed])
if not all(checks.values()):
raise RuntimeError(
f"historical reference pairing failed: "
f"{variant}/{seed}: {checks}"
)
pairing[f"{variant}:{seed}"] = {
"passed": True,
"checks": checks,
"run_environment": environment_metadata(value),
"reference_environment": environment_metadata(references[seed]),
"metadata_equal": (
environment_metadata(value)
== environment_metadata(references[seed])
),
}
replay = read_result(args.replay, PROTOCOL_ID)
replay_contract = manifest["replay"]
if (
replay["run_kind"] != "replay"
or replay["variant"] != replay_contract["variant"]
or replay["seed"] != replay_contract["seed"]
or replay["steps"] != manifest["formal_steps"]
or not replay["forward_intervention"]["passed"]
):
raise RuntimeError("invalid replay identity/audit")
formal_primary = runs[
(replay_contract["variant"], replay_contract["seed"])
]
formal_payload = scientific_replay_payload(formal_primary)
replay_payload = scientific_replay_payload(replay)
replay_exact = formal_payload == replay_payload
if not replay_exact:
raise RuntimeError("primary formal/replay scientific payload mismatch")
metrics_by_cell: dict[tuple[str, int, int], dict[str, Any]] = {}
for seed, reference in references.items():
for step in steps:
metrics_by_cell[("learned_reference", seed, step)] = (
spectrum_metrics(
exactly_one(reference["diagnostics"], step),
spike_layers,
epsilon,
)
)
for (variant, seed), value in runs.items():
if tuple(item["step"] for item in value["diagnostics"]) != steps:
raise RuntimeError(f"diagnostic schedule drift: {variant}/{seed}")
for step in steps:
metrics_by_cell[(variant, seed, step)] = spectrum_metrics(
exactly_one(value["diagnostics"], step),
spike_layers,
epsilon,
)
runs_by_variant = {
variant: {seed: runs[(variant, seed)] for seed in seeds}
for variant in variants
}
qualities = {
variant: quality_gate(
variant_runs,
references,
step=primary_step,
per_seed_maximum=manifest["thresholds"][
"final_bpc_delta_per_seed_maximum"
],
mean_maximum=manifest["thresholds"][
"final_bpc_delta_mean_maximum"
],
)
for variant, variant_runs in runs_by_variant.items()
}
effects = {
variant: variant_effect(
variant,
variant_runs,
references,
metrics_by_cell,
step=primary_step,
threshold=manifest["thresholds"]["material_relative_drop"],
quality=qualities[variant],
)
for variant, variant_runs in runs_by_variant.items()
}
primary = effects[manifest["primary_variant"]]
if primary["attenuation_passed"] and primary["quality"]["passed"]:
status = (
"forward_training_attenuation_established_within_reduced_protocol"
)
elif primary["attenuation_passed"]:
status = "quality_guard_failed"
elif primary["quality"]["passed"]:
status = "attenuation_not_established"
else:
status = "attenuation_and_quality_failed"
secondary = {
variant: (
"secondary_material_response"
if effect["material_response_passed"]
else "secondary_response_not_established"
)
for variant, effect in effects.items()
if variant != manifest["primary_variant"]
}
interaction = interaction_map(metrics_by_cell, seeds, steps)
trajectories = []
final_spectra = []
for variant in ("learned_reference",) + variants:
for seed in seeds:
for step in steps:
record = metrics_by_cell[(variant, seed, step)]
reference = metrics_by_cell[
("learned_reference", seed, step)
]
trajectories.append(
{
"variant": variant,
"seed": seed,
"step": step,
"spike_mean": record["spike_mean"],
"reference_mean": record["reference_mean"],
"spike_contrast": record["spike_contrast"],
"peak_normalized": record["peak_normalized"],
"relative_drop": {
metric: (
reference[metric] - record[metric]
)
/ reference[metric]
for metric in METRICS
},
}
)
final = metrics_by_cell[(variant, seed, primary_step)]
final_spectra.append(
{
"variant": variant,
"seed": seed,
**final,
}
)
input_files = {
"manifest": {
"path": str(args.manifest),
"sha256": file_sha256(args.manifest),
},
"formal": [
{
"variant": variant,
"seed": seed,
"path": str(run_paths[(variant, seed)]),
"sha256": file_sha256(run_paths[(variant, seed)]),
}
for variant, seed in sorted(runs)
],
"references": [
{
"seed": seed,
"path": str(reference_paths[seed]),
"sha256": file_sha256(reference_paths[seed]),
}
for seed in seeds
],
"replay": {
"path": str(args.replay),
"sha256": file_sha256(args.replay),
},
}
aggregate = {
"schema_version": 1,
"protocol_id": PROTOCOL_ID,
"status": status,
"scope": (
"depth-32 reduced Block AttnRes train-time architecture "
"ablation; not a real Kimi-K3 checkpoint result"
),
"primary_step": primary_step,
"spike_layers_1based": list(spike_layers),
"thresholds": manifest["thresholds"],
"input_files": input_files,
"historical_pairing": pairing,
"replay": {
"passed": replay_exact,
"scientific_payload_sha256": canonical_sha256(formal_payload),
"excluded": [
"run_kind",
"timing",
"self hashes",
"manifest path strings",
"GPU/version metadata",
],
},
"primary": primary,
"secondary_status": secondary,
"effects": effects,
"interaction": interaction,
"trajectories": trajectories,
"final_spectra": final_spectra,
"processed_target_bytes": manifest["new_target_bytes"],
"historical_reference_target_bytes": (
manifest["historical_reference_target_bytes"]
),
"reporting_boundary": (
"C can change through spike-window numerator and the 27-layer "
"reference denominator; layers 26-28 are intervened but belong "
"to the denominator."
),
}
write_hashed(args.aggregate_output, aggregate)
compact = {
"schema_version": 1,
"protocol_id": PROTOCOL_ID,
"status": status,
"primary_step": primary_step,
"spike_layers_1based": list(spike_layers),
"thresholds": manifest["thresholds"],
"primary": primary,
"secondary_status": secondary,
"effects": effects,
"interaction": interaction,
"trajectories": trajectories,
"final_spectra": final_spectra,
"replay": aggregate["replay"],
"processed_target_bytes": manifest["new_target_bytes"],
"reporting_boundary": aggregate["reporting_boundary"],
"aggregate_sha256": aggregate["canonical_sha256_without_self"],
}
write_hashed(args.compact_output, compact)
reproduction = {
"schema_version": 1,
"protocol_id": PROTOCOL_ID,
"passed": replay_exact,
"formal_variant": replay_contract["variant"],
"seed": replay_contract["seed"],
"formal_file_sha256": file_sha256(
run_paths[
(replay_contract["variant"], replay_contract["seed"])
]
),
"replay_file_sha256": file_sha256(args.replay),
"scientific_payload_sha256": canonical_sha256(formal_payload),
"excluded_fields": aggregate["replay"]["excluded"],
}
write_hashed(args.reproduction_output, reproduction)
print(
json.dumps(
{
"status": status,
"primary_attenuation": {
"passed_cells": primary["passed_cells"],
"required_cells": primary["required_cells"],
},
"primary_quality": {
"passed_checks": primary["quality"]["passed_checks"],
"required_checks": primary["quality"]["required_checks"],
},
"replay_exact": replay_exact,
"aggregate": str(args.aggregate_output),
"compact": str(args.compact_output),
},
ensure_ascii=False,
indent=2,
)
)
if __name__ == "__main__":
main()