h3-blackwell-runtime/tools/compare_block0_qkv.py
2026-08-12 21:11:02 +07:00

46 lines
3 KiB
Python

"""Compare prepared block-0 QKV tensors and Sage3 output with ComfyUI."""
import argparse
import torch
from h3_blackwell_runtime.attention import apply_split_half_rope, rms_norm
from h3_blackwell_runtime.block import modulate_segments
from h3_blackwell_runtime.checkpoint import H3Checkpoint
from h3_blackwell_runtime.denoiser import H3PackedDenoiser
from h3_blackwell_runtime.rope import h3_rope_rotation
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()
capture_dir = args.capture_dir
inputs = torch.load(f"{capture_dir}/input.pt", map_location="cuda", weights_only=False)
expected_qkv = torch.load(f"{capture_dir}/block0_qkv_prepared.pt", map_location="cuda", weights_only=False)
expected_raw = torch.load(f"{capture_dir}/block0_qkv_raw.pt", map_location="cuda", weights_only=False)
expected_norm1 = torch.load(f"{capture_dir}/block0_norm1.pt", map_location="cuda", weights_only=False)
expected_attention = torch.load(f"{capture_dir}/block0_attention.pt", map_location="cuda", weights_only=False)
checkpoint = H3Checkpoint(args.model)
model = H3PackedDenoiser.from_checkpoint(checkpoint).eval()
block = model.backbone.blocks[0]
shift_msa, scale_msa, *_ = model.backbone.adaln[0](inputs["timesteps"])
hidden = modulate_segments(rms_norm(inputs["hidden"], block.norm1_weight, block.norm_eps), shift_msa, scale_msa, inputs["segments"])
sequence = hidden.shape[0]
inner = block.attention.heads * block.attention.head_dim
rotation = h3_rope_rotation(inputs["position_ids"], model.backbone.inv_freq, hidden.dtype)
with torch.inference_mode():
raw_q, raw_k, raw_v = block.attention.qkv_proj(hidden).split(inner, dim=-1)
comfy_raw_q, comfy_raw_k, comfy_raw_v = block.attention.qkv_proj(expected_norm1).split(inner, dim=-1)
q, k, v = raw_q, raw_k, raw_v
q = apply_split_half_rope(rms_norm(q.view(1, sequence, 56, 128), block.attention.q_norm_weight, block.attention.eps), rotation).transpose(1, 2).contiguous()
k = apply_split_half_rope(rms_norm(k.view(1, sequence, 56, 128), block.attention.k_norm_weight, block.attention.eps), rotation).transpose(1, 2).contiguous()
v = v.view(1, sequence, 56, 128).transpose(1, 2).contiguous()
from sageattn3 import sageattn3_blackwell
attention = block.attention.out_proj(sageattn3_blackwell(q, k, v, is_causal=False).transpose(1, 2).reshape(sequence, inner).contiguous())
for name, actual, expected in (("raw_q", raw_q, expected_raw["q"]), ("raw_k", raw_k, expected_raw["k"]), ("raw_v", raw_v, expected_raw["v"]), ("comfy_norm_raw_q", comfy_raw_q, expected_raw["q"]), ("comfy_norm_raw_k", comfy_raw_k, expected_raw["k"]), ("comfy_norm_raw_v", comfy_raw_v, expected_raw["v"]), ("q", q, expected_qkv["q"]), ("k", k, expected_qkv["k"]), ("v", v, expected_qkv["v"]), ("attention", attention, expected_attention)):
delta = (actual.float() - expected.float()).abs()
print(f"{name} max_abs={delta.max().item():.6g} mean_abs={delta.mean().item():.6g}")