[FIX] V2 shape mismatch bug — v_padded used num_heads for kv_h tensor

Bug: After GQA broadcast optimization, v_perm was [kv_h, seq_len, d]
in the GQA path, but unconditional v_padded allocation used num_heads:
  v_padded = torch.zeros((num_heads, padded_len, head_size))
  v_padded[:, :seq_len, :] = v_perm  # [24, padded, d] vs [4, seq, d] → CRASH

Fix: v_padded/v_parts allocation is now inside the non-GQA else branch.
GQA branch uses its own v_padded_kv with correct [kv_h, padded, d] shape.

This was a real runtime bug — V2 would have crashed on first call
for any GQA model (Qwen3.6, Llama, etc.).
This commit is contained in:
Claude
2026-07-31 03:52:23 +00:00
parent cd0d9e1a91
commit 2316199c97

View File

@@ -158,16 +158,10 @@ def paged_attention_v2_pytorch(
# Will handle GQA in the bmm below via broadcast
else:
v_perm = v_flat.permute(1, 0, 2).float().contiguous() # [H, seq_len, d]
if padded_len > seq_len:
v_padded = torch.zeros(
(num_heads, padded_len, head_size),
dtype=v_perm.dtype, device=v_perm.device)
v_padded[:, :seq_len, :] = v_perm
else:
v_padded = v_perm
v_parts = v_padded.view(num_heads, num_partitions, _PARTITION_SIZE, head_size)
# Weighted V sum per partition — GQA broadcast (avoid 2.4GB expansion)
# Weighted V sum per partition
# NOTE: v_perm shape differs by GQA mode:
# GQA: v_perm = v_kv = [kv_h, seq_len, d]
# No GQA: v_perm = [H, seq_len, d]
# scores_exp: [H, P, part_sz] → [kv_h, gqa, P, part_sz]
# v_perm: [kv_h, seq_len, d] → [kv_h, P, part_sz, d]
if gqa_ratio > 1:
@@ -189,6 +183,15 @@ def paged_attention_v2_pytorch(
).squeeze(3) # [kv_h, gqa, P, d]
part_out = part_out_grouped.reshape(num_heads, num_partitions, head_size)
else:
# Non-GQA: v_perm is [H, seq_len, d], pad and reshape normally
if padded_len > seq_len:
v_padded = torch.zeros(
(num_heads, padded_len, head_size),
dtype=v_perm.dtype, device=v_perm.device)
v_padded[:, :seq_len, :] = v_perm
else:
v_padded = v_perm
v_parts = v_padded.view(num_heads, num_partitions, _PARTITION_SIZE, head_size)
HP = num_heads * num_partitions
scores_exp_flat = scores_exp.reshape(HP, 1, _PARTITION_SIZE)
v_parts_flat = v_parts.reshape(HP, _PARTITION_SIZE, head_size)