From 2316199c97fbb99a0df33ec3cc6c54711d9416aa Mon Sep 17 00:00:00 2001 From: Claude Date: Fri, 31 Jul 2026 03:52:23 +0000 Subject: [PATCH] =?UTF-8?q?[FIX]=20V2=20shape=20mismatch=20bug=20=E2=80=94?= =?UTF-8?q?=20v=5Fpadded=20used=20num=5Fheads=20for=20kv=5Fh=20tensor?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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.). --- paged_attention_v2_pytorch.py | 23 +++++++++++++---------- 1 file changed, 13 insertions(+), 10 deletions(-) diff --git a/paged_attention_v2_pytorch.py b/paged_attention_v2_pytorch.py index b3cf22c4..b097f49d 100644 --- a/paged_attention_v2_pytorch.py +++ b/paged_attention_v2_pytorch.py @@ -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)