Use explicit C (not x.shape[-1]) in _merge_tokens
This commit is contained in:
parent
99a7550a51
commit
571d3541e8
1 changed files with 3 additions and 4 deletions
|
|
@ -678,11 +678,10 @@ class Qwen3VL32BVision(nn.Module):
|
||||||
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])
|
||||||
merge = 2
|
merge = 2
|
||||||
print(f"[merge_tokens] x.shape={tuple(x.shape)}, t={t}, h={h}, w={w}, "
|
target = (t, h // merge, merge, w // merge, merge, C)
|
||||||
f"target=({t}, {h//merge}, {merge}, {w//merge}, {merge}, {x.shape[-1]}), "
|
return x.view(target).permute(0, 1, 3, 2, 4, 5).reshape(-1, C // (merge * merge))
|
||||||
f"prod={t*(h//merge)*merge*(w//merge)*merge*x.shape[-1]}", flush=True)
|
|
||||||
return x.view((t, h // merge, merge, w // merge, merge, x.shape[-1])).permute(0, 1, 3, 2, 4, 5).reshape(-1, -1)
|
|
||||||
|
|
||||||
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)."""
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue