124 lines
4.1 KiB
Diff
124 lines
4.1 KiB
Diff
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()
|