[OPT] GQA broadcast in V2 — eliminate 1GB/step memory allocation
Qwen3.6: num_heads=24, num_kv_heads=4, gqa_ratio=6, head_dim=256 Before (expand GQA then bmm): k_flat: [100K, 4, 256] → expand to [100K, 24, 256] → contiguous Memory: 100K × 24 × 256 × 2B = 1.2GB allocated per decode step Then: [24, 256, 100K] @ [24, 1, 256]^T → scores After (broadcast without materializing): k_kv: [100K, 4, 256] → [4, 256, 100K] (no expansion) q: [24, 256] → [4, 6, 1, 256] scores: matmul([4, 6, 1, 256], [4, 1, 256, 100K]) → [4, 6, 100K] Broadcasting handles GQA — K stays at kv_heads size. Memory: 100K × 4 × 256 × 2B = 200MB (6x reduction) For 100K context generating 1000 tokens: Old: 1000 × 1.2GB = 1.2TB total memory traffic for GQA expansion alone New: 1000 × 200MB = 200GB total (saved 1TB of unnecessary data movement) V weighted sum still needs GQA expansion (V @ scores requires matching dims), but the dominant cost (Q @ K^T) is now broadcast.
This commit is contained in:
@@ -99,33 +99,32 @@ def paged_attention_v2_pytorch(
|
||||
if v_scale != 1.0:
|
||||
v_flat = v_flat.float().mul_(v_scale)
|
||||
|
||||
# GQA: expand (zero-copy view) then reshape to contiguous for bmm
|
||||
# [seq_len, kv_h, d] → [seq_len, H, d]
|
||||
if gqa_ratio > 1:
|
||||
k_all = (k_flat.unsqueeze(2)
|
||||
.expand(-1, -1, gqa_ratio, -1)
|
||||
.reshape(seq_len, num_heads, head_size))
|
||||
v_all = (v_flat.unsqueeze(2)
|
||||
.expand(-1, -1, gqa_ratio, -1)
|
||||
.reshape(seq_len, num_heads, head_size))
|
||||
else:
|
||||
k_all = k_flat
|
||||
v_all = v_flat
|
||||
|
||||
# =============================================================
|
||||
# Phase 1: ALL partitions in ONE bmm (CCCL transform_reduce pattern)
|
||||
# GQA broadcast: avoid materializing the expanded KV tensor
|
||||
#
|
||||
# Instead of: for p in range(195): bmm(Q, K_p)
|
||||
# We do: scores = Q @ K_all^T → [H, seq_len]
|
||||
# reshape to [H, P, part_sz] → partition-wise softmax
|
||||
#
|
||||
# This is one kernel launch vs 195.
|
||||
# Qwen3.6: H=24, kv_h=4, gqa_ratio=6, head_dim=256
|
||||
# Old: expand kv_h→H then contiguous → allocates seq_len×H×d (1.2GB at 100K)
|
||||
# New: reshape Q as [kv_h, gqa, 1, d], K as [kv_h, 1, d, seq_len]
|
||||
# → bmm with broadcasting → [kv_h, gqa, 1, seq_len]
|
||||
# → reshape to [H, seq_len]
|
||||
# Saves: gqa_ratio × memory (6x for Qwen3.6 = 1GB per decode step)
|
||||
# =============================================================
|
||||
q = query[seq_idx].float() # [H, d]
|
||||
|
||||
# Q @ K^T: [H, 1, d] @ [H, d, seq_len] → [H, 1, seq_len] → [H, seq_len]
|
||||
k_t = k_all.permute(1, 2, 0).float().contiguous() # [H, d, seq_len]
|
||||
scores_all = torch.bmm(q.unsqueeze(1), k_t).squeeze(1) * scale # [H, seq_len]
|
||||
if gqa_ratio > 1:
|
||||
# K: [seq_len, kv_h, d] → [kv_h, d, seq_len] (no GQA expansion)
|
||||
k_kv = k_flat.permute(1, 2, 0).float().contiguous() # [kv_h, d, seq_len]
|
||||
v_kv = v_flat.permute(1, 0, 2).float().contiguous() # [kv_h, seq_len, d]
|
||||
|
||||
# Q: [H, d] → [kv_h, gqa, 1, d]
|
||||
q_grouped = q.view(num_kv_heads, gqa_ratio, 1, head_size)
|
||||
|
||||
# Scores: [kv_h, gqa, 1, d] @ [kv_h, 1, d, seq_len] → [kv_h, gqa, 1, seq_len]
|
||||
scores_all = torch.matmul(q_grouped, k_kv.unsqueeze(1)).squeeze(2) # [kv_h, gqa, seq_len]
|
||||
scores_all = scores_all.reshape(num_heads, seq_len) * scale # [H, seq_len]
|
||||
else:
|
||||
k_t = k_flat.permute(1, 2, 0).float().contiguous() # [H, d, seq_len]
|
||||
scores_all = torch.bmm(q.unsqueeze(1), k_t).squeeze(1) * scale # [H, seq_len]
|
||||
|
||||
# Alibi bias (if needed)
|
||||
if alibi_slopes is not None:
|
||||
@@ -154,7 +153,11 @@ def paged_attention_v2_pytorch(
|
||||
|
||||
# Weighted values per partition: need V reshaped the same way
|
||||
# V: [seq_len, H, d] → pad → [padded_len, H, d] → [H, P, part_sz, d]
|
||||
v_perm = v_all.permute(1, 0, 2).float().contiguous() # [H, seq_len, d]
|
||||
if gqa_ratio > 1:
|
||||
v_perm = v_kv # already [kv_h, seq_len, d], no GQA expansion needed
|
||||
# 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),
|
||||
@@ -164,13 +167,37 @@ def paged_attention_v2_pytorch(
|
||||
v_padded = v_perm
|
||||
v_parts = v_padded.view(num_heads, num_partitions, _PARTITION_SIZE, head_size)
|
||||
|
||||
# Weighted sum: [H, P, 1, part_sz] @ [H, P, part_sz, d] → [H, P, 1, d] → [H, P, d]
|
||||
# Reshape for batched bmm: [H*P, 1, part_sz] @ [H*P, part_sz, d] → [H*P, 1, d]
|
||||
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)
|
||||
part_out_flat = torch.bmm(scores_exp_flat, v_parts_flat) # [HP, 1, d]
|
||||
part_out = part_out_flat.view(num_heads, num_partitions, head_size) # [H, P, d]
|
||||
# Weighted V sum per partition
|
||||
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]
|
||||
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, :]
|
||||
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)
|
||||
else:
|
||||
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)
|
||||
part_out_flat = torch.bmm(scores_exp_flat, v_parts_flat) # [HP, 1, d]
|
||||
part_out = part_out_flat.view(num_heads, num_partitions, head_size) # [H, P, d]
|
||||
|
||||
# Store partition results
|
||||
max_logits[seq_idx, :, :num_partitions] = part_max
|
||||
|
||||
Reference in New Issue
Block a user