Use merge_dim as normalized_shape in _VisionPatchMerger (4*1152 for 2x2 merge)

This commit is contained in:
Daniel Maddern 2026-08-19 22:39:05 +07:00
parent b085ef02e8
commit 68c49fdc95

View file

@ -322,8 +322,10 @@ class _VisionPatchMerger(nn.Module):
self.out_hidden_size = out_hidden_size
def forward(self, x: torch.Tensor) -> torch.Tensor:
# x: [seq, hidden] (already spatially merged upstream).
x = F.layer_norm(x, (VISION_HIDDEN,), weight=self.norm_weight, bias=self.norm_bias, eps=1e-6)
# 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)
return F.linear(F.gelu(F.linear(x, self.fc1_weight, self.fc1_bias), approximate="tanh"), self.fc2_weight, self.fc2_bias)