"""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()