research: lock AttnRes spike diagnostic runner

This commit is contained in:
wuyang
2026-07-30 10:54:08 +08:00
parent 204f599304
commit 54c80fb115
3 changed files with 1316 additions and 0 deletions
+178
View File
@@ -0,0 +1,178 @@
#!/usr/bin/env python3
"""Rebuild the Round 06 scoping table from frozen Round 05 raw files."""
from __future__ import annotations
import argparse
import json
import math
import statistics
from pathlib import Path
from typing import Any, Callable
SEEDS = (2026073001, 2026073002, 2026073003)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument(
"--raw-dir",
type=Path,
default=Path(__file__).resolve().parents[1]
/ "attnres_gradient"
/ "results"
/ "raw",
)
return parser.parse_args()
def pearson(left: list[float], right: list[float]) -> float:
left_mean = statistics.fmean(left)
right_mean = statistics.fmean(right)
numerator = sum(
(x - left_mean) * (y - right_mean)
for x, y in zip(left, right)
)
left_square = sum((value - left_mean) ** 2 for value in left)
right_square = sum((value - right_mean) ** 2 for value in right)
return numerator / math.sqrt(left_square * right_square)
def mixer_metrics(value: dict[str, Any]) -> dict[str, float | int]:
weights = value["mean_weights"]
sources = value["sources"]
return {
"sources": sources,
"latest": weights[-1],
"maximum": max(weights),
"normalized_entropy": (
1.0
if sources == 1
else value["entropy_mean"] / math.log(sources)
),
}
def average(rows: list[dict[str, Any]], key: Callable[[dict[str, Any]], float]) -> float:
return statistics.fmean(key(row) for row in rows)
def main() -> None:
args = parse_args()
rows = []
source_hashes = {}
for seed in SEEDS:
path = (
args.raw_dir
/ f"formal-depth-32-block-seed-{seed}.json"
)
run = json.loads(path.read_text())
source_hashes[str(seed)] = run["canonical_sha256_without_self"]
diagnostic = next(
item for item in run["diagnostics"] if item["step"] == 8000
)
gradients = diagnostic["activation_grad_rms_by_block"]
gradient_mean = statistics.fmean(gradients)
for layer in range(32):
rows.append(
{
"seed": seed,
"layer": layer + 1,
"group": layer // 4 + 1,
"offset": layer % 4 + 1,
"normalized_gradient": gradients[layer] / gradient_mean,
"attention": mixer_metrics(
diagnostic["depth_weights"][2 * layer]
),
"mlp": mixer_metrics(
diagnostic["depth_weights"][2 * layer + 1]
),
}
)
layer_means = []
for layer in range(1, 33):
selected = [row for row in rows if row["layer"] == layer]
layer_means.append(
{
"layer": layer,
"group": selected[0]["group"],
"offset": selected[0]["offset"],
"normalized_gradient": average(
selected, lambda row: row["normalized_gradient"]
),
"attention_latest": average(
selected, lambda row: row["attention"]["latest"]
),
"mlp_latest": average(
selected, lambda row: row["mlp"]["latest"]
),
"attention_normalized_entropy": average(
selected,
lambda row: row["attention"]["normalized_entropy"],
),
"mlp_normalized_entropy": average(
selected, lambda row: row["mlp"]["normalized_entropy"]
),
"per_seed_normalized_gradient": [
row["normalized_gradient"] for row in selected
],
}
)
correlations = {}
for label, selected in (
("all_layers", rows),
("layers_19_28", [
row for row in rows if 19 <= row["layer"] <= 28
]),
):
gradients = [row["normalized_gradient"] for row in selected]
correlations[label] = {
"points": len(selected),
"attention_latest": pearson(
gradients,
[row["attention"]["latest"] for row in selected],
),
"mlp_latest": pearson(
gradients, [row["mlp"]["latest"] for row in selected]
),
"attention_normalized_entropy": pearson(
gradients,
[
row["attention"]["normalized_entropy"]
for row in selected
],
),
"mlp_normalized_entropy": pearson(
gradients,
[row["mlp"]["normalized_entropy"] for row in selected],
),
"attention_maximum": pearson(
gradients,
[row["attention"]["maximum"] for row in selected],
),
"mlp_maximum": pearson(
gradients, [row["mlp"]["maximum"] for row in selected]
),
}
print(
json.dumps(
{
"schema_version": 1,
"identity": "exploratory Round 05 scoping; not confirmatory Round 06",
"source_canonical_sha256": source_hashes,
"layers_19_28": layer_means[18:28],
"correlations": correlations,
},
ensure_ascii=False,
indent=2,
sort_keys=True,
)
)
if __name__ == "__main__":
main()