h3-blackwell-runtime/tools/compare_attention_backends.py
2026-08-12 14:12:42 +07:00

34 lines
1.5 KiB
Python

"""Compare direct attention kernels against captured Comfy block-0 output."""
import time
import torch
import torch.nn.functional as functional
from h3_blackwell_runtime.checkpoint import H3Checkpoint
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
payload = torch.load("/artifacts/capture/block0_qkv_prepared.pt", map_location="cuda", weights_only=False)
expected = torch.load("/artifacts/capture/block0_attention.pt", map_location="cuda", weights_only=False)
q, k, v = payload["q"], payload["k"], payload["v"]
model = H3PackedDenoiser.from_checkpoint(H3Checkpoint("/models/minimax_h3_ref2va_pruned_nvfp4.safetensors")).eval()
out_proj = model.backbone.blocks[0].attention.out_proj
for name in ("sdpa", "sage2", "sage3"):
torch.cuda.synchronize()
start = time.perf_counter()
if name == "sdpa":
output = functional.scaled_dot_product_attention(q, k, v, is_causal=False)
elif name == "sage2":
from sageattention import sageattn
output = sageattn(q, k, v, is_causal=False, tensor_layout="HND", smooth_k=False)
else:
from sageattn3 import sageattn3_blackwell
output = sageattn3_blackwell(q, k, v, is_causal=False)
torch.cuda.synchronize()
output = out_proj(output.transpose(1, 2).reshape(q.shape[0], q.shape[2], -1).reshape(q.shape[2], -1).contiguous())
delta = (output.float() - expected.float()).abs()
print(f"{name} elapsed_s={time.perf_counter() - start:.3f} max_abs={delta.max().item():.6g} mean_abs={delta.mean().item():.6g}")