"""Compare direct attention kernels against captured Comfy block-0 output.""" import argparse import time import torch import torch.nn.functional as functional from h3_blackwell_runtime.checkpoint import H3Checkpoint from h3_blackwell_runtime.denoiser import H3PackedDenoiser 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) q, k, v = payload["q"], payload["k"], payload["v"] model = H3PackedDenoiser.from_checkpoint(H3Checkpoint(args.model)).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}")