Use merge_dim as normalized_shape in _VisionPatchMerger (4*1152 for 2x2 merge)
This commit is contained in:
parent
b085ef02e8
commit
68c49fdc95
1 changed files with 4 additions and 2 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue