h3-blackwell-runtime/tools/validate_sage2_vprep.py
2026-08-25 20:30:22 +07:00

118 lines
3.8 KiB
Python

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