Files
AHA-L2A-Qwen3-1.7B-repro/recipe/scripts/summarize_qwen1p7b_router_granularity_20260714.py
ModelHub XC c13d63b439 初始化项目,由ModelHub XC社区提供模型
Model: jiamingshan/AHA-L2A-Qwen3-1.7B-repro
Source: Original Platform
2026-07-21 11:06:13 +08:00

404 lines
17 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/usr/bin/env python3
"""Audit and summarize the matched shared-gate vs KV-head-gate sweep."""
from __future__ import annotations
import argparse
import csv
import glob
import hashlib
import json
import subprocess
from collections import defaultdict
from pathlib import Path
import matplotlib.pyplot as plt
SUITES = ("RULER-local-full13", "BabiLong", "HELMET-ICL", "MRCR")
THRESHOLDS = (0.45, 0.50, 0.525, 0.55, 0.575, 0.60, 0.625, 0.65)
ARMS = ("token", "token_kv_head")
SLUGS = {
"RULER-local-full13": "ruler",
"BabiLong": "babilong",
"HELMET-ICL": "helmet",
"MRCR": "mrcr",
}
SCORE_DEFINITIONS = {
"RULER-local-full13": "unweighted macro of lm-eval 8192 string-match scores over the local 13-config set",
"BabiLong": "unweighted macro exact-match accuracy over qa1-qa5",
"HELMET-ICL": "unweighted macro label exact-match over five ICL configurations",
"MRCR": "unweighted macro of prefix-check plus SequenceMatcher scores over 2/4/8-needle configurations",
}
def threshold_slug(value: float) -> str:
return {
0.45: "045",
0.50: "050",
0.525: "0525",
0.55: "055",
0.575: "0575",
0.60: "060",
0.625: "0625",
0.65: "065",
}[value]
def one(pattern: str) -> Path:
paths = [Path(path) for path in glob.glob(pattern, recursive=True)]
if len(paths) != 1:
raise RuntimeError(f"expected one path for {pattern}, found {paths}")
return paths[0]
def read_jsonl(path: Path) -> list[dict]:
return [json.loads(line) for line in path.read_text().splitlines() if line.strip()]
def collect_scores(root: Path) -> dict[tuple[str, str], dict]:
rows: dict[tuple[str, str], dict] = {}
for split in ("a", "b"):
result = one(str(root / f"ruler8k_splits/{split}/lm_eval/**/results_*.json"))
payload = json.loads(result.read_text())
for config, metrics in payload["results"].items():
if config in payload["n-samples"]:
rows[("RULER-local-full13", config)] = {
"score": float(metrics["8192,none"]),
"n": int(payload["n-samples"][config]["effective"]),
}
result = one(str(root / "babilong8k_qa1_qa5_n50/**/results_*.json"))
payload = json.loads(result.read_text())
for config, metrics in payload["results"].items():
if config in payload["n-samples"]:
rows[("BabiLong", config)] = {
"score": float(metrics["acc,none"]),
"n": int(payload["n-samples"][config]["effective"]),
}
for suite, filename in (
("HELMET-ICL", "helmet_icl8k_n50.jsonl"),
("MRCR", "mrcr_8k_2_4_8needle_n10.jsonl"),
):
grouped: dict[str, list[dict]] = defaultdict(list)
for row in read_jsonl(root / filename):
grouped[row["config"]].append(row)
for config, values in grouped.items():
rows[(suite, config)] = {
"score": sum(float(row["score"]) for row in values) / len(values),
"n": len(values),
}
return rows
def sample_signatures(root: Path, score_rows: dict[tuple[str, str], dict]) -> dict[str, str]:
signatures: dict[str, str] = {}
for suite, rel_dirs in (
("RULER-local-full13", ("ruler8k_splits/a/lm_eval", "ruler8k_splits/b/lm_eval")),
("BabiLong", ("babilong8k_qa1_qa5_n50",)),
):
configs = [config for row_suite, config in score_rows if row_suite == suite]
for config in configs:
matches = []
for rel_dir in rel_dirs:
matches.extend(root.glob(f"{rel_dir}/**/samples_{config}_*.jsonl"))
if len(matches) != 1:
raise RuntimeError(f"expected one sample log for {suite}/{config}: {matches}")
values = [
(
row.get("doc_id"),
row.get("doc_hash"),
row.get("prompt_hash"),
row.get("target_hash"),
)
for row in read_jsonl(matches[0])
]
signatures[f"{suite}/{config}"] = hashlib.sha256(
json.dumps(values, sort_keys=True).encode()
).hexdigest()
for suite, filename in (
("HELMET-ICL", "helmet_icl8k_n50.jsonl"),
("MRCR", "mrcr_8k_2_4_8needle_n10.jsonl"),
):
grouped: dict[str, list[tuple]] = defaultdict(list)
for row in read_jsonl(root / filename):
grouped[row["config"]].append(
(row["row_id"], row.get("prompt_sha256"), row.get("target"))
)
for config, values in grouped.items():
signatures[f"{suite}/{config}"] = hashlib.sha256(
json.dumps(values, sort_keys=True).encode()
).hexdigest()
return signatures
def read_sparsity(root: Path, suite: str) -> dict:
if suite == "RULER-local-full13":
parts = [
json.loads((root / f"sparsity_ruler_{split}.json").read_text())
for split in ("a", "b")
]
sparse = sum(int(part["sparse_decisions"]) for part in parts)
total = sum(int(part["total_decisions"]) for part in parts)
granularities = {part.get("router_granularity") for part in parts}
native = sum(int(part.get("native_router_decisions", 0)) for part in parts)
phases: dict[str, dict[str, int]] = defaultdict(lambda: {"sparse": 0, "total": 0})
for part in parts:
for phase, values in part.get("by_phase", {}).items():
phases[phase]["sparse"] += int(values["sparse_decisions"])
phases[phase]["total"] += int(values["total_decisions"])
else:
part = json.loads((root / f"sparsity_{SLUGS[suite]}.json").read_text())
sparse = int(part["sparse_decisions"])
total = int(part["total_decisions"])
granularities = {part.get("router_granularity")}
native = int(part.get("native_router_decisions", 0))
phases = {
phase: {
"sparse": int(values["sparse_decisions"]),
"total": int(values["total_decisions"]),
}
for phase, values in part.get("by_phase", {}).items()
}
return {
"sparse_decisions": sparse,
"total_decisions": total,
"sparsity": sparse / total,
"full_attention_usage": 1.0 - sparse / total,
"router_granularities": sorted(value for value in granularities if value),
"native_router_decisions": native or None,
"by_phase": phases,
}
def sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for block in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(block)
return digest.hexdigest()
def command_output(command: list[str], cwd: Path) -> str:
try:
return subprocess.run(
command, cwd=cwd, check=True, text=True, capture_output=True
).stdout.strip()
except (OSError, subprocess.CalledProcessError):
return "unavailable in downloaded HF snapshot"
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument(
"--repo", type=Path, default=Path("/data/sjm/AHA/AHA-Qwen3")
)
parser.add_argument(
"--eval-root",
type=Path,
default=Path("/data/sjm/AHA/AHA-Qwen3/experiments/qwen3_1p7b_router_granularity_20260714/eval"),
)
parser.add_argument(
"--baseline-root",
type=Path,
default=Path("/data/sjm/AHA/AHA-Qwen3/experiments/qwen3_1p7b_fourbench_8k_broad_20260713/vanilla"),
)
parser.add_argument("--output-dir", type=Path)
parser.add_argument("--input-dir", type=Path)
args = parser.parse_args()
output = args.output_dir or args.eval_root.parent / "summary"
output.mkdir(parents=True, exist_ok=True)
baseline = collect_scores(args.baseline_root)
baseline_signatures = sample_signatures(args.baseline_root, baseline)
curve_rows = []
per_config = []
points = []
alignment = {}
done_files = ("ruler_a.DONE", "ruler_b.DONE", "babilong.DONE", "helmet.DONE", "mrcr.DONE")
for arm in ARMS:
for threshold in THRESHOLDS:
method = f"{arm}_t{threshold_slug(threshold)}"
root = args.eval_root / method
missing = [name for name in done_files if not (root / name).exists()]
if missing:
raise RuntimeError(f"incomplete {method}: missing {missing}")
scores = collect_scores(root)
if set(scores) != set(baseline):
raise RuntimeError(f"configuration mismatch for {method}")
if any(scores[key]["n"] != baseline[key]["n"] for key in baseline):
raise RuntimeError(f"sample-count mismatch for {method}")
signatures = sample_signatures(root, scores)
mismatches = sorted(
key for key in baseline_signatures if signatures.get(key) != baseline_signatures[key]
)
alignment[method] = {"aligned": not mismatches, "mismatches": mismatches}
if mismatches:
raise RuntimeError(f"prompt/target hash mismatch for {method}: {mismatches}")
suite_rows = []
sparse_sum = total_sum = 0
for suite in SUITES:
keys = sorted(key for key in baseline if key[0] == suite)
baseline_score = sum(baseline[key]["score"] for key in keys) / len(keys)
score = sum(scores[key]["score"] for key in keys) / len(keys)
sparsity = read_sparsity(root, suite)
expected_granularity = arm
if sparsity["router_granularities"] != [expected_granularity]:
raise RuntimeError(
f"{method}/{suite} recorded {sparsity['router_granularities']}, "
f"expected {[expected_granularity]}"
)
sparse_sum += sparsity["sparse_decisions"]
total_sum += sparsity["total_decisions"]
row = {
"arm": arm,
"router_granularity": arm,
"threshold": threshold,
"benchmark": suite,
"config_count": len(keys),
"samples_per_config": ",".join(
str(value) for value in sorted({baseline[key]["n"] for key in keys})
),
"total_samples": sum(baseline[key]["n"] for key in keys),
"tuned_vanilla_score": baseline_score,
"quality_floor_95pct": 0.95 * baseline_score,
"score": score,
"retention": score / baseline_score if baseline_score else None,
"quality_pass": score >= 0.95 * baseline_score,
**{key: sparsity[key] for key in (
"sparse_decisions", "total_decisions", "sparsity",
"full_attention_usage", "native_router_decisions", "by_phase"
)},
}
curve_rows.append(row)
suite_rows.append(row)
for key in keys:
per_config.append({
"arm": arm,
"threshold": threshold,
"benchmark": suite,
"config": key[1],
"n": baseline[key]["n"],
"tuned_vanilla_score": baseline[key]["score"],
"score": scores[key]["score"],
})
points.append({
"arm": arm,
"threshold": threshold,
"all_suite_quality_pass": all(row["quality_pass"] for row in suite_rows),
"decision_weighted_sparsity": sparse_sum / total_sum,
"full_attention_usage": 1.0 - sparse_sum / total_sum,
"sparse_decisions": sparse_sum,
"total_decisions": total_sum,
"suites": suite_rows,
})
headline = {}
for arm in ARMS:
valid = [
point for point in points
if point["arm"] == arm and point["all_suite_quality_pass"]
]
if valid:
best = max(valid, key=lambda point: point["decision_weighted_sparsity"])
headline[arm] = {"found": True, **best}
else:
headline[arm] = {
"found": False,
"statement": "No measured threshold preserved at least 95% of tuned vanilla on every suite.",
}
input_dir = args.input_dir or args.baseline_root.parent / "inputs"
frozen_inputs = [
input_dir / "helmet_icl_8k_n50_per_config.jsonl",
input_dir / "mrcr_8k_2_4_8needle_n10_per_config.jsonl",
]
payload = {
"protocol": {
"model": "Qwen3-1.7B tuned vanilla",
"context_length": 8192,
"arms": {
"token": "L2A-style shared-gate: one native gate per token/layer",
"token_kv_head": "AHA: one native gate per token/KV-head/layer",
},
"local_attention": {"sink_tokens": 64, "recent_tokens": 256},
"inference": "strict AHA routing in both prefill and decode; no force-full heads or full-decode fallback",
"thresholds": list(THRESHOLDS),
"quality_rule": "every suite macro score >= 95% of aligned tuned-vanilla macro score",
"headline_rule": "highest decision-weighted measured sparsity among thresholds passing every suite",
"sparsity_definition": "hard local routes / token x KV-head x layer effective decisions",
"suite_scope": "RULER local full13 means this repository's fixed 13-config set, not complete upstream RULER",
},
"score_definitions": SCORE_DEFINITIONS,
"headline": headline,
"points": points,
"curves": curve_rows,
"per_config": per_config,
"alignment": {"all_aligned": all(value["aligned"] for value in alignment.values()), "methods": alignment},
"appendix_note": (
"Upstream L2A results are not merged into this table because architecture, "
"training data, objective, and evaluation protocols differ."
),
"reproducibility": {
"git_commit": command_output(["git", "rev-parse", "HEAD"], args.repo),
"git_status_short": command_output(["git", "status", "--short"], args.repo),
"container": "dmtd-repro",
"container_image": command_output(
["docker", "inspect", "-f", "{{.Config.Image}}", "dmtd-repro"], args.repo
),
"baseline_root": str(args.baseline_root),
"eval_root": str(args.eval_root),
"frozen_input_sha256": {
str(path): sha256(path) for path in frozen_inputs if path.exists()
},
},
}
(output / "router_granularity_results.json").write_text(
json.dumps(payload, indent=2) + "\n"
)
flat_curve_rows = [
{key: value for key, value in row.items() if key != "by_phase"}
for row in curve_rows
]
with (output / "quality_sparsity_curves.csv").open("w", newline="") as handle:
writer = csv.DictWriter(handle, fieldnames=list(flat_curve_rows[0]))
writer.writeheader()
writer.writerows(flat_curve_rows)
with (output / "per_config_scores.csv").open("w", newline="") as handle:
writer = csv.DictWriter(handle, fieldnames=list(per_config[0]))
writer.writeheader()
writer.writerows(per_config)
figure, axes = plt.subplots(2, 2, figsize=(12, 9), constrained_layout=True)
colors = {"token": "#d62728", "token_kv_head": "#1f77b4"}
labels = {"token": "L2A-style shared-gate", "token_kv_head": "AHA KV-head gate"}
for axis, suite in zip(axes.flat, SUITES):
for arm in ARMS:
values = sorted(
(row for row in curve_rows if row["arm"] == arm and row["benchmark"] == suite),
key=lambda row: row["sparsity"],
)
axis.plot(
[100 * row["sparsity"] for row in values],
[100 * row["retention"] for row in values],
marker="o", color=colors[arm], label=labels[arm],
)
for row in values:
axis.annotate(f"{row['threshold']:.3g}", (100 * row["sparsity"], 100 * row["retention"]), fontsize=7)
axis.axhline(95, color="black", linestyle="--", linewidth=1)
axis.set_title(suite)
axis.set_xlabel("Measured effective sparsity (%)")
axis.set_ylabel("Retention vs tuned vanilla (%)")
axis.grid(alpha=0.25)
axes.flat[0].legend()
figure.suptitle("Qwen3-1.7B matched router-granularity qualitysparsity curves")
figure.savefig(output / "quality_sparsity_curves.png", dpi=180)
plt.close(figure)
print(json.dumps(headline, indent=2))
if __name__ == "__main__":
main()