Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 30 additions & 0 deletions kernels/prefill_attention/conversion-notes.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,30 @@
# prefill_attention conversion notes

## Tensor-descriptor conversion

- Source: original.py → tensor_descriptor.py (kernel `_prefill_attention_kernel_td`)
- Raw pointer arithmetic (`off_q`/`off_k`/`off_v`/`off_o`, `k_ptrs`/`v_ptrs`
advanced per iteration) replaced with 3D `tl.make_tensor_descriptor` views over
the packed `(seq_len, heads, head_dim)` layout. Each descriptor is rebased at
`cur_batch_in_all_start_index * stride_*bs` so its coordinates run
`0..cur_batch_seq_len`; the batch start is a scalar folded into the base
pointer, not a per-row runtime index, so plain `desc.load`/`desc.store` suffice
(no gather/scatter needed).
- **Tail masks dropped:** the original's seq-len (`offs_m < seq_len`,
`pos_k < seq_len` on the load) and head-dim (`mask_d = offs_d < Lk`) masks on
Q/K/V loads and the O store are redundant — the descriptor `shape` carries both
boundaries and zero-fills OOB. Zero is the additive identity for the `tl.dot`
accumulation and for masked-out `qk` (overwritten by the causal `tl.where`), so
the drop is safe.
- **Compute masks kept:** causal / sliding-window / valid-position masking
(`tl.where(mask, qk * sm_scale, -1.0e8)`) is attention *semantics*, not a tail
fill, so it is preserved verbatim.
- K is loaded as `(BLOCK_N, BLOCK_DMODEL)` and transposed via `tl.trans` before
`tl.dot(q, k)`. Descriptors require the last dim contiguous, so K cannot be
loaded pre-transposed the way the original did via `strides=(1, stride_kbs)`.
- Last-dim (§4) is satisfied: the contiguous axis is `head_dim` (`BLOCK_DMODEL`),
≥ 16 bytes; the length-1 head axis is the middle dim.
- **Signature change:** added `num_q_heads` / `num_kv_heads` runtime args (needed
for the descriptor `shape`, which the original never materialized). The wrapper
passes them only when the target kernel declares them (`arg_names` check) and
registers the TMA allocator, so it still drives the original kernel unchanged.
167 changes: 0 additions & 167 deletions kernels/prefill_attention/kernel.ktir

This file was deleted.

Original file line number Diff line number Diff line change
Expand Up @@ -2,38 +2,30 @@
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
#
# Tensor-descriptor conversion of prefill attention _fwd_kernel.
# Original: vllm/v1/attention/ops/triton_prefill_attention.py
#
# Changes from original:
# - Q/K/V/O accesses use 3D tensor descriptors over the full
# (total_tokens, num_*_heads, head_dim) tensors so the descriptor
# base is the input pointer (16-byte aligned trivially).
# - K is loaded as (BLOCK_N, BLOCK_DMODEL) and transposed via tl.trans
# before tl.dot(q, k_t). Tensor descriptors require the last dim
# contiguous, so K cannot be loaded pre-transposed (as block pointers
# allowed via strides=(1, stride_kbs)).
# - tl.advance is replaced with offset arithmetic per loop iteration.
# - Causal and sequence-length masking are still applied post-load via
# tl.where, identical to the block-pointer version.

import math
import torch
# Original: kernels/prefill_attention/original.py
# Changes summarized in kernels/prefill_attention/conversion-notes.md.

import triton
import triton.language as tl

RCP_LN2 = 1.0 / math.log(2.0)


