[OPT] Complete Triton V2 Phase 1 — paged K/V gather from prefix_prefill.py pattern
Phase 1 kernel (_paged_attn_v2_partition_kernel) now has complete
paged K/V gather implementation, adapted from prefix_prefill.py:
K gather:
bn = tl.load(block_tables + seq*stride + (token//block_size)*stride)
off_k = bn * stride_kc_b + kv_head * stride_kc_h +
(d//x) * stride_kc_dx + (token%block_size) * stride_kc_bs +
(d%x) * stride_kc_x
k = tl.load(key_cache + off_k, mask=valid)
V gather (simpler layout):
off_v = bn * stride_vc_b + kv_head * stride_vc_h +
d * stride_vc_d + (token%block_size) * stride_vc_bs
Online softmax (Flash Attention pattern):
m_i_new = max(m_i, max(scores))
alpha = exp(m_i - m_i_new)
acc = acc * alpha * l_i / l_i_new + (p/l_i_new * beta) @ V
Key difference from prefix_prefill.py:
- BLOCK_M=1 (decode: 1 query token) vs BLOCK_M>1 (prefill)
- q @ k is dot product [D]•[D,N] → [N], not matrix [M,D]@[D,N] → [M,N]
- head_dim=256 support: BLOCK_N=32 (vs 64 for head_dim=128)
32×256×2×2 = 32KB ≤ 48KB SMEM ✓
Integration: Triton V2 tried first, PyTorch V2 as fallback.
If Triton works on BI-V100: single GPU launch for all partitions
(grid = num_seqs × num_heads × num_partitions = 1 × 24 × 200 = 4800 blocks)
vs PyTorch's 3 bmm launches.
This commit is contained in:
@@ -6,6 +6,7 @@ WORKDIR /workspace/
|
||||
# Copy all scripts and the V2 module
|
||||
COPY ./qwen3_6_scripts /workspace/qwen3_6_scripts
|
||||
COPY ./paged_attention_v2_pytorch.py /workspace/paged_attention_v2_pytorch.py
|
||||
COPY ./paged_attention_v2_triton.py /workspace/paged_attention_v2_triton.py
|
||||
|
||||
# Run baseline patches (model registration, xformers fallback, tool parser, etc.)
|
||||
RUN cd ./qwen3_6_scripts && ./patch_ops.sh
|
||||
|
||||
Reference in New Issue
Block a user