35 lines
1.7 KiB
Python
35 lines
1.7 KiB
Python
"""Compare direct Qwen layer-0 projections with a Comfy capture."""
|
|
|
|
import argparse
|
|
|
|
import torch
|
|
import torch.nn.functional as functional
|
|
|
|
from h3_blackwell_runtime.qwen3vl_text import Qwen3VL32BTextEncoder
|
|
|
|
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument("--capture-dir", required=True)
|
|
parser.add_argument("--checkpoint", default="/text-encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors")
|
|
args = parser.parse_args()
|
|
|
|
encoder = Qwen3VL32BTextEncoder(args.checkpoint).eval()
|
|
layer = encoder.layers[0]
|
|
embeds = torch.load(f"{args.capture_dir}/qwen_input_embeds.pt", map_location="cuda", weights_only=False)
|
|
qkv_expected = torch.load(f"{args.capture_dir}/qwen0_qkv.pt", map_location="cuda", weights_only=False)
|
|
mlp_expected = torch.load(f"{args.capture_dir}/qwen0_mlp_projections.pt", map_location="cuda", weights_only=False)
|
|
|
|
with torch.inference_mode():
|
|
norm1 = layer.input_layernorm(embeds.to(encoder.dtype))
|
|
qkv_actual = {"q": layer.q_proj(norm1), "k": layer.k_proj(norm1), "v": layer.v_proj(norm1)}
|
|
post_attention = torch.load(f"{args.capture_dir}/qwen0_post_attention.pt", map_location="cuda", weights_only=False)
|
|
norm2 = layer.post_attention_layernorm(post_attention.to(encoder.dtype))
|
|
gate = layer.gate_proj(norm2)
|
|
up = layer.up_proj(norm2)
|
|
activated = functional.silu(gate) * up
|
|
mlp_actual = {"gate": gate, "up": up, "activated": activated, "down": layer.down_proj(activated)}
|
|
|
|
for group, actual, expected in (("qkv", qkv_actual, qkv_expected), ("mlp", mlp_actual, mlp_expected)):
|
|
for name, value in actual.items():
|
|
delta = (value.float() - expected[name].float()).abs()
|
|
print(f"{group}.{name} shape={tuple(value.shape)} max_abs={delta.max().item():.6g} mean_abs={delta.mean().item():.6g}")
|