35 lines
1.5 KiB
Python
35 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}")
|