212 lines
9.4 KiB
Python
212 lines
9.4 KiB
Python
"""Compare direct and Comfy Qwen3-VL vision outputs on one keyframe."""
|
|
|
|
import argparse
|
|
from pathlib import Path
|
|
import sys
|
|
|
|
import numpy as np
|
|
from PIL import Image
|
|
from safetensors import safe_open
|
|
import torch
|
|
from torch.nn import functional as F
|
|
|
|
from h3_blackwell_runtime.qwen3vl_vision import (
|
|
Qwen3VL32BVision,
|
|
_apply_rope_vision,
|
|
process_image,
|
|
resize_keyframe,
|
|
)
|
|
|
|
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument("--image", required=True)
|
|
parser.add_argument("--checkpoint", default="/text-encoders/qwen3vl_32b_minimax_h3_nvfp4_awq.safetensors")
|
|
parser.add_argument("--width", type=int, default=384)
|
|
parser.add_argument("--height", type=int, default=384)
|
|
parser.add_argument("--comfy-path", default="/opt/ComfyUI")
|
|
parser.add_argument("--save-reference")
|
|
parser.add_argument("--dtype", choices=("float16", "bfloat16", "float32"), default="bfloat16")
|
|
parser.add_argument("--captured-reference")
|
|
args = parser.parse_args()
|
|
|
|
sys.path.insert(0, args.comfy_path)
|
|
import comfy.ops # noqa: E402
|
|
from comfy.ldm.modules.attention import optimized_attention_for_device # noqa: E402
|
|
from comfy.text_encoders.qwen3vl import ( # noqa: E402
|
|
QWEN3VL_VISION,
|
|
QWEN3VL_VISION_COMMON,
|
|
Qwen3VLVisionModel,
|
|
)
|
|
from comfy.text_encoders.qwen_vl import process_qwen2vl_images # noqa: E402
|
|
from comfy.text_encoders.llama import apply_rope # noqa: E402
|
|
|
|
|
|
def report(name, actual, expected):
|
|
actual = actual.detach()
|
|
expected = expected.detach().to(actual.device)
|
|
delta = (actual.float() - expected.float()).abs()
|
|
print({
|
|
"stage": name,
|
|
"shape": tuple(actual.shape),
|
|
"actual_dtype": str(actual.dtype),
|
|
"expected_dtype": str(expected.dtype),
|
|
"mean_delta": float(delta.mean()),
|
|
"max_delta": float(delta.max()),
|
|
}, flush=True)
|
|
|
|
|
|
image = Image.open(args.image).convert("RGB")
|
|
pixels = torch.from_numpy(np.asarray(image).copy()).unsqueeze(0).cuda().float().div(255.0)
|
|
pixels = resize_keyframe(pixels, args.width, args.height)
|
|
|
|
direct_flatten, direct_grid = process_image(pixels)
|
|
reference_flatten, reference_grid = process_qwen2vl_images(
|
|
pixels,
|
|
patch_size=16,
|
|
image_mean=[0.5, 0.5, 0.5],
|
|
image_std=[0.5, 0.5, 0.5],
|
|
)
|
|
report("flatten_patches", direct_flatten, reference_flatten)
|
|
print({"stage": "grid", "direct": direct_grid.tolist(), "reference": reference_grid.tolist()}, flush=True)
|
|
|
|
dtype = {"float16": torch.float16, "bfloat16": torch.bfloat16, "float32": torch.float32}[args.dtype]
|
|
config = {
|
|
**QWEN3VL_VISION_COMMON,
|
|
**QWEN3VL_VISION["qwen3vl_32b"],
|
|
"out_hidden_size": 5120,
|
|
}
|
|
reference = Qwen3VLVisionModel(
|
|
config,
|
|
device="cuda",
|
|
dtype=dtype,
|
|
ops=comfy.ops.disable_weight_init,
|
|
).to("cuda").eval()
|
|
with safe_open(args.checkpoint, framework="pt", device="cuda") as checkpoint:
|
|
print({
|
|
"stage": "checkpoint_dtypes",
|
|
"text_norm": str(checkpoint.get_tensor("model.layers.0.input_layernorm.weight").dtype),
|
|
"vision_norm": str(checkpoint.get_tensor("visual.blocks.0.norm1.weight").dtype),
|
|
"vision_patch": str(checkpoint.get_tensor("visual.patch_embed.proj.weight").dtype),
|
|
}, flush=True)
|
|
visual_state = {
|
|
name.removeprefix("visual."): checkpoint.get_tensor(name).to(dtype)
|
|
for name in checkpoint.keys()
|
|
if name.startswith("visual.")
|
|
}
|
|
reference.load_state_dict(visual_state, strict=True)
|
|
del visual_state
|
|
|
|
direct = Qwen3VL32BVision(args.checkpoint, device="cuda", dtype=dtype).eval()
|
|
with torch.inference_mode():
|
|
direct_x = direct.patch_embed(direct_flatten.cuda().to(dtype))
|
|
direct_patch_embed = direct_x
|
|
reference_x = reference.patch_embed(reference_flatten.cuda().to(dtype))
|
|
report("patch_embed", direct_x, reference_x)
|
|
direct_pos = direct.fast_pos_embed_interpolate(direct_grid).to(direct_x.device)
|
|
reference_pos = reference.fast_pos_embed_interpolate(reference_grid).to(reference_x.device)
|
|
report("position_embed", direct_pos, reference_pos)
|
|
direct_x = direct_x + direct_pos
|
|
direct_vision_input = direct_x
|
|
reference_x = reference_x + reference_pos
|
|
report("vision_input", direct_x, reference_x)
|
|
|
|
direct_rotary = direct.rot_pos_emb(direct_grid.to(direct_x.device)).reshape(direct_x.shape[0], -1)
|
|
reference_rotary = reference.rot_pos_emb(reference_grid).to(reference_x.device).reshape(reference_x.shape[0], -1)
|
|
report("rotary", direct_rotary, reference_rotary)
|
|
|
|
def position_tuple(rotary):
|
|
embedding = torch.cat((rotary, rotary), dim=-1)
|
|
cosine = embedding.cos().unsqueeze(-2)
|
|
sine = embedding.sin().unsqueeze(-2)
|
|
split = sine.shape[-1] // 2
|
|
return cosine, sine[..., :split], -sine[..., split:]
|
|
|
|
direct_position = position_tuple(direct_rotary)
|
|
reference_position = position_tuple(reference_rotary)
|
|
cu_seqlens = F.pad(
|
|
torch.repeat_interleave(direct_grid[:, 1] * direct_grid[:, 2], direct_grid[:, 0]).cumsum(0, dtype=torch.int32),
|
|
(1, 0),
|
|
value=0,
|
|
)
|
|
optimized_attention = optimized_attention_for_device(reference_x.device, mask=False, small_input=True)
|
|
|
|
direct_block0 = direct.blocks[0]
|
|
reference_block0 = reference.blocks[0]
|
|
direct_norm = F.layer_norm(
|
|
direct_x,
|
|
(direct_x.shape[-1],),
|
|
weight=direct_block0.norm1_weight,
|
|
bias=direct_block0.norm1_bias,
|
|
eps=1e-6,
|
|
)
|
|
reference_norm = reference_block0.norm1(reference_x)
|
|
report("block0_norm1", direct_norm, reference_norm)
|
|
direct_qkv = F.linear(direct_norm, direct_block0.attn.qkv_weight, direct_block0.attn.qkv_bias)
|
|
reference_qkv = reference_block0.attn.qkv(reference_norm)
|
|
report("block0_qkv", direct_qkv, reference_qkv)
|
|
direct_q, direct_k, direct_v = direct_qkv.reshape(direct_x.shape[0], 3, 16, 72).permute(1, 0, 2, 3).unbind(0)
|
|
reference_q, reference_k, reference_v = reference_qkv.reshape(reference_x.shape[0], 3, 16, 72).permute(1, 0, 2, 3).unbind(0)
|
|
direct_q, direct_k = _apply_rope_vision(direct_q.float(), direct_k.float(), direct_position)
|
|
direct_q, direct_k = direct_q.to(dtype), direct_k.to(dtype)
|
|
reference_q, reference_k = apply_rope(reference_q, reference_k, reference_position)
|
|
report("block0_rope_q", direct_q, reference_q)
|
|
report("block0_rope_k", direct_k, reference_k)
|
|
direct_attention_heads = F.scaled_dot_product_attention(
|
|
direct_q.transpose(0, 1).unsqueeze(0),
|
|
direct_k.transpose(0, 1).unsqueeze(0),
|
|
direct_v.transpose(0, 1).unsqueeze(0),
|
|
)
|
|
direct_attention = direct_attention_heads.transpose(1, 2).reshape(1, direct_x.shape[0], -1)
|
|
reference_attention = optimized_attention(
|
|
reference_q.transpose(0, 1).unsqueeze(0),
|
|
reference_k.transpose(0, 1).unsqueeze(0),
|
|
reference_v.transpose(0, 1).unsqueeze(0),
|
|
16,
|
|
skip_reshape=True,
|
|
)
|
|
report("block0_attention", direct_attention, reference_attention)
|
|
direct_projected = F.linear(direct_attention[0], direct_block0.attn.proj_weight, direct_block0.attn.proj_bias)
|
|
reference_projected = reference_block0.attn.proj(reference_attention)[0]
|
|
report("block0_projected", direct_projected, reference_projected)
|
|
|
|
direct_deepstack = []
|
|
direct_blocks = []
|
|
reference_deepstack = []
|
|
for index, (direct_block, reference_block) in enumerate(zip(direct.blocks, reference.blocks)):
|
|
direct_x = direct_block(direct_x, cu_seqlens, direct_position)
|
|
direct_blocks.append(direct_x)
|
|
reference_x = reference_block(
|
|
reference_x,
|
|
cu_seqlens,
|
|
reference_position,
|
|
optimized_attention=optimized_attention,
|
|
)
|
|
report(f"block_{index:02d}", direct_x, reference_x)
|
|
if index in direct.deepstack_visual_indexes:
|
|
merger_index = direct.deepstack_visual_indexes.index(index)
|
|
direct_deepstack.append(direct.deepstack_merger_list[merger_index](direct_x))
|
|
reference_deepstack.append(reference.deepstack_merger_list[merger_index](reference_x))
|
|
direct_merged = direct.merger(direct_x)
|
|
reference_merged = reference.merger(reference_x)
|
|
report("merged", direct_merged, reference_merged)
|
|
for index, (actual, expected) in enumerate(zip(direct_deepstack, reference_deepstack)):
|
|
report(f"deepstack_{index}", actual, expected)
|
|
if args.captured_reference:
|
|
captured = torch.load(args.captured_reference, map_location="cuda", weights_only=False)
|
|
report("loaded_comfy_pixel_values", direct_flatten, captured["pixel_values"])
|
|
print({"stage": "loaded_comfy_grid", "direct": direct_grid.tolist(), "expected": captured["grid"].tolist()})
|
|
trace_path = Path(args.captured_reference).with_name(Path(args.captured_reference).name.replace("qwen_vision_", "qwen_vision_trace_"))
|
|
trace = torch.load(trace_path, map_location="cuda", weights_only=False)
|
|
report("loaded_comfy_patch_embed", direct_patch_embed, trace["patch_embed"])
|
|
report("loaded_comfy_position_embed", direct_pos, trace["position_embed"])
|
|
report("loaded_comfy_vision_input", direct_vision_input, trace["vision_input"])
|
|
for index, block_output in enumerate(direct_blocks):
|
|
report(f"loaded_comfy_block_{index:02d}", block_output, trace[f"block_{index:02d}"])
|
|
report("loaded_comfy_merged", direct_merged, captured["merged"])
|
|
for index, (actual, expected) in enumerate(zip(direct_deepstack, captured["deepstack"])):
|
|
report(f"loaded_comfy_deepstack_{index}", actual, expected)
|
|
if args.save_reference:
|
|
torch.save({
|
|
"merged": reference_merged.detach().cpu(),
|
|
"deepstack": [value.detach().cpu() for value in reference_deepstack],
|
|
}, args.save_reference)
|