2026-08-12 14:12:42 +07:00
|
|
|
"""Compare direct attention kernels against captured Comfy block-0 output."""
|
|
|
|
|
|
2026-08-12 21:11:02 +07:00
|
|
|
import argparse
|
2026-08-12 14:12:42 +07:00
|
|
|
import time
|
|
|
|
|
|
|
|
|
|
import torch
|
|
|
|
|
|
2026-08-14 20:28:56 +07:00
|
|
|
from h3_blackwell_runtime.attention import AVAILABLE_BACKENDS, run_attention
|
2026-08-12 14:12:42 +07:00
|
|
|
from h3_blackwell_runtime.checkpoint import H3Checkpoint
|
|
|
|
|
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
|
|
|
|
|
|
|
|
|
|
|
2026-08-12 21:11:02 +07:00
|
|
|
parser = argparse.ArgumentParser()
|
|
|
|
|
parser.add_argument("--capture-dir", default="/artifacts/capture")
|
|
|
|
|
parser.add_argument("--model", default="/models/minimax_h3_ref2va_pruned_nvfp4.safetensors")
|
|
|
|
|
args = parser.parse_args()
|
|
|
|
|
|
|
|
|
|
payload = torch.load(f"{args.capture_dir}/block0_qkv_prepared.pt", map_location="cuda", weights_only=False)
|
|
|
|
|
expected = torch.load(f"{args.capture_dir}/block0_attention.pt", map_location="cuda", weights_only=False)
|
2026-08-12 14:12:42 +07:00
|
|
|
q, k, v = payload["q"], payload["k"], payload["v"]
|
2026-08-12 21:11:02 +07:00
|
|
|
model = H3PackedDenoiser.from_checkpoint(H3Checkpoint(args.model)).eval()
|
2026-08-12 14:12:42 +07:00
|
|
|
out_proj = model.backbone.blocks[0].attention.out_proj
|
|
|
|
|
|
2026-08-14 20:28:56 +07:00
|
|
|
for name in AVAILABLE_BACKENDS:
|
2026-08-12 14:12:42 +07:00
|
|
|
torch.cuda.synchronize()
|
|
|
|
|
start = time.perf_counter()
|
2026-08-14 20:28:56 +07:00
|
|
|
try:
|
|
|
|
|
output = run_attention(q, k, v, backend=name, 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}")
|
|
|
|
|
except Exception as exc:
|
|
|
|
|
print(f"{name} error={type(exc).__name__}: {exc}")
|