strata-headquotient-q25 / verify_export.py
nur-dev's picture
Add files using upload-large-folder tool
e69b72a verified
Raw
History Blame Contribute Delete
4.33 kB
#!/usr/bin/env python3
"""Verify the public Q25 checkpoint, static plan, and evaluation bundle."""
from __future__ import annotations
from collections import Counter
import hashlib
import json
from pathlib import Path
import numpy as np
import torch
ROOT = Path(__file__).resolve().parent
EXPECTED_DENSE = "10e8559713ef1d951c604605f8f3666a027a25a341363d0c17006f628cc38c1f"
EXPECTED_Q25 = "51120c7ecca5234a4e7cee424c199582e444eb04e434a2d7f08263e3ccc74a90"
def sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(16 << 20), b""):
digest.update(chunk)
return digest.hexdigest()
def verify_manifest() -> int:
rows = json.loads((ROOT / "inventory.json").read_text(encoding="utf-8"))
for relative, expected in rows.items():
path = ROOT / relative
if not path.is_file():
raise AssertionError(f"missing release file: {relative}")
if path.stat().st_size != expected["bytes"] or sha256(path) != expected["sha256"]:
raise AssertionError(f"release file changed: {relative}")
return len(rows)
def main() -> None:
files = verify_manifest()
dense_path = ROOT / "checkpoints/dense_model.pt"
q25_path = ROOT / "checkpoints/q25_export.pt"
if sha256(dense_path) != EXPECTED_DENSE or sha256(q25_path) != EXPECTED_Q25:
raise AssertionError("checkpoint digest differs from the article manifest")
plan = json.loads((ROOT / "selection/Q25/plan.json").read_text(encoding="utf-8"))
assignments = {
(int(row["layer"]), int(row["head"])): row["mode"] for row in plan["groups"]
}
if len(assignments) != 96:
raise AssertionError("Q25 plan does not contain 96 unique heads")
plan_modes = Counter(assignments.values())
if plan_modes != Counter({"LOCAL": 81, "LOCAL_GRAPH": 15}):
raise AssertionError(f"unexpected plan modes: {plan_modes}")
state = torch.load(q25_path, map_location="cpu", weights_only=True)
state_modes: dict[tuple[int, int], int] = {}
for layer in range(24):
key = f"base_model.blocks.{layer}.attention.modes"
values = state[key].tolist()
if len(values) != 16:
raise AssertionError(f"layer {layer} does not contain 16 head modes")
for head, mode in enumerate(values):
state_modes[(layer, head)] = int(mode)
counts = Counter(state_modes.values())
if counts != Counter({0: 288, 2: 81, 1: 15}):
raise AssertionError(f"checkpoint mode counts differ: {counts}")
for identity, mode in assignments.items():
expected = 1 if mode == "LOCAL_GRAPH" else 2
if state_modes[identity] != expected:
raise AssertionError(f"checkpoint/plan mismatch at {identity}")
exported_layers = {
int(key.split(".")[2])
for key in state
if key.endswith("attention.global_head_indices")
}
if exported_layers != {2, 4, 5, 17, 18, 19, 20, 21, 22, 23}:
raise AssertionError(f"unexpected physically packed layers: {exported_layers}")
with np.load(ROOT / "evaluation/q25_470_documents.npz") as bundle:
if bundle["tokens"].shape != (470, 8192) or bundle["tokens"].dtype != np.uint16:
raise AssertionError("evaluation token bundle has the wrong shape or dtype")
languages = Counter(str(item) for item in bundle["languages"])
if languages != Counter({"zh": 300, "en": 50, "de": 50, "es": 50, "ar": 20}):
raise AssertionError(f"evaluation language support differs: {languages}")
if len(set(str(item) for item in bundle["document_sha256"])) != 470:
raise AssertionError("evaluation documents are duplicated")
parameter_elements = sum(value.numel() for value in state.values())
print(json.dumps({
"verified_files": files,
"dense_checkpoint_sha256": EXPECTED_DENSE,
"q25_checkpoint_sha256": EXPECTED_Q25,
"checkpoint_state_elements_including_buffers": parameter_elements,
"global_heads": counts[0],
"local_heads": counts[2],
"local_graph_heads": counts[1],
"physically_packed_layers": sorted(exported_layers),
"evaluation_documents": 470,
"status": "verified",
}, indent=2))
if __name__ == "__main__":
main()