[codex] Centralize more inline Triton kernels (#27429)

This commit is contained in:
Xiaoyu Zhang
2026-06-11 22:17:26 +08:00
committed by GitHub
parent d9110d971e
commit d571e076fa
8 changed files with 448 additions and 407 deletions
@@ -2,8 +2,6 @@ import logging
from typing import Optional, Union from typing import Optional, Union
import torch import torch
import triton
import triton.language as tl
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.attention.mamba.causal_conv1d_triton import PAD_SLOT_ID from sglang.srt.layers.attention.mamba.causal_conv1d_triton import PAD_SLOT_ID
@@ -14,6 +12,7 @@ from sglang.srt.layers.attention.mamba.mamba2_metadata import (
) )
from sglang.srt.layers.attention.mamba.mamba_state_scatter_triton import ( from sglang.srt.layers.attention.mamba.mamba_state_scatter_triton import (
fused_mamba_state_scatter_with_mask, fused_mamba_state_scatter_with_mask,
track_mamba_states_if_needed,
) )
from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool
@@ -26,108 +25,6 @@ from sglang.srt.speculative.spec_info import SpecInput
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Kernel to track mamba states if needed based on track mask
@triton.jit
def track_mamba_state_if_needed_kernel(
conv_states_ptr,
ssm_states_ptr,
cache_indices_ptr,
mamba_track_mask_ptr,
mamba_track_indices_ptr,
conv_state_stride_0, # stride for first dimension (batch/pool index)
ssm_state_stride_0, # stride for first dimension (batch/pool index)
conv_state_numel_per_row: tl.constexpr, # total elements per row
ssm_state_numel_per_row: tl.constexpr, # total elements per row
BLOCK_SIZE: tl.constexpr,
):
"""
Track conv_states and ssm_states rows based on track mask.
This kernel replaces a Python loop that copies state tensors for mamba attention.
For each batch element, if the track mask is True, it copies the entire row from
the source index (cache_indices[i]) to the destination index (mamba_track_indices[i]).
Grid: (batch_size,)
Each block handles one batch element, using multiple threads to copy data in parallel.
"""
batch_idx = tl.program_id(0)
# Load the copy mask for this batch element
track_mask = tl.load(mamba_track_mask_ptr + batch_idx)
# Early exit if we don't need to track
if not track_mask:
return
# Load source and destination indices
src_idx = tl.load(cache_indices_ptr + batch_idx)
dst_idx = tl.load(mamba_track_indices_ptr + batch_idx)
# Copy conv_states
# Each thread handles BLOCK_SIZE elements
for offset in range(0, conv_state_numel_per_row, BLOCK_SIZE):
element_indices = offset + tl.arange(0, BLOCK_SIZE)
mask = element_indices < conv_state_numel_per_row
src_ptr = conv_states_ptr + src_idx * conv_state_stride_0 + element_indices
dst_ptr = conv_states_ptr + dst_idx * conv_state_stride_0 + element_indices
data = tl.load(src_ptr, mask=mask, other=0.0)
tl.store(dst_ptr, data, mask=mask)
# Copy ssm_states
for offset in range(0, ssm_state_numel_per_row, BLOCK_SIZE):
element_indices = offset + tl.arange(0, BLOCK_SIZE)
mask = element_indices < ssm_state_numel_per_row
src_ptr = ssm_states_ptr + src_idx * ssm_state_stride_0 + element_indices
dst_ptr = ssm_states_ptr + dst_idx * ssm_state_stride_0 + element_indices
data = tl.load(src_ptr, mask=mask, other=0.0)
tl.store(dst_ptr, data, mask=mask)
def track_mamba_states_if_needed(
conv_states: torch.Tensor,
ssm_states: torch.Tensor,
cache_indices: torch.Tensor,
mamba_track_mask: torch.Tensor,
mamba_track_indices: torch.Tensor,
batch_size: int,
):
"""
Track mamba states using Triton kernel for better performance.
Args:
conv_states: Convolution states tensor [pool_size, ...]
ssm_states: SSM states tensor [pool_size, ...]
cache_indices: Source indices for each batch element [batch_size]
mamba_track_mask: Boolean mask indicating which elements to track [batch_size]
mamba_track_indices: Indices to track for each batch element [batch_size]
batch_size: Number of batch elements
"""
conv_state_numel_per_row = conv_states[0].numel()
ssm_state_numel_per_row = ssm_states[0].numel()
# Choose BLOCK_SIZE based on the size of the data
BLOCK_SIZE = 1024
# Launch kernel with batch_size blocks
grid = (batch_size,)
track_mamba_state_if_needed_kernel[grid](
conv_states,
ssm_states,
cache_indices,
mamba_track_mask,
mamba_track_indices,
conv_states.stride(0),
ssm_states.stride(0),
conv_state_numel_per_row,
ssm_state_numel_per_row,
BLOCK_SIZE,
)
class MambaAttnBackendBase(AttentionBackend): class MambaAttnBackendBase(AttentionBackend):
def __init__(self, model_runner: ModelRunner): def __init__(self, model_runner: ModelRunner):
super().__init__() super().__init__()
@@ -11,6 +11,107 @@ import triton
import triton.language as tl import triton.language as tl
@triton.jit
def track_mamba_state_if_needed_kernel(
conv_states_ptr,
ssm_states_ptr,
cache_indices_ptr,
mamba_track_mask_ptr,
mamba_track_indices_ptr,
conv_state_stride_0, # stride for first dimension (batch/pool index)
ssm_state_stride_0, # stride for first dimension (batch/pool index)
conv_state_numel_per_row: tl.constexpr, # total elements per row
ssm_state_numel_per_row: tl.constexpr, # total elements per row
BLOCK_SIZE: tl.constexpr,
):
"""
Track conv_states and ssm_states rows based on track mask.
This kernel replaces a Python loop that copies state tensors for mamba attention.
For each batch element, if the track mask is True, it copies the entire row from
the source index (cache_indices[i]) to the destination index (mamba_track_indices[i]).
Grid: (batch_size,)
Each block handles one batch element, using multiple threads to copy data in parallel.
"""
batch_idx = tl.program_id(0)
# Load the copy mask for this batch element
track_mask = tl.load(mamba_track_mask_ptr + batch_idx)
# Early exit if we don't need to track
if not track_mask:
return
# Load source and destination indices
src_idx = tl.load(cache_indices_ptr + batch_idx)
dst_idx = tl.load(mamba_track_indices_ptr + batch_idx)
# Copy conv_states
# Each thread handles BLOCK_SIZE elements
for offset in range(0, conv_state_numel_per_row, BLOCK_SIZE):
element_indices = offset + tl.arange(0, BLOCK_SIZE)
mask = element_indices < conv_state_numel_per_row
src_ptr = conv_states_ptr + src_idx * conv_state_stride_0 + element_indices
dst_ptr = conv_states_ptr + dst_idx * conv_state_stride_0 + element_indices
data = tl.load(src_ptr, mask=mask, other=0.0)
tl.store(dst_ptr, data, mask=mask)
# Copy ssm_states
for offset in range(0, ssm_state_numel_per_row, BLOCK_SIZE):
element_indices = offset + tl.arange(0, BLOCK_SIZE)
mask = element_indices < ssm_state_numel_per_row
src_ptr = ssm_states_ptr + src_idx * ssm_state_stride_0 + element_indices
dst_ptr = ssm_states_ptr + dst_idx * ssm_state_stride_0 + element_indices
data = tl.load(src_ptr, mask=mask, other=0.0)
tl.store(dst_ptr, data, mask=mask)
def track_mamba_states_if_needed(
conv_states: torch.Tensor,
ssm_states: torch.Tensor,
cache_indices: torch.Tensor,
mamba_track_mask: torch.Tensor,
mamba_track_indices: torch.Tensor,
batch_size: int,
):
"""
Track mamba states using Triton kernel for better performance.
Args:
conv_states: Convolution states tensor [pool_size, ...]
ssm_states: SSM states tensor [pool_size, ...]
cache_indices: Source indices for each batch element [batch_size]
mamba_track_mask: Boolean mask indicating which elements to track [batch_size]
mamba_track_indices: Indices to track for each batch element [batch_size]
batch_size: Number of batch elements
"""
conv_state_numel_per_row = conv_states[0].numel()
ssm_state_numel_per_row = ssm_states[0].numel()
# Choose BLOCK_SIZE based on the size of the data
BLOCK_SIZE = 1024
# Launch kernel with batch_size blocks
grid = (batch_size,)
track_mamba_state_if_needed_kernel[grid](
conv_states,
ssm_states,
cache_indices,
mamba_track_mask,
mamba_track_indices,
conv_states.stride(0),
ssm_states.stride(0),
conv_state_numel_per_row,
ssm_state_numel_per_row,
BLOCK_SIZE,
)
@triton.jit @triton.jit
def _fused_mamba_state_scatter_with_mask_kernel( def _fused_mamba_state_scatter_with_mask_kernel(
src_ptr, src_ptr,
@@ -44,6 +44,43 @@ def create_flashinfer_kv_indices_triton(
tl.store(kv_indices_ptr + kv_indices_offset + offset, data, mask=mask) tl.store(kv_indices_ptr + kv_indices_offset + offset, data, mask=mask)
@triton.jit
def create_chunked_prefix_cache_kv_indices(
req_to_token_ptr, # (max_batch, max_context_len,)
req_pool_indices_ptr, # (batch_size,)
chunk_start_idx_ptr, # (batch_size,)
chunk_seq_lens_ptr, # (batch_size,)
chunk_cu_seq_lens_ptr, # (batch_size + 1,)
chunk_kv_indices_ptr, # (num_chunk_tokens,)
req_to_token_ptr_stride: tl.constexpr,
):
BLOCK_SIZE: tl.constexpr = 512
pid = tl.program_id(axis=0)
# find the req pool idx, this is for batch to token
req_pool_index = tl.load(req_pool_indices_ptr + pid)
chunk_kv_indices_offset = tl.load(chunk_cu_seq_lens_ptr + pid)
# get the token positions of current chunk
chunk_start_pos = tl.load(chunk_start_idx_ptr + pid).to(tl.int32)
chunk_seq_len = tl.load(chunk_seq_lens_ptr + pid).to(tl.int32)
num_loop = tl.cdiv(chunk_seq_len, BLOCK_SIZE)
for i in range(num_loop):
offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
mask = offset < chunk_seq_len
data = tl.load(
req_to_token_ptr
+ req_pool_index * req_to_token_ptr_stride
+ chunk_start_pos
+ offset,
mask=mask,
)
tl.store(
chunk_kv_indices_ptr + chunk_kv_indices_offset + offset, data, mask=mask
)
def get_num_page_per_block_flashmla(page_size: int = 64) -> int: def get_num_page_per_block_flashmla(page_size: int = 64) -> int:
num_page_per_block = _FLASHMLA_CREATE_KV_BLOCK_SIZE // page_size num_page_per_block = _FLASHMLA_CREATE_KV_BLOCK_SIZE // page_size
return num_page_per_block return num_page_per_block
@@ -101,6 +101,229 @@ def pad_sequence_with_mask(
return B, output, attn_mask return B, output, attn_mask
@triton.jit
def pad_draft_extend_query_kernel(
q_ptr, # Input query tensor [total_seq_len, num_heads, head_dim]
padded_q_ptr, # Output padded query tensor [batch_size, max_seq_len, num_heads, head_dim]
seq_lens_q_ptr, # Sequence lengths for each sequence [batch_size]
cumsum_ptr, # Cumulative sum of sequence lengths [batch_size + 1]
batch_size,
max_seq_len,
num_heads,
head_dim,
BLOCK_SIZE: tl.constexpr,
):
"""Triton kernel for padding draft extended query tensor with parallelized head and dim processing."""
# Use 3D program IDs: (batch_seq, head_block, dim_block)
batch_seq_pid = tl.program_id(0)
head_pid = tl.program_id(1)
dim_pid = tl.program_id(2)
batch_id = batch_seq_pid // max_seq_len
seq_pos = batch_seq_pid % max_seq_len
if batch_id >= batch_size:
return
# Load sequence length for this batch
seq_len = tl.load(seq_lens_q_ptr + batch_id)
if seq_pos >= seq_len:
return
# Load cumulative sum to get start position in input tensor
input_start = tl.load(cumsum_ptr + batch_id)
input_pos = input_start + seq_pos
# Calculate head and dim block ranges
head_start = head_pid * BLOCK_SIZE
head_end = tl.minimum(head_start + BLOCK_SIZE, num_heads)
head_mask = tl.arange(0, BLOCK_SIZE) < (head_end - head_start)
dim_start = dim_pid * BLOCK_SIZE
dim_end = tl.minimum(dim_start + BLOCK_SIZE, head_dim)
dim_mask = tl.arange(0, BLOCK_SIZE) < (dim_end - dim_start)
# Calculate input offset
input_offset = (
input_pos * num_heads * head_dim
+ (head_start + tl.arange(0, BLOCK_SIZE))[:, None] * head_dim
+ (dim_start + tl.arange(0, BLOCK_SIZE))[None, :]
)
# Load data
data = tl.load(
q_ptr + input_offset,
mask=head_mask[:, None] & dim_mask[None, :],
other=0.0,
)
# Calculate output offset
output_offset = (
batch_id * max_seq_len * num_heads * head_dim
+ seq_pos * num_heads * head_dim
+ (head_start + tl.arange(0, BLOCK_SIZE))[:, None] * head_dim
+ (dim_start + tl.arange(0, BLOCK_SIZE))[None, :]
)
# Store data
tl.store(
padded_q_ptr + output_offset,
data,
mask=head_mask[:, None] & dim_mask[None, :],
)
def pad_draft_extend_query(
q: torch.Tensor,
padded_q: torch.Tensor,
seq_lens_q: torch.Tensor,
cu_seqlens_q: torch.Tensor,
) -> torch.Tensor:
"""Pad draft extended query using Triton kernel."""
batch_size = cu_seqlens_q.shape[0] - 1
max_seq_len_q = padded_q.shape[1]
num_heads = padded_q.shape[2]
head_dim = padded_q.shape[3]
# Launch Triton kernel with 3D grid for parallelized head and dim processing
BLOCK_SIZE = 64
num_head_blocks = triton.cdiv(num_heads, BLOCK_SIZE)
num_dim_blocks = triton.cdiv(head_dim, BLOCK_SIZE)
grid = (batch_size * max_seq_len_q, num_head_blocks, num_dim_blocks)
pad_draft_extend_query_kernel[grid](
q_ptr=q,
padded_q_ptr=padded_q,
seq_lens_q_ptr=seq_lens_q,
cumsum_ptr=cu_seqlens_q,
batch_size=batch_size,
max_seq_len=max_seq_len_q,
num_heads=num_heads,
head_dim=head_dim,
BLOCK_SIZE=BLOCK_SIZE,
)
return padded_q
@triton.jit
def unpad_draft_extend_output_kernel(
raw_out_ptr, # Input raw output tensor (batch_size, token_per_batch, tp_q_head_num, v_head_dim)
output_ptr, # Output tensor (-1, tp_q_head_num, v_head_dim)
num_accept_tokens_ptr, # Accept lengths for each sequence [batch_size]
cumsum_ptr, # Cumulative sum of accept lengths [batch_size + 1]
batch_size,
token_per_batch,
tp_q_head_num,
v_head_dim,
BLOCK_SIZE: tl.constexpr,
):
"""Triton kernel for unpadding draft extended output tensor with parallelized head and dim processing."""
batch_seq_pid = tl.program_id(0)
head_pid = tl.program_id(1)
dim_pid = tl.program_id(2)
batch_id = batch_seq_pid // token_per_batch
seq_pos = batch_seq_pid % token_per_batch
if batch_id >= batch_size:
return
# Load accept length for this batch
accept_len = tl.load(num_accept_tokens_ptr + batch_id)
if seq_pos >= accept_len:
return
# Load cumulative sum to get start position in output tensor
output_start = tl.load(cumsum_ptr + batch_id)
output_pos = output_start + seq_pos
# Calculate head and dim block ranges
head_start = head_pid * BLOCK_SIZE
head_end = tl.minimum(head_start + BLOCK_SIZE, tp_q_head_num)
head_mask = tl.arange(0, BLOCK_SIZE) < (head_end - head_start)
dim_start = dim_pid * BLOCK_SIZE
dim_end = tl.minimum(dim_start + BLOCK_SIZE, v_head_dim)
dim_mask = tl.arange(0, BLOCK_SIZE) < (dim_end - dim_start)
# Calculate input offset: (batch_id, seq_pos, head_id, dim_id)
input_offset = (
batch_id * token_per_batch * tp_q_head_num * v_head_dim
+ seq_pos * tp_q_head_num * v_head_dim
+ (head_start + tl.arange(0, BLOCK_SIZE))[:, None] * v_head_dim
+ (dim_start + tl.arange(0, BLOCK_SIZE))[None, :]
)
# Load data
data = tl.load(
raw_out_ptr + input_offset,
mask=head_mask[:, None] & dim_mask[None, :],
other=0.0,
)
output_offset = (
output_pos * tp_q_head_num * v_head_dim
+ (head_start + tl.arange(0, BLOCK_SIZE))[:, None] * v_head_dim
+ (dim_start + tl.arange(0, BLOCK_SIZE))[None, :]
)
# Store data
tl.store(
output_ptr + output_offset,
data,
mask=head_mask[:, None] & dim_mask[None, :],
)
def unpad_draft_extend_output(
raw_out: torch.Tensor,
cu_seqlens_q: torch.Tensor,
seq_lens_q: torch.Tensor,
sum_seq_lens_q: int,
unpad_output_buffer: torch.Tensor | None = None,
) -> torch.Tensor:
"""Unpad draft extended output using Triton kernel."""
# raw_out: (batch_size, token_per_batch, layer.tp_q_head_num, layer.v_head_dim)
batch_size = seq_lens_q.shape[0]
token_per_batch = raw_out.shape[1] # max_seq_len
tp_q_head_num = raw_out.shape[2] # num_heads
v_head_dim = raw_out.shape[3] # head_dim
total_tokens = sum_seq_lens_q
# Check if we're in CUDA graph mode (buffers are pre-allocated)
if unpad_output_buffer is not None:
# Use pre-allocated buffer for CUDA graph compatibility
output = unpad_output_buffer[:total_tokens, :, :].to(dtype=raw_out.dtype)
else:
# Dynamic allocation for non-CUDA graph mode
output = torch.empty(
(total_tokens, tp_q_head_num, v_head_dim),
dtype=raw_out.dtype,
device=raw_out.device,
)
# Launch Triton kernel with 3D grid for parallelized head and dim processing
BLOCK_SIZE = 64
num_head_blocks = triton.cdiv(tp_q_head_num, BLOCK_SIZE)
num_dim_blocks = triton.cdiv(v_head_dim, BLOCK_SIZE)
grid = (batch_size * token_per_batch, num_head_blocks, num_dim_blocks)
unpad_draft_extend_output_kernel[grid](
raw_out_ptr=raw_out,
output_ptr=output,
num_accept_tokens_ptr=seq_lens_q,
cumsum_ptr=cu_seqlens_q,
batch_size=batch_size,
token_per_batch=token_per_batch,
tp_q_head_num=tp_q_head_num,
v_head_dim=v_head_dim,
BLOCK_SIZE=BLOCK_SIZE,
)
return output[:total_tokens, :, :]
@triton.jit @triton.jit
def seqlens_expand_kernel( def seqlens_expand_kernel(
extend_seq_lens_ptr, # [N] extend_seq_lens_ptr, # [N]
@@ -1,9 +1,9 @@
from __future__ import annotations
""" """
Support attention backend for TRTLLM MLA kernels from flashinfer. Support attention backend for TRTLLM MLA kernels from flashinfer.
""" """
from __future__ import annotations
import logging import logging
import math import math
from dataclasses import dataclass from dataclasses import dataclass
@@ -11,7 +11,6 @@ from typing import TYPE_CHECKING, Optional, Union
import torch import torch
import triton import triton
import triton.language as tl
from sglang.jit_kernel.fixup_zero_kv import fixup_zero_kv_rows from sglang.jit_kernel.fixup_zero_kv import fixup_zero_kv_rows
from sglang.srt.environ import envs from sglang.srt.environ import envs
@@ -19,11 +18,19 @@ from sglang.srt.layers.attention.flashinfer_mla_backend import (
FlashInferMLAAttnBackend, FlashInferMLAAttnBackend,
FlashInferMLAMultiStepDraftBackend, FlashInferMLAMultiStepDraftBackend,
) )
from sglang.srt.layers.attention.utils import ( from sglang.srt.layers.attention.triton_ops.kv_indices import (
concat_mla_absorb_q_general,
create_flashmla_kv_indices_triton, create_flashmla_kv_indices_triton,
get_num_kv_index_blocks_flashmla, get_num_kv_index_blocks_flashmla,
get_num_page_per_block_flashmla, get_num_page_per_block_flashmla,
)
from sglang.srt.layers.attention.triton_ops.pad import (
pad_draft_extend_query as pad_draft_extend_query_triton,
)
from sglang.srt.layers.attention.triton_ops.pad import (
unpad_draft_extend_output as unpad_draft_extend_output_triton,
)
from sglang.srt.layers.attention.utils import (
concat_mla_absorb_q_general,
mla_quantize_and_rope_for_fp8, mla_quantize_and_rope_for_fp8,
) )
from sglang.srt.layers.dp_attention import get_attention_tp_size from sglang.srt.layers.dp_attention import get_attention_tp_size
@@ -56,150 +63,6 @@ DEFAULT_WORKSPACE_SIZE_MB = 150 # Memory workspace size in MB
TRTLLM_BLOCK_CONSTRAINT = 128 TRTLLM_BLOCK_CONSTRAINT = 128
@triton.jit
def pad_draft_extend_query_kernel(
q_ptr, # Input query tensor [total_seq_len, num_heads, head_dim]
padded_q_ptr, # Output padded query tensor [batch_size, max_seq_len, num_heads, head_dim]
seq_lens_q_ptr, # Sequence lengths for each sequence [batch_size]
cumsum_ptr, # Cumulative sum of sequence lengths [batch_size + 1]
batch_size,
max_seq_len,
num_heads,
head_dim,
BLOCK_SIZE: tl.constexpr,
):
"""Triton kernel for padding draft extended query tensor with parallelized head and dim processing."""
# Use 3D program IDs: (batch_seq, head_block, dim_block)
batch_seq_pid = tl.program_id(0)
head_pid = tl.program_id(1)
dim_pid = tl.program_id(2)
batch_id = batch_seq_pid // max_seq_len
seq_pos = batch_seq_pid % max_seq_len
if batch_id >= batch_size:
return
# Load sequence length for this batch
seq_len = tl.load(seq_lens_q_ptr + batch_id)
if seq_pos >= seq_len:
return
# Load cumulative sum to get start position in input tensor
input_start = tl.load(cumsum_ptr + batch_id)
input_pos = input_start + seq_pos
# Calculate head and dim block ranges
head_start = head_pid * BLOCK_SIZE
head_end = tl.minimum(head_start + BLOCK_SIZE, num_heads)
head_mask = tl.arange(0, BLOCK_SIZE) < (head_end - head_start)
dim_start = dim_pid * BLOCK_SIZE
dim_end = tl.minimum(dim_start + BLOCK_SIZE, head_dim)
dim_mask = tl.arange(0, BLOCK_SIZE) < (dim_end - dim_start)
# Calculate input offset
input_offset = (
input_pos * num_heads * head_dim
+ (head_start + tl.arange(0, BLOCK_SIZE))[:, None] * head_dim
+ (dim_start + tl.arange(0, BLOCK_SIZE))[None, :]
)
# Load data
data = tl.load(
q_ptr + input_offset,
mask=head_mask[:, None] & dim_mask[None, :],
other=0.0,
)
# Calculate output offset
output_offset = (
batch_id * max_seq_len * num_heads * head_dim
+ seq_pos * num_heads * head_dim
+ (head_start + tl.arange(0, BLOCK_SIZE))[:, None] * head_dim
+ (dim_start + tl.arange(0, BLOCK_SIZE))[None, :]
)
# Store data
tl.store(
padded_q_ptr + output_offset,
data,
mask=head_mask[:, None] & dim_mask[None, :],
)
@triton.jit
def unpad_draft_extend_output_kernel(
raw_out_ptr, # Input raw output tensor (batch_size, token_per_batch, tp_q_head_num, v_head_dim)
output_ptr, # Output tensor (-1, tp_q_head_num, v_head_dim)
num_accept_tokens_ptr, # Accept lengths for each sequence [batch_size]
cumsum_ptr, # Cumulative sum of accept lengths [batch_size + 1]
batch_size,
token_per_batch,
tp_q_head_num,
v_head_dim,
BLOCK_SIZE: tl.constexpr,
):
"""Triton kernel for unpadding draft extended output tensor with parallelized head and dim processing."""
batch_seq_pid = tl.program_id(0)
head_pid = tl.program_id(1)
dim_pid = tl.program_id(2)
batch_id = batch_seq_pid // token_per_batch
seq_pos = batch_seq_pid % token_per_batch
if batch_id >= batch_size:
return
# Load accept length for this batch
accept_len = tl.load(num_accept_tokens_ptr + batch_id)
if seq_pos >= accept_len:
return
# Load cumulative sum to get start position in output tensor
output_start = tl.load(cumsum_ptr + batch_id)
output_pos = output_start + seq_pos
# Calculate head and dim block ranges
head_start = head_pid * BLOCK_SIZE
head_end = tl.minimum(head_start + BLOCK_SIZE, tp_q_head_num)
head_mask = tl.arange(0, BLOCK_SIZE) < (head_end - head_start)
dim_start = dim_pid * BLOCK_SIZE
dim_end = tl.minimum(dim_start + BLOCK_SIZE, v_head_dim)
dim_mask = tl.arange(0, BLOCK_SIZE) < (dim_end - dim_start)
# Calculate input offset: (batch_id, seq_pos, head_id, dim_id)
input_offset = (
batch_id * token_per_batch * tp_q_head_num * v_head_dim
+ seq_pos * tp_q_head_num * v_head_dim
+ (head_start + tl.arange(0, BLOCK_SIZE))[:, None] * v_head_dim
+ (dim_start + tl.arange(0, BLOCK_SIZE))[None, :]
)
# Load data
data = tl.load(
raw_out_ptr + input_offset,
mask=head_mask[:, None] & dim_mask[None, :],
other=0.0,
)
output_offset = (
output_pos * tp_q_head_num * v_head_dim
+ (head_start + tl.arange(0, BLOCK_SIZE))[:, None] * v_head_dim
+ (dim_start + tl.arange(0, BLOCK_SIZE))[None, :]
)
# Store data
tl.store(
output_ptr + output_offset,
data,
mask=head_mask[:, None] & dim_mask[None, :],
)
def _quantize_fp8_qkv(q, k, v, layer): def _quantize_fp8_qkv(q, k, v, layer):
q = q.to(torch.float8_e4m3fn) q = q.to(torch.float8_e4m3fn)
@@ -593,7 +456,6 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
bs = forward_batch.batch_size bs = forward_batch.batch_size
if in_capture: if in_capture:
num_tokens = forward_batch.positions.numel() num_tokens = forward_batch.positions.numel()
seq_lens_cpu = forward_batch.seq_lens.cpu()
self._init_cuda_graph_metadata( self._init_cuda_graph_metadata(
bs, bs,
num_tokens, num_tokens,
@@ -718,29 +580,12 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
cu_seqlens_q: torch.Tensor, cu_seqlens_q: torch.Tensor,
) -> torch.Tensor: ) -> torch.Tensor:
"""Pad draft extended query using Triton kernel.""" """Pad draft extended query using Triton kernel."""
batch_size = cu_seqlens_q.shape[0] - 1 return pad_draft_extend_query_triton(
max_seq_len_q = padded_q.shape[1] q,
num_heads = padded_q.shape[2] padded_q,
head_dim = padded_q.shape[3] seq_lens_q,
cu_seqlens_q,
# Launch Triton kernel with 3D grid for parallelized head and dim processing
BLOCK_SIZE = 64
num_head_blocks = triton.cdiv(num_heads, BLOCK_SIZE)
num_dim_blocks = triton.cdiv(head_dim, BLOCK_SIZE)
grid = (batch_size * max_seq_len_q, num_head_blocks, num_dim_blocks)
pad_draft_extend_query_kernel[grid](
q_ptr=q,
padded_q_ptr=padded_q,
seq_lens_q_ptr=seq_lens_q,
cumsum_ptr=cu_seqlens_q,
batch_size=batch_size,
max_seq_len=max_seq_len_q,
num_heads=num_heads,
head_dim=head_dim,
BLOCK_SIZE=BLOCK_SIZE,
) )
return padded_q
def unpad_draft_extend_output( def unpad_draft_extend_output(
self, self,
@@ -750,45 +595,13 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
sum_seq_lens_q: int, sum_seq_lens_q: int,
) -> torch.Tensor: ) -> torch.Tensor:
"""Unpad draft extended output using Triton kernel.""" """Unpad draft extended output using Triton kernel."""
# raw_out: (batch_size, token_per_batch, layer.tp_q_head_num, layer.v_head_dim) return unpad_draft_extend_output_triton(
batch_size = seq_lens_q.shape[0] raw_out,
token_per_batch = raw_out.shape[1] # max_seq_len cu_seqlens_q,
tp_q_head_num = raw_out.shape[2] # num_heads seq_lens_q,
v_head_dim = raw_out.shape[3] # head_dim sum_seq_lens_q,
total_tokens = sum_seq_lens_q self.unpad_output_buffer,
# Check if we're in CUDA graph mode (buffers are pre-allocated)
if self.unpad_output_buffer is not None:
# Use pre-allocated buffer for CUDA graph compatibility
output = self.unpad_output_buffer[:total_tokens, :, :].to(
dtype=raw_out.dtype
) )
else:
# Dynamic allocation for non-CUDA graph mode
output = torch.empty(
(total_tokens, tp_q_head_num, v_head_dim),
dtype=raw_out.dtype,
device=raw_out.device,
)
# Launch Triton kernel with 3D grid for parallelized head and dim processing
BLOCK_SIZE = 64
num_head_blocks = triton.cdiv(tp_q_head_num, BLOCK_SIZE)
num_dim_blocks = triton.cdiv(v_head_dim, BLOCK_SIZE)
grid = (batch_size * token_per_batch, num_head_blocks, num_dim_blocks)
unpad_draft_extend_output_kernel[grid](
raw_out_ptr=raw_out,
output_ptr=output,
num_accept_tokens_ptr=seq_lens_q,
cumsum_ptr=cu_seqlens_q,
batch_size=batch_size,
token_per_batch=token_per_batch,
tp_q_head_num=tp_q_head_num,
v_head_dim=v_head_dim,
BLOCK_SIZE=BLOCK_SIZE,
)
return output[:total_tokens, :, :]
def _compute_decode_bmm1_scale(self, layer: RadixAttention) -> float: def _compute_decode_bmm1_scale(self, layer: RadixAttention) -> float:
"""BMM1 scale q_scale * k_scale * softmax_scale. k_scale only """BMM1 scale q_scale * k_scale * softmax_scale. k_scale only
@@ -4,11 +4,12 @@
from typing import List, Optional from typing import List, Optional
import torch import torch
import triton
import triton.language as tl
from sglang.srt.environ import envs from sglang.srt.environ import envs
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton from sglang.srt.layers.attention.triton_ops.kv_indices import (
create_chunked_prefix_cache_kv_indices,
create_flashinfer_kv_indices_triton,
)
from sglang.srt.model_executor.forward_context import ( from sglang.srt.model_executor.forward_context import (
get_req_to_token_pool, get_req_to_token_pool,
get_token_to_kv_pool, get_token_to_kv_pool,
@@ -213,40 +214,3 @@ class ForwardBatchDeepSeekMHAMixin:
) )
self.mha_one_shot_kv_indices = kv_indices self.mha_one_shot_kv_indices = kv_indices
return kv_indices return kv_indices
@triton.jit
def create_chunked_prefix_cache_kv_indices(
req_to_token_ptr, # (max_batch, max_context_len,)
req_pool_indices_ptr, # (batch_size,)
chunk_start_idx_ptr, # (batch_size,)
chunk_seq_lens_ptr, # (batch_size,)
chunk_cu_seq_lens_ptr, # (batch_size + 1,)
chunk_kv_indices_ptr, # (num_chunk_tokens,)
req_to_token_ptr_stride: tl.constexpr,
):
BLOCK_SIZE: tl.constexpr = 512
pid = tl.program_id(axis=0)
# find the req pool idx, this is for batch to token
req_pool_index = tl.load(req_pool_indices_ptr + pid)
chunk_kv_indices_offset = tl.load(chunk_cu_seq_lens_ptr + pid)
# get the token positions of current chunk
chunk_start_pos = tl.load(chunk_start_idx_ptr + pid).to(tl.int32)
chunk_seq_len = tl.load(chunk_seq_lens_ptr + pid).to(tl.int32)
num_loop = tl.cdiv(chunk_seq_len, BLOCK_SIZE)
for i in range(num_loop):
offset = tl.arange(0, BLOCK_SIZE) + i * BLOCK_SIZE
mask = offset < chunk_seq_len
data = tl.load(
req_to_token_ptr
+ req_pool_index * req_to_token_ptr_stride
+ chunk_start_pos
+ offset,
mask=mask,
)
tl.store(
chunk_kv_indices_ptr + chunk_kv_indices_offset + offset, data, mask=mask
)
+3 -53
View File
@@ -18,8 +18,6 @@ from typing import (
import torch import torch
import torch.nn as nn import torch.nn as nn
import torch.nn.functional as F import torch.nn.functional as F
import triton
import triton.language as tl
import sglang.srt.models.deepseek_v2 as deepseek_v2 import sglang.srt.models.deepseek_v2 as deepseek_v2
from sglang.jit_kernel.dsv4 import ( from sglang.jit_kernel.dsv4 import (
@@ -114,6 +112,9 @@ from sglang.srt.models.deepseek_common.amd.deepseek_v4_fused_mhc import (
) )
from sglang.srt.models.deepseek_common.utils import _use_aiter_bpreshuffle_gfx95 from sglang.srt.models.deepseek_common.utils import _use_aiter_bpreshuffle_gfx95
from sglang.srt.models.deepseek_v2 import ParallelLMHead, _is_cuda, _is_hip, _is_npu from sglang.srt.models.deepseek_v2 import ParallelLMHead, _is_cuda, _is_hip, _is_npu
from sglang.srt.models.triton_ops.deepseek_v4 import (
rms_normalize_triton as rms_normalize_triton,
)
if not _is_hip: if not _is_hip:
from sglang.srt.layers.utils.cp_utils import ( from sglang.srt.layers.utils.cp_utils import (
@@ -258,57 +259,6 @@ bcg_deepseek_v4_attention_with_output = eager_on_graph(True)(
) )
@triton.jit
def _rms_normalize_kernel(
x_ptr,
weight_ptr,
eps,
stride_row,
dim,
BLOCK_SIZE: tl.constexpr,
HAS_WEIGHT: tl.constexpr,
):
pid = tl.program_id(0)
offs = tl.arange(0, BLOCK_SIZE)
mask = offs < dim
base = pid * stride_row
x = tl.load(x_ptr + base + offs, mask=mask, other=0.0).to(tl.float32)
mean_sq = tl.sum(x * x, axis=0) / dim
rms_inv = tl.rsqrt(mean_sq + eps)
out = x * rms_inv
if HAS_WEIGHT:
weight = tl.load(weight_ptr + offs, mask=mask, other=0.0)
out = out * weight
tl.store(x_ptr + base + offs, out, mask=mask)
def rms_normalize_triton(
x: torch.Tensor, eps: float, weight: torch.Tensor = None
) -> torch.Tensor:
dim = x.shape[-1]
x_flat = x.view(-1, dim)
num_rows = x_flat.shape[0]
BLOCK_SIZE = triton.next_power_of_2(dim)
grid = (num_rows,)
_rms_normalize_kernel[grid](
x_flat,
weight,
eps,
x_flat.stride(0),
dim,
BLOCK_SIZE=BLOCK_SIZE,
HAS_WEIGHT=(weight is not None),
)
return x
class MQALayer(nn.Module): class MQALayer(nn.Module):
def __init__( def __init__(
self, self,
@@ -0,0 +1,56 @@
from typing import Optional
import torch
import triton
import triton.language as tl
@triton.jit
def _rms_normalize_kernel(
x_ptr,
weight_ptr,
eps,
stride_row,
dim,
BLOCK_SIZE: tl.constexpr,
HAS_WEIGHT: tl.constexpr,
):
pid = tl.program_id(0)
offs = tl.arange(0, BLOCK_SIZE)
mask = offs < dim
base = pid * stride_row
x = tl.load(x_ptr + base + offs, mask=mask, other=0.0).to(tl.float32)
mean_sq = tl.sum(x * x, axis=0) / dim
rms_inv = tl.rsqrt(mean_sq + eps)
out = x * rms_inv
if HAS_WEIGHT:
weight = tl.load(weight_ptr + offs, mask=mask, other=0.0)
out = out * weight
tl.store(x_ptr + base + offs, out, mask=mask)
def rms_normalize_triton(
x: torch.Tensor, eps: float, weight: Optional[torch.Tensor] = None
) -> torch.Tensor:
dim = x.shape[-1]
x_flat = x.view(-1, dim)
num_rows = x_flat.shape[0]
BLOCK_SIZE = triton.next_power_of_2(dim)
grid = (num_rows,)
_rms_normalize_kernel[grid](
x_flat,
weight,
eps,
x_flat.stride(0),
dim,
BLOCK_SIZE=BLOCK_SIZE,
HAS_WEIGHT=(weight is not None),
)
return x