39 lines
2.8 KiB
Python
39 lines
2.8 KiB
Python
"""Compare prepared block-0 QKV tensors and Sage3 output with ComfyUI."""
|
|
|
|
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
|
|
|
|
|
|
capture_dir = "/artifacts/capture"
|
|
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("/models/minimax_h3_ref2va_pruned_nvfp4.safetensors")
|
|
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}")
|