"""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)