diff --git a/src/h3_blackwell_runtime/qwen3vl_vision.py b/src/h3_blackwell_runtime/qwen3vl_vision.py index 8e584a0..e4b5a9b 100644 --- a/src/h3_blackwell_runtime/qwen3vl_vision.py +++ b/src/h3_blackwell_runtime/qwen3vl_vision.py @@ -674,11 +674,13 @@ class Qwen3VL32BVision(nn.Module): @staticmethod def _merge_tokens(x: torch.Tensor, grid_thw: torch.Tensor) -> torch.Tensor: """Flatten 2x2 spatial-merge blocks (the reference's pre-merge flatten).""" - t, h, w = (int(v) for v in grid_thw.tolist()[0]) + # grid_thw is a 2D long tensor [[t, h, w], ...]; take the first image. + first = grid_thw[0] + t = int(first[0].item()) + h = int(first[1].item()) + w = int(first[2].item()) merge = 2 - if x.ndim != 2 or x.shape[0] != t * h * w: - print(f"[merge_tokens] x.shape={tuple(x.shape)}, grid={(t, h, w)}, expected N={t*h*w}") - return x.view(t, h // merge, merge, w // merge, merge, -1).permute(0, 1, 3, 2, 4, 5).reshape(-1, -1) + return x.view((t, h // merge, merge, w // merge, merge, x.shape[-1])).permute(0, 1, 3, 2, 4, 5).reshape(-1, -1) def forward(self, flatten_patches: torch.Tensor, grid_thw: torch.Tensor) -> tuple[torch.Tensor, list[torch.Tensor]]: """Run the visual tower -> (merged, deepstack)."""