diff --git a/tools/validate_sage2_vprep.py b/tools/validate_sage2_vprep.py new file mode 100644 index 0000000..e644d63 --- /dev/null +++ b/tools/validate_sage2_vprep.py @@ -0,0 +1,118 @@ +"""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()