h3-blackwell-runtime/research/vortex_exact_attention/tools/capture_canonical_fixtures.py

144 lines
5.6 KiB
Python
Raw Normal View History

2026-08-26 14:44:28 +07:00
"""Capture self-contained canonical Vortex Exact Attention Q/K/V fixtures."""
from __future__ import annotations
import argparse
import hashlib
import json
import platform
import sys
from pathlib import Path
from types import SimpleNamespace
import torch
PROJECT = Path(__file__).resolve().parents[1]
ROOT = PROJECT.parents[1]
sys.path.insert(0, str(ROOT / "tools"))
from profile_attention_path import prepare_qkv, representative_attention_inputs # noqa: E402
EXPECTED_OUTPUT_SHA256 = "4c666c20f5f8f651158a2ced33ccff08f3bada07665c595b99008d171db30574"
EXPECTED_CHECKPOINT_SHA256 = "72fa9269ce551fb63ff42a32d9b46d0c122e84b4b2c511e22fa698287b088f70"
def sha256_file(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(16 * 1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def tensor_sha256(value: torch.Tensor) -> str:
data = value.detach().contiguous().view(torch.uint8).cpu().numpy()
return hashlib.sha256(memoryview(data)).hexdigest()
def save_tensor(output_dir: Path, name: str, value: torch.Tensor) -> dict:
host = value.detach().contiguous().cpu()
tensor_hash = tensor_sha256(host)
path = output_dir / f"{name}.pt"
torch.save({"tensor": host}, path)
loaded = torch.load(path, map_location="cpu", weights_only=True)["tensor"]
loaded_hash = tensor_sha256(loaded)
if loaded_hash != tensor_hash or not torch.equal(loaded, host):
raise RuntimeError(f"fixture reload verification failed for {name}")
return {
"path": str(path),
"size_bytes": path.stat().st_size,
"file_sha256": sha256_file(path),
"tensor_sha256": tensor_hash,
"shape": list(host.shape),
"stride": list(host.stride()),
"dtype": str(host.dtype),
"contiguous": host.is_contiguous(),
"reload_verified": True,
}
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--output-dir", type=Path, required=True)
parser.add_argument("--model-path", type=Path, default=Path("/models/minimax_h3_fl2va_pruned_nvfp4.safetensors"))
parser.add_argument("--sage-mainloop-binary", type=Path, required=True)
parser.add_argument("--sage-fused-binary", type=Path, required=True)
parser.add_argument("--image", required=True)
parser.add_argument("--device", default="cuda")
return parser.parse_args()
def main() -> None:
args = parse_args()
args.output_dir.mkdir(parents=True, exist_ok=True)
checkpoint_hash = sha256_file(args.model_path)
if checkpoint_hash != EXPECTED_CHECKPOINT_SHA256:
raise RuntimeError(f"checkpoint SHA-256 mismatch: {checkpoint_hash}")
workload = SimpleNamespace(
model_path=str(args.model_path), device=args.device, width=1344, height=768,
frames=124, steps=12, sampler_step=1, seed=440420, text_tokens=100,
block_index=24, attention="sage2",
)
block, hidden, rotation, segments, metadata = representative_attention_inputs(workload)
with torch.inference_mode():
q, k, v, _ = prepare_qkv(block, hidden, rotation, None)
from sageattention import sageattn
output = sageattn(q, k, v, tensor_layout="NHD", is_causal=False, smooth_k=False)
torch.cuda.synchronize()
output_hash = tensor_sha256(output)
if output_hash != EXPECTED_OUTPUT_SHA256:
raise RuntimeError(
f"public Sage2 output mismatch: expected {EXPECTED_OUTPUT_SHA256}, got {output_hash}"
)
fixtures = {
"q": save_tensor(args.output_dir, "q_prepared_bf16_nhd", q),
"k": save_tensor(args.output_dir, "k_prepared_bf16_nhd", k),
"v": save_tensor(args.output_dir, "v_bf16_nhd", v),
"output": save_tensor(args.output_dir, "sage2_output_bf16_nhd", output),
}
manifest = {
"schema": "vortex-exact-canonical-fixtures",
"version": 1,
"status": "captured_and_reload_verified",
"reference": {
"sageattention_version": "2.2.0",
"sageattention_commit": "d1a57a546c3d395b1ffcbeecc66d81db76f3b4b5",
"expected_output_sha256": EXPECTED_OUTPUT_SHA256,
"actual_output_sha256": output_hash,
"byte_exact": True,
},
"workload": metadata,
"segments": segments,
"checkpoint": {
"path": str(args.model_path), "size_bytes": args.model_path.stat().st_size,
"sha256": checkpoint_hash,
},
"sage_binaries": {
"mainloop": {"path": str(args.sage_mainloop_binary), "size_bytes": args.sage_mainloop_binary.stat().st_size, "sha256": sha256_file(args.sage_mainloop_binary)},
"fused_preparation": {"path": str(args.sage_fused_binary), "size_bytes": args.sage_fused_binary.stat().st_size, "sha256": sha256_file(args.sage_fused_binary)},
},
"environment": {
"image": args.image,
"gpu": torch.cuda.get_device_name(),
"compute_capability": list(torch.cuda.get_device_capability()),
"torch": torch.__version__,
"cuda": torch.version.cuda,
"driver": torch.cuda.driver_version() if hasattr(torch.cuda, "driver_version") else None,
"python": platform.python_version(),
"argv": sys.argv,
},
"fixtures": fixtures,
}
manifest_path = args.output_dir / "manifest.json"
manifest_path.write_text(json.dumps(manifest, indent=2) + "\n", encoding="utf-8")
print(json.dumps(manifest, indent=2), flush=True)
if __name__ == "__main__":
main()