[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:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user