h3-blackwell-runtime/tools/compare_qwen0_loaded_modules.py
2026-08-12 22:43:43 +07:00

30 lines
1.4 KiB
Python

"""Replay direct layer-0 projections from actual loaded-Comfy module inputs."""
import argparse
import torch
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]
qkv = torch.load(f"{args.capture_dir}/qwen0_loaded_qkv.pt", map_location="cuda", weights_only=False)
mlp = torch.load(f"{args.capture_dir}/qwen0_loaded_mlp.pt", map_location="cuda", weights_only=False)
with torch.inference_mode():
qkv_actual = {"q": layer.q_proj(qkv["input"].to(encoder.dtype)), "k": layer.k_proj(qkv["input"].to(encoder.dtype)), "v": layer.v_proj(qkv["input"].to(encoder.dtype))}
gate = layer.gate_proj(mlp["input"].to(encoder.dtype))
up = layer.up_proj(mlp["input"].to(encoder.dtype))
mlp_actual = {"gate": gate, "up": up, "down": layer.down_proj(mlp["output"]["activated"].to(encoder.dtype))}
for group, actual, captured in (("qkv", qkv_actual, qkv), ("mlp", mlp_actual, mlp)):
print(f"{group}.loaded_metadata={captured['metadata']}")
for name, value in actual.items():
delta = (value.float() - captured["output"][name].float()).abs()
print(f"{group}.{name} max_abs={delta.max().item():.6g} mean_abs={delta.mean().item():.6g}")