[OPT] Complete GQA broadcast — V weighted sum also avoids expansion

Previous commit broadcast Q@K^T (saved 1GB/step).
This commit broadcasts scores@V too (saves 2GB/step).

Before: V expanded from [kv_h, padded_len, d] to [H, padded_len, d]
  4×100K×256×4B → 24×100K×256×4B = 400MB → 2.4GB allocation

After: broadcast matmul at kv_h level
  se: [kv_h, gqa, P, 1, part_sz] @ V: [kv_h, 1, P, part_sz, d]
  → [kv_h, gqa, P, 1, d] → reshape to [H, P, d]
  V stays at kv_h size: 400MB (no 2.4GB allocation)

Total per-decode-step memory for 100K context:
  Before all GQA opts: 3.6GB (K expansion + V expansion)
  After: 600MB (6x total reduction from GQA ratio=6)

This is the CCCL insight applied: transform_reduce with a compound type.
Instead of expanding to full head count then reducing, keep the reduction
at the minimal group size and broadcast the grouping dimension.
This commit is contained in:
Claude
2026-07-30 16:15:34 +00:00
parent 8951d74936
commit d9bbef54d8

View File

@@ -167,31 +167,27 @@ def paged_attention_v2_pytorch(
v_padded = v_perm
v_parts = v_padded.view(num_heads, num_partitions, _PARTITION_SIZE, head_size)
# Weighted V sum per partition
# Weighted V sum per partition — GQA broadcast (avoid 2.4GB expansion)
# 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:
# v_parts: [kv_h, P, part_sz, d] — no GQA expansion
# scores_exp: [H, P, part_sz] → [kv_h, gqa, P, part_sz]
se_grouped = scores_exp.view(num_kv_heads, gqa_ratio, num_partitions, _PARTITION_SIZE)
# [kv_h, gqa, P, 1, part_sz] @ [kv_h, 1, P, part_sz, d] → [kv_h, gqa, P, 1, d]
# Simpler: loop over partitions (they're already separate in softmax)
# Actually just reshape both for HP bmm:
HP = num_heads * num_partitions
se_flat = scores_exp.reshape(HP, 1, _PARTITION_SIZE)
# Need V in [H, P, part_sz, d] — must expand V for GQA here
v_expanded = (v_perm.unsqueeze(1)
.expand(-1, gqa_ratio, -1, -1)
.reshape(num_heads, -1, head_size)) # [H, padded_len, d]
# V: pad and reshape to [kv_h, P, part_sz, d]
if padded_len > seq_len:
v_padded_h = torch.zeros(
(num_heads, padded_len, head_size),
dtype=v_expanded.dtype, device=v_expanded.device)
v_padded_h[:, :seq_len, :] = v_expanded[:, :seq_len, :]
v_padded_kv = torch.zeros(
(num_kv_heads, padded_len, head_size),
dtype=v_kv.dtype, device=v_kv.device)
v_padded_kv[:, :seq_len, :] = v_kv
else:
v_padded_h = v_expanded
v_parts_h = v_padded_h.view(num_heads, num_partitions, _PARTITION_SIZE, head_size)
vp_flat = v_parts_h.reshape(HP, _PARTITION_SIZE, head_size)
part_out_flat = torch.bmm(se_flat, vp_flat) # [HP, 1, d]
part_out = part_out_flat.view(num_heads, num_partitions, head_size)
v_padded_kv = v_kv
v_parts_kv = v_padded_kv.view(num_kv_heads, num_partitions, _PARTITION_SIZE, head_size)
# Broadcast: [kv_h, gqa, P, 1, part_sz] @ [kv_h, 1, P, part_sz, d]
# → [kv_h, gqa, P, 1, d]
part_out_grouped = torch.matmul(
se_grouped.unsqueeze(3), # [kv_h, gqa, P, 1, part_sz]
v_parts_kv.unsqueeze(1) # [kv_h, 1, P, part_sz, d]
).squeeze(3) # [kv_h, gqa, P, d]
part_out = part_out_grouped.reshape(num_heads, num_partitions, head_size)
else:
HP = num_heads * num_partitions
scores_exp_flat = scores_exp.reshape(HP, 1, _PARTITION_SIZE)