@triton.jit
def _fwd_kernel_block_ptr(
Q, K, V,
def _prefill_attention_kernel_td(
Q,
K,
V,
sm_scale,
B_Start_Loc, B_Seqlen,
B_Start_Loc,
B_Seqlen,
Out,
stride_qbs, stride_qh,
stride_kbs, stride_kh,
stride_vbs, stride_vh,
stride_obs, stride_oh,
stride_qbs,
stride_qh,
stride_kbs,
stride_kh,
stride_vbs,
stride_vh,
stride_obs,
stride_oh,
num_q_heads,
num_kv_heads,
kv_group_num: tl.constexpr,
Expand All @@ -45,6 +37,8 @@ def _fwd_kernel_block_ptr(
SLIDING_WINDOW_K: tl.constexpr,
Lk: tl.constexpr,
):
"""Flash-attention prefill over a packed (total_tokens, heads, head_dim)
layout, using tensor descriptors for all Q/K/V/O accesses."""
cur_batch = tl.program_id(0)
cur_head = tl.program_id(1)
start_m = tl.program_id(2)
Expand All @@ -58,11 +52,18 @@ def _fwd_kernel_block_ptr(
offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
offs_n = tl.arange(0, BLOCK_N)

# Rebase each descriptor at this batch's first token so descriptor
# coordinates run 0..cur_batch_seq_len. The batch start is a scalar offset
# folded into the base pointer, not a per-row runtime index, so plain
# desc.load / desc.store suffice (no gather).
q_base = Q + cur_batch_in_all_start_index * stride_qbs
k_base = K + cur_batch_in_all_start_index * stride_kbs
v_base = V + cur_batch_in_all_start_index * stride_vbs
o_base = Out + cur_batch_in_all_start_index * stride_obs

# 3D descriptors over (seq_len, heads, head_dim). The last (head_dim) axis
# carries BLOCK_DMODEL * dtype_bytes >= 16 bytes; the head axis selects one
# head via a length-1 tile.
q_desc = tl.make_tensor_descriptor(
q_base,
shape=[cur_batch_seq_len, num_q_heads, Lk],
Expand All @@ -89,48 +90,82 @@ def _fwd_kernel_block_ptr(
)

q_row0 = start_m * BLOCK_M
# OOB rows/head-dim lanes are zero-filled by the descriptor; the seq-len and
# head-dim tail masks the original applied here are redundant.
q = q_desc.load([q_row0, cur_head, 0]).reshape([BLOCK_M, BLOCK_DMODEL])

# initialize pointer to m and l
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float("inf")
l_i = tl.zeros([BLOCK_M], dtype=tl.float32)
acc = tl.zeros([BLOCK_M, BLOCK_DMODEL], dtype=tl.float32)

block_mask = tl.where(block_start_loc < cur_batch_seq_len, 1, 0)

# Calculate the end position for attention computation
end_n = cur_batch_seq_len
# Apply causal attention pruning
end_n = tl.minimum(end_n, (start_m + 1) * BLOCK_M) if IS_CAUSAL else end_n

start_n_limit = 0
end_n_limit = block_mask * end_n

for start_n in range(start_n_limit, end_n_limit, BLOCK_N):
pos_q = offs_m[:, None]
pos_k = start_n + offs_n[None, :]
# -- prepare attention mask ----
# These are *compute* masks (causal / sliding-window / valid-position),
# not tail masks. The descriptor handles OOB fill; these encode the
# attention semantics and are kept verbatim from the original.
pos_q = offs_m[:, None] # Query positions [BLOCK_M, 1]
pos_k = start_n + offs_n[None, :] # Key positions [1, BLOCK_N]

# Valid sequence mask
mask = pos_k < cur_batch_seq_len
# Causal mask
if IS_CAUSAL:
mask &= pos_q >= pos_k

k_row0 = start_n
k_tile = k_desc.load([k_row0, cur_kv_head, 0]).reshape([BLOCK_N, BLOCK_DMODEL])
k_t = tl.trans(k_tile)

qk = tl.dot(q, k_t)
# Bidirectional sliding window masks
sliding_mask_q = (
pos_q - pos_k <= SLIDING_WINDOW_Q if SLIDING_WINDOW_Q > 0 else None
)
sliding_mask_k = (
pos_k - pos_q <= SLIDING_WINDOW_K if SLIDING_WINDOW_K > 0 else None
)
if sliding_mask_q is not None:
mask &= sliding_mask_q
if sliding_mask_k is not None:
mask &= sliding_mask_k

start_n = tl.multiple_of(start_n, BLOCK_N)
# -- compute qk ----
# Descriptors require the last dim contiguous, so K is loaded as
# (BLOCK_N, BLOCK_DMODEL) and transposed before the dot (the original
# loaded it pre-transposed via strides).
k_tile = k_desc.load([start_n, cur_kv_head, 0]).reshape(
[BLOCK_N, BLOCK_DMODEL]
)
k = tl.trans(k_tile)

qk = tl.dot(q, k)
qk = tl.where(mask, qk * sm_scale, -1.0e8)
m_ij = tl.maximum(m_i, tl.max(qk, 1))
qk -= m_ij[:, None]
p = tl.math.exp2(qk)
l_ij = tl.sum(p, 1)

# -- update m_i and l_i
alpha = tl.math.exp2(m_i - m_ij)
l_i = l_i * alpha + l_ij
# -- update output accumulator --
acc = acc * alpha[:, None]

v = v_desc.load([k_row0, cur_kv_head, 0]).reshape([BLOCK_N, BLOCK_DMODEL])
# update acc
v = v_desc.load([start_n, cur_kv_head, 0]).reshape([BLOCK_N, BLOCK_DMODEL])
p = p.to(v.dtype)
acc = tl.dot(p, v, acc)
# update m_i
m_i = m_ij

acc = acc / l_i[:, None]
acc = acc.to(Out.dtype.element_ty)

# The store clamps at shape=[cur_batch_seq_len, ...], so the partial tail

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

would be good to try to lower it with the triton fork to see if these index comparisons and tl.where will be lowered properly.

# row tile never writes past the sequence — no output mask needed.
o_desc.store([q_row0, cur_head, 0], acc.reshape([BLOCK_M, 1, BLOCK_DMODEL]))
4 changes: 4 additions & 0 deletions kernels/prefill_attention/wrapper.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,10 @@ def context_attention_fwd(
grid = (batch, head, triton.cdiv(max_input_len, BLOCK))
num_warps = 4 if Lk <= 64 else 8

# The tensor-descriptor kernel needs the head counts to build its 3D
# descriptor shapes (the original derived head indexing purely from
# strides). Supply them only when the target kernel declares them, and
# register the TMA allocator that make_tensor_descriptor requires.
extra = {}
if "num_q_heads" in kernel_fn.arg_names:
ensure_triton_allocator()
Expand Down
Loading