Two-kernel design following vllm's paged_attention_v2_kernel.cu: Phase 1: _paged_attn_v2_partition_kernel grid = (num_seqs, num_heads, num_partitions) Each instance: Q[head] @ K[partition]^T → softmax → @ V[partition] Status: SKELETON — paged K/V gather from indirect block_tables is complex in Triton (requires scatter/gather through block_tables). Currently falls back to PyTorch partition loop. Phase 2: _paged_attn_v2_reduce_kernel grid = (num_seqs, num_heads) Each instance: log-sum-exp reduction across partitions Status: COMPLETE — replaces Python einsum with single Triton launch. Algorithm: global_max → rescale → weighted sum (same pattern as CCCL summary_statistics binary_op for combining partial statistics). SMEM: Phase 1 needs BLOCK_N=64 × head_dim=128 × 2B × 2 = 32KB ≤ 48KB. Phase 2 needs no SMEM (partitions fit in registers). The Phase 1 paged gather is the hard part. The key_cache layout [blocks, kv_heads, head_dim/x, block_size, x] requires: 1. block_tables[seq, token // block_size] → physical_block_id 2. key_cache[physical_block_id, kv_head, :, token % block_size, :] This is indirect indexed access — possible in Triton via tl.load with computed offsets, but needs careful stride arithmetic.
305 lines
12 KiB
Python
305 lines
12 KiB
Python
"""
|
|
paged_attention_v2_triton.py — Triton kernel for PagedAttention V2 on BI-V100
|
|
================================================================================
|
|
|
|
Replaces the Python partition loop with a single Triton kernel launch.
|
|
|
|
Phase 1 kernel: paged_attn_v2_partition
|
|
grid = (num_seqs, num_heads, num_partitions)
|
|
Each program instance computes attention for one (seq, head, partition).
|
|
|
|
Algorithm per instance:
|
|
1. Load Q vector for this (seq, head): [head_dim]
|
|
2. Load K/V from paged cache for this partition's token range
|
|
3. Compute QK^T scores, online softmax max + sum
|
|
4. Compute weighted V output
|
|
5. Store: tmp_output[seq, head, part, :], exp_sums[seq, head, part], max_logits[seq, head, part]
|
|
|
|
Phase 2 kernel: paged_attn_v2_reduce
|
|
grid = (num_seqs, num_heads)
|
|
Each program instance reduces across partitions for one (seq, head).
|
|
|
|
Algorithm:
|
|
1. Load max_logits[seq, head, :num_parts] → find global_max
|
|
2. Rescale: weights[p] = exp(max[p] - global_max) * sum[p]
|
|
3. Normalize and weighted sum of tmp_output
|
|
|
|
SMEM analysis:
|
|
Phase 1: K tile [BLOCK_N, head_dim] + V tile [BLOCK_N, head_dim] in SMEM
|
|
At BLOCK_N=64, head_dim=128, fp16: 64*128*2*2 = 32KB ≤ 48KB ✓
|
|
Phase 2: No SMEM needed (max_partitions ≈ 200, fits in registers)
|
|
|
|
Deploy:
|
|
This kernel requires Triton to be functional on BI-V100.
|
|
patch_enable_triton.py already enables Triton with try/fallback.
|
|
If Triton works, this kernel replaces the Python V2 for decode.
|
|
If Triton doesn't work, fall back to paged_attention_v2_pytorch.py.
|
|
"""
|
|
|
|
import torch
|
|
import triton
|
|
import triton.language as tl
|
|
from typing import Optional
|
|
|
|
|
|
@triton.jit
|
|
def _paged_attn_v2_partition_kernel(
|
|
# Outputs
|
|
tmp_output_ptr, # [num_seqs, num_heads, max_num_parts, head_size]
|
|
exp_sums_ptr, # [num_seqs, num_heads, max_num_parts]
|
|
max_logits_ptr, # [num_seqs, num_heads, max_num_parts]
|
|
# Inputs
|
|
query_ptr, # [num_seqs, num_heads, head_size]
|
|
key_cache_ptr, # [num_blocks, num_kv_heads, head_size/x, block_size, x]
|
|
value_cache_ptr, # [num_blocks, num_kv_heads, head_size, block_size]
|
|
block_tables_ptr, # [num_seqs, max_blocks_per_seq]
|
|
seq_lens_ptr, # [num_seqs]
|
|
# Scalars
|
|
scale,
|
|
num_kv_heads,
|
|
block_size,
|
|
max_blocks_per_seq,
|
|
max_num_parts,
|
|
# Strides
|
|
stride_qt_s, stride_qt_h, stride_qt_d,
|
|
stride_kc_b, stride_kc_h, stride_kc_dx, stride_kc_bs, stride_kc_x,
|
|
stride_vc_b, stride_vc_h, stride_vc_d, stride_vc_bs,
|
|
stride_bt_s, stride_bt_b,
|
|
stride_to_s, stride_to_h, stride_to_p, stride_to_d,
|
|
stride_es_s, stride_es_h, stride_es_p,
|
|
# Constants
|
|
PARTITION_SIZE: tl.constexpr,
|
|
HEAD_DIM: tl.constexpr,
|
|
BLOCK_N: tl.constexpr, # KV tokens processed per inner loop iteration
|
|
X_PACK: tl.constexpr, # key cache packing factor (16 // element_size)
|
|
):
|
|
"""Phase 1: Per-partition attention computation.
|
|
|
|
Each program computes attention for one (seq, head, partition).
|
|
Iterates over BLOCK_N tokens at a time within the partition.
|
|
Uses online softmax (Flash Attention style) to compute max, sum, and weighted V.
|
|
"""
|
|
seq_idx = tl.program_id(0)
|
|
head_idx = tl.program_id(1)
|
|
part_idx = tl.program_id(2)
|
|
|
|
seq_len = tl.load(seq_lens_ptr + seq_idx)
|
|
|
|
# This partition's token range
|
|
part_start = part_idx * PARTITION_SIZE
|
|
part_end = tl.minimum(part_start + PARTITION_SIZE, seq_len)
|
|
|
|
if part_start >= seq_len:
|
|
# This partition is beyond the sequence length — write -inf/0
|
|
tl.store(max_logits_ptr + seq_idx * stride_es_s + head_idx * stride_es_h + part_idx * stride_es_p,
|
|
float('-inf'))
|
|
tl.store(exp_sums_ptr + seq_idx * stride_es_s + head_idx * stride_es_h + part_idx * stride_es_p,
|
|
0.0)
|
|
return
|
|
|
|
# GQA: map head_idx to kv_head_idx
|
|
num_queries_per_kv = (tl.program_id(1) + 1) # placeholder — need actual num_heads/num_kv_heads
|
|
kv_head_idx = head_idx // (stride_qt_h // stride_kc_h) if stride_kc_h > 0 else head_idx # TODO: fix GQA mapping
|
|
|
|
# Load query: [HEAD_DIM]
|
|
q_offsets = seq_idx * stride_qt_s + head_idx * stride_qt_h + tl.arange(0, HEAD_DIM) * stride_qt_d
|
|
q = tl.load(query_ptr + q_offsets).to(tl.float32)
|
|
|
|
# Online softmax state
|
|
m_i = float('-inf') # running max
|
|
l_i = 0.0 # running sum of exp
|
|
# Accumulator for weighted V: [HEAD_DIM]
|
|
acc = tl.zeros([HEAD_DIM], dtype=tl.float32)
|
|
|
|
# Iterate over KV tokens in this partition, BLOCK_N at a time
|
|
for token_start in range(part_start, part_end, BLOCK_N):
|
|
token_end = tl.minimum(token_start + BLOCK_N, part_end)
|
|
n_tokens = token_end - token_start
|
|
|
|
# For each token, find its physical block and offset
|
|
token_offsets = tl.arange(0, BLOCK_N)
|
|
valid_mask = token_offsets < n_tokens
|
|
|
|
global_token_ids = token_start + token_offsets
|
|
block_indices = global_token_ids // block_size
|
|
within_block_offsets = global_token_ids % block_size
|
|
|
|
# Look up physical block numbers from block_table
|
|
bt_offsets = seq_idx * stride_bt_s + block_indices * stride_bt_b
|
|
physical_blocks = tl.load(block_tables_ptr + bt_offsets, mask=valid_mask, other=0)
|
|
|
|
# Load K for these tokens: need to gather from paged cache
|
|
# K shape: [num_blocks, num_kv_heads, head_size/x, block_size, x]
|
|
# For each token, load K[physical_block, kv_head, :, within_block_offset, :]
|
|
# → [BLOCK_N, HEAD_DIM]
|
|
|
|
# Compute QK^T scores for this chunk
|
|
# scores[n] = sum_d(q[d] * k[n, d]) * scale
|
|
# This requires loading K values — which is complex with paged layout
|
|
# TODO: implement the actual paged K gather in Triton
|
|
# For now, this is a skeleton showing the algorithm structure
|
|
|
|
# --- Placeholder: scores computation ---
|
|
# In a full implementation, we would:
|
|
# 1. For each token n in [0, BLOCK_N):
|
|
# a. physical_block = block_tables[seq, global_token_ids[n] // block_size]
|
|
# b. offset = global_token_ids[n] % block_size
|
|
# c. k[n, :] = key_cache[physical_block, kv_head, :, offset, :].reshape(HEAD_DIM)
|
|
# 2. scores = q @ k.T * scale
|
|
# 3. Online softmax update
|
|
# 4. Load V similarly, accumulate weighted V
|
|
pass
|
|
|
|
# Store results
|
|
tl.store(max_logits_ptr + seq_idx * stride_es_s + head_idx * stride_es_h + part_idx * stride_es_p,
|
|
m_i)
|
|
tl.store(exp_sums_ptr + seq_idx * stride_es_s + head_idx * stride_es_h + part_idx * stride_es_p,
|
|
l_i)
|
|
|
|
# Store accumulated output
|
|
out_offsets = (seq_idx * stride_to_s + head_idx * stride_to_h +
|
|
part_idx * stride_to_p + tl.arange(0, HEAD_DIM) * stride_to_d)
|
|
tl.store(tmp_output_ptr + out_offsets, acc.to(tmp_output_ptr.dtype.element_ty))
|
|
|
|
|
|
@triton.jit
|
|
def _paged_attn_v2_reduce_kernel(
|
|
# Output
|
|
output_ptr, # [num_seqs, num_heads, head_size]
|
|
# Inputs
|
|
tmp_output_ptr, # [num_seqs, num_heads, max_num_parts, head_size]
|
|
exp_sums_ptr, # [num_seqs, num_heads, max_num_parts]
|
|
max_logits_ptr, # [num_seqs, num_heads, max_num_parts]
|
|
seq_lens_ptr, # [num_seqs]
|
|
# Scalars
|
|
max_num_parts,
|
|
# Strides
|
|
stride_out_s, stride_out_h, stride_out_d,
|
|
stride_to_s, stride_to_h, stride_to_p, stride_to_d,
|
|
stride_es_s, stride_es_h, stride_es_p,
|
|
# Constants
|
|
PARTITION_SIZE: tl.constexpr,
|
|
HEAD_DIM: tl.constexpr,
|
|
MAX_NUM_PARTS: tl.constexpr,
|
|
):
|
|
"""Phase 2: Cross-partition reduction.
|
|
|
|
Each program reduces across partitions for one (seq, head).
|
|
Numerically stable log-sum-exp combination.
|
|
|
|
This corresponds to CCCL's summary_statistics binary_op pattern:
|
|
combining partial statistics from independent segments.
|
|
"""
|
|
seq_idx = tl.program_id(0)
|
|
head_idx = tl.program_id(1)
|
|
|
|
seq_len = tl.load(seq_lens_ptr + seq_idx)
|
|
num_parts = (seq_len + PARTITION_SIZE - 1) // PARTITION_SIZE
|
|
|
|
# Load all partition max_logits and exp_sums
|
|
part_offsets = tl.arange(0, MAX_NUM_PARTS)
|
|
valid_mask = part_offsets < num_parts
|
|
|
|
ml_base = seq_idx * stride_es_s + head_idx * stride_es_h
|
|
part_max = tl.load(max_logits_ptr + ml_base + part_offsets * stride_es_p,
|
|
mask=valid_mask, other=float('-inf'))
|
|
part_sum = tl.load(exp_sums_ptr + ml_base + part_offsets * stride_es_p,
|
|
mask=valid_mask, other=0.0)
|
|
|
|
# Global max across partitions
|
|
global_max = tl.max(part_max, axis=0)
|
|
|
|
# Rescale: weights[p] = exp(max[p] - global_max) * sum[p]
|
|
rescale = tl.exp(part_max - global_max) * part_sum
|
|
total = tl.sum(rescale, axis=0)
|
|
weights = rescale / total # [MAX_NUM_PARTS]
|
|
|
|
# Weighted combination of partition outputs
|
|
# For each dimension d in HEAD_DIM:
|
|
# output[d] = sum_p(weights[p] * tmp_output[seq, head, p, d])
|
|
for d in range(HEAD_DIM):
|
|
to_base = seq_idx * stride_to_s + head_idx * stride_to_h + d * stride_to_d
|
|
part_vals = tl.load(tmp_output_ptr + to_base + part_offsets * stride_to_p,
|
|
mask=valid_mask, other=0.0)
|
|
val = tl.sum(weights * part_vals, axis=0)
|
|
tl.store(output_ptr + seq_idx * stride_out_s + head_idx * stride_out_h + d * stride_out_d,
|
|
val)
|
|
|
|
|
|
def paged_attention_v2_triton(
|
|
output: torch.Tensor,
|
|
exp_sums: torch.Tensor,
|
|
max_logits: torch.Tensor,
|
|
tmp_output: torch.Tensor,
|
|
query: torch.Tensor,
|
|
key_cache: torch.Tensor,
|
|
value_cache: torch.Tensor,
|
|
num_kv_heads: int,
|
|
scale: float,
|
|
block_tables: torch.Tensor,
|
|
seq_lens: torch.Tensor,
|
|
block_size: int,
|
|
max_seq_len: int,
|
|
alibi_slopes: Optional[torch.Tensor],
|
|
kv_cache_dtype: str = "auto",
|
|
k_scale: float = 1.0,
|
|
v_scale: float = 1.0,
|
|
**kwargs,
|
|
) -> None:
|
|
"""Triton-based PagedAttention V2.
|
|
|
|
NOTE: The Phase 1 kernel's K/V gather from paged cache is a skeleton.
|
|
The paged cache layout (key_cache: [blocks, kv_heads, head_dim/x, block_size, x])
|
|
requires indirect memory access (gather via block_tables) which is complex
|
|
in Triton. The Phase 2 reduction kernel is complete.
|
|
|
|
Current status:
|
|
Phase 1: SKELETON — falls back to PyTorch partition loop
|
|
Phase 2: COMPLETE — Triton reduction kernel
|
|
|
|
When Phase 1 is complete, this will be a single-launch V2:
|
|
grid = (num_seqs, num_heads, max_num_partitions) for Phase 1
|
|
grid = (num_seqs, num_heads) for Phase 2
|
|
"""
|
|
num_seqs, num_heads, head_size = query.shape
|
|
max_num_parts = tmp_output.shape[2]
|
|
|
|
PARTITION_SIZE = 512
|
|
BLOCK_N = 64 # Must fit in SMEM: BLOCK_N * head_dim * 2B * 2 ≤ 48KB
|
|
|
|
# --- Phase 1: Use PyTorch for now (Triton K/V gather skeleton above) ---
|
|
# TODO: Complete the Triton Phase 1 kernel with proper paged K/V gather
|
|
from paged_attention_v2_pytorch import paged_attention_v2_pytorch
|
|
paged_attention_v2_pytorch(
|
|
output, exp_sums, max_logits, tmp_output,
|
|
query, key_cache, value_cache,
|
|
num_kv_heads, scale, block_tables, seq_lens,
|
|
block_size, max_seq_len, alibi_slopes,
|
|
kv_cache_dtype, k_scale, v_scale,
|
|
)
|
|
# Phase 1 writes tmp_output, exp_sums, max_logits
|
|
# Phase 2 below will re-reduce them (redundant but correct)
|
|
|
|
# --- Phase 2: Triton reduction kernel ---
|
|
# This replaces the Python einsum reduction with a single Triton launch
|
|
MAX_NUM_PARTS_CONST = triton.next_power_of_2(max_num_parts)
|
|
if MAX_NUM_PARTS_CONST > 1024:
|
|
MAX_NUM_PARTS_CONST = 1024 # Safety cap
|
|
|
|
grid_reduce = (num_seqs, num_heads)
|
|
_paged_attn_v2_reduce_kernel[grid_reduce](
|
|
output,
|
|
tmp_output, exp_sums, max_logits, seq_lens,
|
|
max_num_parts,
|
|
# output strides
|
|
output.stride(0), output.stride(1), output.stride(2),
|
|
# tmp_output strides
|
|
tmp_output.stride(0), tmp_output.stride(1), tmp_output.stride(2), tmp_output.stride(3),
|
|
# exp_sums strides (same layout as max_logits)
|
|
exp_sums.stride(0), exp_sums.stride(1), exp_sums.stride(2),
|
|
# Constants
|
|
PARTITION_SIZE=PARTITION_SIZE,
|
|
HEAD_DIM=head_size,
|
|
MAX_NUM_PARTS=MAX_NUM_PARTS_CONST,
|
|
)
|