From 571d3541e899d997b0dc018a64530a27f0816cf7 Mon Sep 17 00:00:00 2001 From: Daniel Maddern Date: Wed, 19 Aug 2026 22:34:16 +0700 Subject: [PATCH] Use explicit C (not x.shape[-1]) in _merge_tokens --- src/h3_blackwell_runtime/qwen3vl_vision.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/src/h3_blackwell_runtime/qwen3vl_vision.py b/src/h3_blackwell_runtime/qwen3vl_vision.py index 1391f01..0887a85 100644 --- a/src/h3_blackwell_runtime/qwen3vl_vision.py +++ b/src/h3_blackwell_runtime/qwen3vl_vision.py @@ -678,11 +678,10 @@ class Qwen3VL32BVision(nn.Module): t = int(first[0].item()) h = int(first[1].item()) w = int(first[2].item()) + C = int(x.shape[-1]) merge = 2 - print(f"[merge_tokens] x.shape={tuple(x.shape)}, t={t}, h={h}, w={w}, " - f"target=({t}, {h//merge}, {merge}, {w//merge}, {merge}, {x.shape[-1]}), " - f"prod={t*(h//merge)*merge*(w//merge)*merge*x.shape[-1]}", flush=True) - return x.view((t, h // merge, merge, w // merge, merge, x.shape[-1])).permute(0, 1, 3, 2, 4, 5).reshape(-1, -1) + target = (t, h // merge, merge, w // merge, merge, C) + return x.view(target).permute(0, 1, 3, 2, 4, 5).reshape(-1, C // (merge * merge)) def forward(self, flatten_patches: torch.Tensor, grid_thw: torch.Tensor) -> tuple[torch.Tensor, list[torch.Tensor]]: """Run the visual tower -> (merged, deepstack)."""