Fix vision merger: main=per-patch-norm+2x2-interleave, deepstack=merged-norm
This commit is contained in:
parent
68c49fdc95
commit
4add435bf9
1 changed files with 52 additions and 19 deletions
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue