"""Validate Vortex direct Sage2 V preparation against SageAttention 2.2.0.""" from __future__ import annotations import argparse import json from pathlib import Path import torch from h3_blackwell_runtime.sage2_entry import prepare_v from profile_attention_path import summarize def difference(actual: torch.Tensor, expected: torch.Tensor) -> dict: delta = actual.float() - expected.float() return { "equal": torch.equal(actual, expected), "different_elements": int(torch.count_nonzero(actual != expected).item()), "max_abs": delta.abs().max().item() if delta.numel() else 0.0, "mean_abs": delta.abs().mean().item() if delta.numel() else 0.0, } def measure(fn, *, warmup: int, iterations: int) -> dict: for _ in range(warmup): fn() torch.cuda.synchronize() samples = [] for _ in range(iterations): started = torch.cuda.Event(enable_timing=True) finished = torch.cuda.Event(enable_timing=True) started.record() fn() finished.record() finished.synchronize() samples.append(started.elapsed_time(finished) / 1000.0) return summarize(samples) def baseline(v: torch.Tensor): import sageattention.core as sage_core return sage_core.per_channel_fp8( v, tensor_layout="NHD", scale_max=2.25, smooth_v=False, ) def run_case(sequence: int, heads: int, seed: int, warmup: int, iterations: int) -> dict: generator = torch.Generator(device="cuda").manual_seed(seed) storage = torch.randn( (sequence, heads * 128 * 3), generator=generator, device="cuda", dtype=torch.bfloat16, ) v = storage[:, heads * 128 * 2 :].view(1, sequence, heads, 128) reference_fp8, reference_scale, _ = baseline(v) candidate_fp8, candidate_scale = prepare_v(v) torch.cuda.synchronize() result = { "sequence": sequence, "heads": heads, "stride": list(v.stride()), "fp8": difference(candidate_fp8, reference_fp8), "scale": difference(candidate_scale, reference_scale), } if sequence >= 1024: result["baseline_timing"] = measure( lambda: baseline(v), warmup=warmup, iterations=iterations, ) result["candidate_timing"] = measure( lambda: prepare_v(v), warmup=warmup, iterations=iterations, ) return result def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--output", type=Path, required=True) parser.add_argument( "--lengths", nargs="+", type=int, default=(1, 31, 32, 33, 63, 64, 65, 127, 128, 129, 37760, 37761, 37810), ) parser.add_argument("--heads", type=int, default=56) parser.add_argument("--seed", type=int, default=440420) parser.add_argument("--warmup", type=int, default=3) parser.add_argument("--iterations", type=int, default=10) return parser.parse_args() def main() -> None: args = parse_args() cases = [ run_case(length, args.heads, args.seed + index, args.warmup, args.iterations) for index, length in enumerate(args.lengths) ] report = { "status": "pass" if all(case["fp8"]["equal"] and case["scale"]["equal"] for case in cases) else "fail", "contract": { "tensor_layout": "NHD", "head_dim": 128, "scale_max": 2.25, "output": "Sage2 padded/permuted E4M3 V and FP32 per-channel scale", }, "cases": cases, } args.output.parent.mkdir(parents=True, exist_ok=True) args.output.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8") print(json.dumps(report, indent=2), flush=True) if report["status"] != "pass": raise RuntimeError("Vortex Sage2 V preparation did not match SageAttention 2.2.0") if __name__ == "__main__": main()