From 4add435bf97faefb13468d1dd6fe663419c3cd0a Mon Sep 17 00:00:00 2001 From: Daniel Maddern Date: Wed, 19 Aug 2026 22:48:26 +0700 Subject: [PATCH] Fix vision merger: main=per-patch-norm+2x2-interleave, deepstack=merged-norm --- src/h3_blackwell_runtime/qwen3vl_vision.py | 71 ++++++++++++++++------ 1 file changed, 52 insertions(+), 19 deletions(-) diff --git a/src/h3_blackwell_runtime/qwen3vl_vision.py b/src/h3_blackwell_runtime/qwen3vl_vision.py index e532b35..44ed372 100644 --- a/src/h3_blackwell_runtime/qwen3vl_vision.py +++ b/src/h3_blackwell_runtime/qwen3vl_vision.py @@ -308,11 +308,18 @@ class _VisionBlock(nn.Module): class _VisionPatchMerger(nn.Module): - """Spatial 2x2 merge + projection to text width (main or deepstack).""" + """Qwen3-VL spatial-merge projector (main or deepstack). - def __init__(self, norm_w: torch.Tensor, norm_b: torch.Tensor, fc1_w: torch.Tensor, fc1_b: torch.Tensor, fc2_w: torch.Tensor, fc2_b: torch.Tensor, *, merge_size: int, out_hidden_size: int): + The main merger applies LayerNorm over ``hidden_size`` (1152) BEFORE the + 2x2 spatial merge; the deepstack merger applies LayerNorm over + ``merge_dim`` (4608) AFTER the merge. This is controlled by ``norm_dim``. + """ + + def __init__(self, norm_w: torch.Tensor, norm_b: torch.Tensor, fc1_w: torch.Tensor, fc1_b: torch.Tensor, fc2_w: torch.Tensor, fc2_b: torch.Tensor, *, merge_size: int, out_hidden_size: int, norm_dim: int | None = None): super().__init__() self.merge_dim = VISION_HIDDEN * (merge_size ** 2) + # Default: norm_dim = merge_dim (deepstack style). Main merger overrides. + self.norm_dim = norm_dim if norm_dim is not None else self.merge_dim self.register_buffer("norm_weight", norm_w, persistent=False) self.register_buffer("norm_bias", norm_b, persistent=False) self.register_buffer("fc1_weight", fc1_w, persistent=False) @@ -322,11 +329,15 @@ class _VisionPatchMerger(nn.Module): self.out_hidden_size = out_hidden_size def forward(self, x: torch.Tensor) -> torch.Tensor: - # x: [seq, 4*hidden] (spatially merged upstream: 2x2 blocks concatenated). - if x.shape[-1] != self.merge_dim: - print(f"[merger] x.shape={tuple(x.shape)}, merge_dim={self.merge_dim}, norm_w={tuple(self.norm_weight.shape)}") - x = F.layer_norm(x, (self.merge_dim,), weight=self.norm_weight, bias=self.norm_bias, eps=1e-6) - x = x.view(-1, self.merge_dim) + # x: [t*h*w, hidden] (unmerged patches) for the main merger; + # [t*(h//2)*(w//2), merge_dim] (pre-merged 2x2) for the deepstack merger. + if x.shape[-1] == self.merge_dim: + # Already 2x2-merged (deepstack style): norm over merge_dim. + x = F.layer_norm(x, (self.merge_dim,), weight=self.norm_weight, bias=self.norm_bias, eps=1e-6) + else: + # Main merger: per-patch LayerNorm over hidden, then group 2x2 into merge_dim. + x = F.layer_norm(x, (x.shape[-1],), weight=self.norm_weight, bias=self.norm_bias, eps=1e-6) + x = x.view(-1, self.merge_dim) return F.linear(F.gelu(F.linear(x, self.fc1_weight, self.fc1_bias), approximate="tanh"), self.fc2_weight, self.fc2_bias) @@ -570,6 +581,7 @@ class Qwen3VL32BVision(nn.Module): get("visual.merger.linear_fc1.weight"), get("visual.merger.linear_fc1.bias"), get("visual.merger.linear_fc2.weight"), get("visual.merger.linear_fc2.bias"), merge_size=self.spatial_merge_size, out_hidden_size=self.out_hidden_size, + norm_dim=VISION_HIDDEN, # main merger: LayerNorm over hidden before 2x2 merge ) self.deepstack_merger_list = nn.ModuleList([ _VisionPatchMerger( @@ -675,15 +687,35 @@ 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).""" + """DeepStack layout: interleave 2x2 spatial neighbours and flatten. + + x: [t*h*w, C] -> [t*(h//2)*(w//2), 4*C]. + """ first = grid_thw[0] t = int(first[0].item()) h = int(first[1].item()) w = int(first[2].item()) C = int(x.shape[-1]) merge = 2 - 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) + x = x.view(t, h // merge, merge, w // merge, merge, C) + x = x.permute(0, 1, 3, 2, 4, 5).contiguous() + return x.reshape(-1, C * merge * merge) + + @staticmethod + def _interleave_2x2(x: torch.Tensor, grid_thw: torch.Tensor) -> torch.Tensor: + """Re-order [t*h*w, C] into 2x2-block-major order (no dim change). + + After this, x.view(-1, 4*C) groups spatially-adjacent 2x2 blocks. + """ + first = grid_thw[0] + t = int(first[0].item()) + h = int(first[1].item()) + w = int(first[2].item()) + C = int(x.shape[-1]) + merge = 2 + x = x.view(t, h // merge, merge, w // merge, merge, C) + x = x.permute(0, 1, 3, 2, 4, 5).contiguous() + return x.reshape(-1, C) # still [t*(h//2)*(w//2)*4, C], block-major def forward(self, flatten_patches: torch.Tensor, grid_thw: torch.Tensor) -> tuple[torch.Tensor, list[torch.Tensor]]: """Run the visual tower -> (merged, deepstack).""" @@ -703,16 +735,17 @@ class Qwen3VL32BVision(nn.Module): deepstack_features = [] for layer_num, block in enumerate(self.blocks): x = block(x, cu_seqlens=cu_seqlens, position_embeddings=position_embeddings) - if layer_num in (0, 8, 16, 24) or layer_num == self.depth - 1: - print(f"[vision] after block {layer_num}: x.shape={tuple(x.shape)}") + # x: [t*h*w, hidden] (unmerged patches). if layer_num in self.deepstack_visual_indexes: - # x should be [N_patches, C]; if not, re-flatten. - if x.ndim != 2: - x = x.reshape(-1, x.shape[-1]) - deepstack_features.append(self.deepstack_merger_list[self.deepstack_visual_indexes.index(layer_num)](self._merge_tokens(x, grid_thw))) - if x.ndim != 2: - x = x.reshape(-1, x.shape[-1]) - return self.merger(self._merge_tokens(x, grid_thw)), deepstack_features + # DeepStack: merge 2x2 first, then project. + deepstack_features.append( + self.deepstack_merger_list[self.deepstack_visual_indexes.index(layer_num)](self._merge_tokens(x, grid_thw)) + ) + # Main merger expects the 2x2-interleaved layout (per-patch norm is done + # over the UNMERGED hidden, but the grouping of 4 patches must be + # spatially-contiguous). Re-order x into 2x2-block-major order first. + x = self._interleave_2x2(x, grid_thw) + return self.merger(x), deepstack_features class _VisionRotary(nn.Module):