Fix vision merger: main=per-patch-norm+2x2-interleave, deepstack=merged-norm

This commit is contained in:
Daniel Maddern 2026-08-19 22:48:26 +07:00
parent 68c49fdc95
commit 4add435bf9

View file

@ -308,11 +308,18 @@ class _VisionBlock(nn.Module):
class _VisionPatchMerger(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__() super().__init__()
self.merge_dim = VISION_HIDDEN * (merge_size ** 2) 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_weight", norm_w, persistent=False)
self.register_buffer("norm_bias", norm_b, persistent=False) self.register_buffer("norm_bias", norm_b, persistent=False)
self.register_buffer("fc1_weight", fc1_w, 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 self.out_hidden_size = out_hidden_size
def forward(self, x: torch.Tensor) -> torch.Tensor: def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: [seq, 4*hidden] (spatially merged upstream: 2x2 blocks concatenated). # x: [t*h*w, hidden] (unmerged patches) for the main merger;
if x.shape[-1] != self.merge_dim: # [t*(h//2)*(w//2), merge_dim] (pre-merged 2x2) for the deepstack merger.
print(f"[merger] x.shape={tuple(x.shape)}, merge_dim={self.merge_dim}, norm_w={tuple(self.norm_weight.shape)}") if x.shape[-1] == self.merge_dim:
x = F.layer_norm(x, (self.merge_dim,), weight=self.norm_weight, bias=self.norm_bias, eps=1e-6) # Already 2x2-merged (deepstack style): norm over merge_dim.
x = x.view(-1, self.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) 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_fc1.weight"), get("visual.merger.linear_fc1.bias"),
get("visual.merger.linear_fc2.weight"), get("visual.merger.linear_fc2.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, 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([ self.deepstack_merger_list = nn.ModuleList([
_VisionPatchMerger( _VisionPatchMerger(
@ -675,15 +687,35 @@ class Qwen3VL32BVision(nn.Module):
@staticmethod @staticmethod
def _merge_tokens(x: torch.Tensor, grid_thw: torch.Tensor) -> torch.Tensor: 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] first = grid_thw[0]
t = int(first[0].item()) t = int(first[0].item())
h = int(first[1].item()) h = int(first[1].item())
w = int(first[2].item()) w = int(first[2].item())
C = int(x.shape[-1]) C = int(x.shape[-1])
merge = 2 merge = 2
target = (t, h // merge, merge, w // merge, merge, C) x = x.view(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.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]]: def forward(self, flatten_patches: torch.Tensor, grid_thw: torch.Tensor) -> tuple[torch.Tensor, list[torch.Tensor]]:
"""Run the visual tower -> (merged, deepstack).""" """Run the visual tower -> (merged, deepstack)."""
@ -703,16 +735,17 @@ class Qwen3VL32BVision(nn.Module):
deepstack_features = [] deepstack_features = []
for layer_num, block in enumerate(self.blocks): for layer_num, block in enumerate(self.blocks):
x = block(x, cu_seqlens=cu_seqlens, position_embeddings=position_embeddings) 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: # x: [t*h*w, hidden] (unmerged patches).
print(f"[vision] after block {layer_num}: x.shape={tuple(x.shape)}")
if layer_num in self.deepstack_visual_indexes: if layer_num in self.deepstack_visual_indexes:
# x should be [N_patches, C]; if not, re-flatten. # DeepStack: merge 2x2 first, then project.
if x.ndim != 2: deepstack_features.append(
x = x.reshape(-1, x.shape[-1]) self.deepstack_merger_list[self.deepstack_visual_indexes.index(layer_num)](self._merge_tokens(x, grid_thw))
deepstack_features.append(self.deepstack_merger_list[self.deepstack_visual_indexes.index(layer_num)](self._merge_tokens(x, grid_thw))) )
if x.ndim != 2: # Main merger expects the 2x2-interleaved layout (per-patch norm is done
x = x.reshape(-1, x.shape[-1]) # over the UNMERGED hidden, but the grouping of 4 patches must be
return self.merger(self._merge_tokens(x, grid_thw)), deepstack_features # 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): class _VisionRotary(nn.Module):