gemini-2.5-pro_triton_dorbxs
gemini-2.5-pro · triton · Apache-2.0
Use it
Vendorable · source mirrored · Apache-2.0View source →
No package. Vendor the mirrored source: 249 lines, Apache-2.0, pinned at da91508.
main.py
curl "https://kernelindex.com/api/v1/implementations/flashinfer-gemini-2-5-pro-triton-dorbxs?include=source"interfacetriton
revisionda915083d4c7
symbolrun
pathmain.py
Compatibility
declared hardwareNVIDIA B200
architecturessm_100
dtypesbf16, fp32, int32
Benchmark evidence
No published measurement for this revision.
No evidence · How evidence levels are derived →
Source and license
sourcehttps://huggingface.co/datasets/flashinfer-ai/flashinfer-trace
commitda915083d4c7c5e61aa3005e3d17ae488e0fc71c
revision digestsha256:79e60611256991bae2f3af286e994b0729ac438273ff289ee29aaf1e51df6cba
license declaredApache-2.0
license concludedApache-2.0
authorsgemini-2.5-pro
imported2026-08-20
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
autotune
@triton.autotune(mma
s_update = tl.dot(k_ckv_tile, q_nope_tile_2d)num-warps = 4
triton.Config({'BLOCK_L': 32, 'BLOCK_DCKV': 64, 'BLOCK_DKPE': 64}, num_warps=4, num_stages=3),stages = 3
triton.Config({'BLOCK_L': 32, 'BLOCK_DCKV': 64, 'BLOCK_DKPE': 64}, num_warps=4, num_stages=3),Kernel source
main.py249 lines
import math
import torch
import triton
import triton.language as tl
@triton.autotune(
configs=[
triton.Config({'BLOCK_L': 32, 'BLOCK_DCKV': 64, 'BLOCK_DKPE': 64}, num_warps=4, num_stages=3),
triton.Config({'BLOCK_L': 64, 'BLOCK_DCKV': 64, 'BLOCK_DKPE': 64}, num_warps=4, num_stages=3),
triton.Config({'BLOCK_L': 32, 'BLOCK_DCKV': 128, 'BLOCK_DKPE': 64}, num_warps=4, num_stages=2),
triton.Config({'BLOCK_L': 64, 'BLOCK_DCKV': 128, 'BLOCK_DKPE': 64}, num_warps=8, num_stages=2),
triton.Config({'BLOCK_L': 128, 'BLOCK_DCKV': 128, 'BLOCK_DKPE': 64}, num_warps=8, num_stages=2),
triton.Config({'BLOCK_L': 32, 'BLOCK_DCKV': 256, 'BLOCK_DKPE': 64}, num_warps=8, num_stages=2),
],
key=['HEAD_DIM_CKV', 'HEAD_DIM_KPE'],
)
@triton.jit
def mla_paged_decode_h16_ckv512_kpe64_ps1_kernel(
# Pointers to tensors
q_nope_ptr, q_pe_ptr, ckv_cache_ptr, kpe_cache_ptr,
kv_indptr_ptr, kv_indices_ptr,
output_ptr, lse_ptr,
# Scalar inputs
sm_scale,
# Strides
q_nope_stride_bs, q_nope_stride_h,
q_pe_stride_bs, q_pe_stride_h,
ckv_cache_stride_n,
kpe_cache_stride_n,
output_stride_bs, output_stride_h,
lse_stride_bs, lse_stride_h,
# Compile-time constants
HEAD_DIM_CKV: tl.constexpr,
HEAD_DIM_KPE: tl.constexpr,
# Tuning parameters
BLOCK_L: tl.constexpr,
BLOCK_DCKV: tl.constexpr,
BLOCK_DKPE: tl.constexpr,
):
"""
Triton kernel for paged multi-level attention decode.
Each program instance computes one head for one batch element.
"""
# Grid computes (batch_size, num_qo_heads)
b_idx = tl.program_id(0)
h_idx = tl.program_id(1)
log2 = 1.4426950408889634 # 1.0 / math.log(2.0)
# 1. --- Get sequence length for this batch element ---
page_beg = tl.load(kv_indptr_ptr + b_idx)
page_end = tl.load(kv_indptr_ptr + b_idx + 1)
L_tokens = page_end - page_beg
# 2. --- Initialize pointers and accumulators ---
q_nope_ptr += b_idx * q_nope_stride_bs + h_idx * q_nope_stride_h
q_pe_ptr += b_idx * q_pe_stride_bs + h_idx * q_pe_stride_h
output_ptr += b_idx * output_stride_bs + h_idx * output_stride_h
lse_ptr += b_idx * lse_stride_bs + h_idx * lse_stride_h
m_i = -float('inf')
l_i = 0.0
NUM_CHUNKS = HEAD_DIM_CKV // BLOCK_DCKV
acc0 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
acc1 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
acc2 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
acc3 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
acc4 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
acc5 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
acc6 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
acc7 = tl.zeros([BLOCK_DCKV], dtype=tl.float32)
# 3. --- Handle empty sequences ---
if L_tokens <= 0:
out_dtype = output_ptr.dtype.element_ty
zero_chunk = tl.zeros([BLOCK_DCKV], dtype=tl.float32).to(out_dtype)
for i in range(NUM_CHUNKS):
tl.store(output_ptr + i * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), zero_chunk)
tl.store(lse_ptr, m_i)
return
# 4. --- Main loop over the KV sequence in blocks ---
for l_start in range(0, L_tokens, BLOCK_L):
l_offs = l_start + tl.arange(0, BLOCK_L)
l_mask = l_offs < L_tokens
indices = tl.load(kv_indices_ptr + page_beg + l_offs, mask=l_mask, other=0)
# --- Compute logits for the block ---
s_block = tl.zeros([BLOCK_L], dtype=tl.float32)
# Contribution from ckv
for d_start in range(0, HEAD_DIM_CKV, BLOCK_DCKV):
d_offs = d_start + tl.arange(0, BLOCK_DCKV)
q_nope_tile = tl.load(q_nope_ptr + d_offs)
k_ckv_ptrs = ckv_cache_ptr + indices[:, None] * ckv_cache_stride_n + d_offs[None, :]
k_ckv_tile = tl.load(k_ckv_ptrs, mask=l_mask[:, None], other=0.0)
# CORRECTNESS FIX: Reshape 1D q_nope_tile to 2D for tl.dot, then squeeze result
q_nope_tile_2d = tl.reshape(q_nope_tile, (BLOCK_DCKV, 1))
s_update = tl.dot(k_ckv_tile, q_nope_tile_2d)
s_block += tl.squeeze(s_update, axis=1)
# Contribution from kpe
for d_start in range(0, HEAD_DIM_KPE, BLOCK_DKPE):
d_offs = d_start + tl.arange(0, BLOCK_DKPE)
q_pe_tile = tl.load(q_pe_ptr + d_offs)
k_kpe_ptrs = kpe_cache_ptr + indices[:, None] * kpe_cache_stride_n + d_offs[None, :]
k_kpe_tile = tl.load(k_kpe_ptrs, mask=l_mask[:, None], other=0.0)
# CORRECTNESS FIX: Reshape 1D q_pe_tile to 2D for tl.dot, then squeeze result
q_pe_tile_2d = tl.reshape(q_pe_tile, (BLOCK_DKPE, 1))
s_update = tl.dot(k_kpe_tile, q_pe_tile_2d)
s_block += tl.squeeze(s_update, axis=1)
# --- Online softmax update ---
s_block = tl.where(l_mask, s_block * sm_scale, -float('inf'))
m_i_old = m_i
m_i = tl.maximum(m_i, tl.max(s_block, axis=0))
# NUMERICAL STABILITY: guard against nan from exp(-inf - (-inf))
s_block_shifted = s_block - m_i
s_block_shifted = tl.where(m_i == -float('inf'), -float('inf'), s_block_shifted)
p_block = tl.exp(s_block_shifted)
l_i_new = tl.sum(p_block, axis=0)
alpha = tl.exp(m_i_old - m_i)
# NUMERICAL STABILITY: if m_i_old == m_i, alpha should be 1.0. Handles -inf case.
alpha = tl.where(m_i_old == m_i, 1.0, alpha)
l_i = alpha * l_i + l_i_new
p_block = p_block.to(ckv_cache_ptr.dtype.element_ty)
# --- Update output accumulator ---
if NUM_CHUNKS == 8:
acc0 *= alpha; acc1 *= alpha; acc2 *= alpha; acc3 *= alpha
acc4 *= alpha; acc5 *= alpha; acc6 *= alpha; acc7 *= alpha
elif NUM_CHUNKS == 4:
acc0 *= alpha; acc1 *= alpha; acc2 *= alpha; acc3 *= alpha
elif NUM_CHUNKS == 2:
acc0 *= alpha; acc1 *= alpha
# Add contribution from the current block (p_block @ v_block)
for i in range(NUM_CHUNKS):
d_offs = i * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV)
v_ckv_ptrs = ckv_cache_ptr + indices[:, None] * ckv_cache_stride_n + d_offs[None, :]
v_ckv_tile = tl.load(v_ckv_ptrs, mask=l_mask[:, None], other=0.0)
# CORRECTNESS FIX: Reshape 1D p_block to 2D for tl.dot, then squeeze result
p_block_2d = tl.reshape(p_block, (1, BLOCK_L))
update_2d = tl.dot(p_block_2d, v_ckv_tile)
update = tl.squeeze(update_2d, axis=0)
if i == 0: acc0 += update
elif i == 1: acc1 += update
elif i == 2: acc2 += update
elif i == 3: acc3 += update
elif i == 4: acc4 += update
elif i == 5: acc5 += update
elif i == 6: acc6 += update
elif i == 7: acc7 += update
# 5. --- Finalize and store results ---
l_i_reciprocal = tl.where(l_i > 0.0, 1.0 / l_i, 0.0)
out_dtype = output_ptr.dtype.element_ty
if NUM_CHUNKS == 8:
tl.store(output_ptr + 0 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc0 * l_i_reciprocal).to(out_dtype))
tl.store(output_ptr + 1 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc1 * l_i_reciprocal).to(out_dtype))
tl.store(output_ptr + 2 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc2 * l_i_reciprocal).to(out_dtype))
tl.store(output_ptr + 3 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc3 * l_i_reciprocal).to(out_dtype))
tl.store(output_ptr + 4 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc4 * l_i_reciprocal).to(out_dtype))
tl.store(output_ptr + 5 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc5 * l_i_reciprocal).to(out_dtype))
tl.store(output_ptr + 6 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc6 * l_i_reciprocal).to(out_dtype))
tl.store(output_ptr + 7 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc7 * l_i_reciprocal).to(out_dtype))
elif NUM_CHUNKS == 4:
tl.store(output_ptr + 0 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc0 * l_i_reciprocal).to(out_dtype))
tl.store(output_ptr + 1 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc1 * l_i_reciprocal).to(out_dtype))
tl.store(output_ptr + 2 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc2 * l_i_reciprocal).to(out_dtype))
tl.store(output_ptr + 3 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc3 * l_i_reciprocal).to(out_dtype))
elif NUM_CHUNKS == 2:
tl.store(output_ptr + 0 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc0 * l_i_reciprocal).to(out_dtype))
tl.store(output_ptr + 1 * BLOCK_DCKV + tl.arange(0, BLOCK_DCKV), (acc1 * l_i_reciprocal).to(out_dtype))
final_lse = (m_i + tl.log(l_i)) * log2
# handle case where l_i is 0, which makes tl.log(l_i) -> -inf
final_lse = tl.where(l_i > 0.0, final_lse, -float('inf'))
tl.store(lse_ptr, final_lse)
def mla_paged_decode_h16_ckv512_kpe64_ps1(q_nope, q_pe, ckv_cache, kpe_cache, kv_indptr, kv_indices, sm_scale):
"""
Wrapper function for the Triton kernel.
Handles device management, grid computation, and kernel launch.
"""
# 1. --- Check inputs and constants ---
batch_size, num_qo_heads, head_dim_ckv = q_nope.shape
head_dim_kpe = q_pe.shape[-1]
assert num_qo_heads == 16, "num_qo_heads must be 16"
assert head_dim_ckv == 512, "head_dim_ckv must be 512"
assert head_dim_kpe == 64, "head_dim_kpe must be 64"
assert ckv_cache.shape[1] == 1, "page_size must be 1"
# 2. --- Device Management ---
input_device = q_nope.device
is_cpu = input_device.type == 'cpu'
if is_cpu:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is not available, but input tensors are on CPU.")
q_nope = q_nope.cuda()
q_pe = q_pe.cuda()
ckv_cache = ckv_cache.cuda()
kpe_cache = kpe_cache.cuda()
kv_indptr = kv_indptr.cuda()
kv_indices = kv_indices.cuda()
# 3. --- Prepare outputs and grid ---
output = torch.empty_like(q_nope)
lse = torch.empty((batch_size, num_qo_heads), dtype=torch.float32, device=q_nope.device)
ckv_cache_squeezed = ckv_cache.squeeze(1)
kpe_cache_squeezed = kpe_cache.squeeze(1)
grid = (batch_size, num_qo_heads)
# 4. --- Launch kernel ---
mla_paged_decode_h16_ckv512_kpe64_ps1_kernel[grid](
q_nope, q_pe, ckv_cache_squeezed, kpe_cache_squeezed,
kv_indptr, kv_indices,
output, lse,
sm_scale,
q_nope.stride(0), q_nope.stride(1),
q_pe.stride(0), q_pe.stride(1),
ckv_cache_squeezed.stride(0),
kpe_cache_squeezed.stride(0),
output.stride(0), output.stride(1),
lse.stride(0), lse.stride(1),
HEAD_DIM_CKV=head_dim_ckv,
HEAD_DIM_KPE=head_dim_kpe,
)
# 5. --- Restore device and return ---
if is_cpu:
output = output.to(input_device)
lse = lse.to(input_device)
return {"output": output, "lse": lse}
def run(*args, **kwargs):
"""
Public entry point. Handles both args and kwargs for flexibility.
"""
return mla_paged_decode_h16_ckv512_kpe64_ps1(*args, **kwargs)scrolls · 249 lines total
Source code from the importing source · Apache-2.0
No published measurement for this revision
JSON