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
212 changes: 209 additions & 3 deletions python/sglang/kernels/ops/mamba/mamba_state_scatter_triton.py
Original file line number Diff line number Diff line change
Expand Up @@ -146,6 +146,208 @@ def track_mamba_states_if_needed(
)


@triton.jit
def _fused_mamba_state_scatter_multi_kernel(
src_ptrs,
dst_ptrs,
dst_indices_ptr,
step_indices_ptr,
dst_indices_stride,
step_indices_stride,
LAYERS: tl.constexpr,
ELEMENTS: tl.constexpr,
BLOCKS: tl.constexpr,
ENTRY_LAYOUTS: tl.constexpr,
SRC_STRIDES: tl.constexpr,
DST_STRIDES: tl.constexpr,
SRC_SIZES: tl.constexpr,
DST_SIZES: tl.constexpr,
BLOCK_SIZE: tl.constexpr,
track_indices_ptr,
track_steps_ptr,
track_indices_stride,
track_steps_stride,
BS: tl.constexpr,
HAS_TRACK: tl.constexpr,
TRACK_BLOCK: tl.constexpr,
):
req = tl.program_id(0).to(tl.int64)
tile = tl.program_id(1).to(tl.int64)
active_dst = tl.load(dst_indices_ptr + req * dst_indices_stride).to(tl.int64)
active_step = tl.load(step_indices_ptr + req * step_indices_stride).to(tl.int64)
if HAS_TRACK:
track_dst = tl.load(track_indices_ptr + req * track_indices_stride).to(tl.int64)
track_step = tl.load(track_steps_ptr + req * track_steps_stride).to(tl.int64)
rows = tl.arange(0, TRACK_BLOCK)
all_track_dst = tl.load(
track_indices_ptr + rows * track_indices_stride, rows < BS, other=-1
).to(tl.int64)
all_track_step = tl.load(
track_steps_ptr + rows * track_steps_stride, rows < BS, other=-1
).to(tl.int64)
first_tile = 0
for i in tl.static_range(len(LAYERS)):
if (tile >= first_tile) & (tile < first_tile + LAYERS[i] * BLOCKS[i]):
overwritten = False
if HAS_TRACK:
valid_track = (
(rows < BS)
& (rows < SRC_SIZES[i][0])
& (all_track_step >= 0)
& (all_track_step < SRC_SIZES[i][1])
)
overwritten = (
tl.sum(
(valid_track & (all_track_dst == active_dst)).to(tl.int32), 0
)
> 0
)
for commit in tl.static_range(2 if HAS_TRACK else 1):
dst_idx = active_dst
step = active_step
enabled = not overwritten
if commit == 1:
dst_idx = track_dst
step = track_step
enabled = True
if (
enabled
& (dst_idx >= 0)
& (dst_idx < DST_SIZES[i])
& (req < SRC_SIZES[i][0])
& (step >= 0)
& (step < SRC_SIZES[i][1])
):
local_tile = tile - first_tile
local_layer = local_tile // BLOCKS[i]
block = local_tile % BLOCKS[i]
offsets = block * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE)
if ENTRY_LAYOUTS[i][0] > 0:
source_offsets = (
offsets // ENTRY_LAYOUTS[i][0] * ENTRY_LAYOUTS[i][1]
+ offsets % ENTRY_LAYOUTS[i][0] * ENTRY_LAYOUTS[i][2]
)
else:
source_offsets = offsets
src_offset = (
local_layer * SRC_STRIDES[i][0]
+ req * SRC_STRIDES[i][1]
+ step * SRC_STRIDES[i][2]
)
dst_offset = (
local_layer * DST_STRIDES[i][0] + dst_idx * DST_STRIDES[i][1]
)
values = tl.load(
src_ptrs[i] + src_offset + source_offsets,
mask=offsets < ELEMENTS[i],
)
tl.store(
dst_ptrs[i] + dst_offset + offsets,
values,
mask=offsets < ELEMENTS[i],
)
first_tile += LAYERS[i] * BLOCKS[i]


def prepare_mamba_state_scatter_multi(state_pairs):
device = state_pairs[0][0].device
for dst, src in state_pairs:
if dst.device != device or src.device != device:
raise ValueError("states and indices must be on the same CUDA device")
if dst.ndim < 2 or src.ndim < 3:
raise ValueError("unexpected state ranks")
if dst.shape[0] != src.shape[0] or dst.shape[2:] != src.shape[3:]:
raise ValueError("state layer and trailing dimensions must match")
_require_entry_contiguous_dst(dst, 2, "fused_mamba_state_scatter_multi")
if not src.is_contiguous() and src.ndim != 5:
raise ValueError("src entries must be contiguous or two-dimensional")
layers = tuple(dst.shape[0] for dst, _ in state_pairs)
elements = tuple(
dst.numel() // (dst.shape[0] * dst.shape[1]) for dst, _ in state_pairs
)
blocks = tuple(triton.cdiv(n, 1024) for n in elements)
entry_layouts = tuple(
(
(0, 0, 0)
if src.is_contiguous()
else (src.shape[-1], src.stride(-2), src.stride(-1))
)
for _, src in state_pairs
)
return (
device,
(
tuple(src for _, src in state_pairs),
tuple(dst for dst, _ in state_pairs),
),
(
layers,
elements,
blocks,
entry_layouts,
tuple(src.stride()[:3] for _, src in state_pairs),
tuple(dst.stride()[:2] for dst, _ in state_pairs),
tuple(src.shape[1:3] for _, src in state_pairs),
tuple(dst.shape[1] for dst, _ in state_pairs),
),
sum(n * b for n, b in zip(layers, blocks)),
)


def fused_mamba_state_scatter_multi(
state_pairs,
dst_indices,
step_indices,
track_indices=None,
track_steps=None,
*,
_metadata=None,
):
if not state_pairs or step_indices.numel() == 0:
return
if dst_indices.ndim != 1 or step_indices.ndim != 1:
raise ValueError("indices must be 1D")
if dst_indices.shape != step_indices.shape:
raise ValueError("indices must have matching shapes")
if (track_indices is None) != (track_steps is None):
raise ValueError("track indices and steps must be supplied together")
indices_to_check = (dst_indices, step_indices)
if track_indices is not None:
if (
track_indices.shape != step_indices.shape
or track_steps.shape != step_indices.shape
):
raise ValueError("track indices must have matching shapes")
indices_to_check += (track_indices, track_steps)
for indices in indices_to_check:
if indices.dtype not in (torch.int32, torch.int64):
raise ValueError("indices must have int32 or int64 dtype")
if not indices.is_cuda or indices.device != step_indices.device:
raise ValueError("indices must be on the same CUDA device")
metadata = _metadata or prepare_mamba_state_scatter_multi(state_pairs)
device, pointers, constants, tiles = metadata
if device != step_indices.device:
raise ValueError("states and indices must be on the same CUDA device")
_fused_mamba_state_scatter_multi_kernel[(step_indices.numel(), tiles)](
*pointers,
dst_indices,
step_indices,
dst_indices.stride(0),
step_indices.stride(0),
*constants,
BLOCK_SIZE=1024,
track_indices_ptr=track_indices,
track_steps_ptr=track_steps,
track_indices_stride=(
track_indices.stride(0) if track_indices is not None else 0
),
track_steps_stride=track_steps.stride(0) if track_steps is not None else 0,
BS=step_indices.numel(),
HAS_TRACK=track_indices is not None,
TRACK_BLOCK=triton.next_power_of_2(step_indices.numel()),
)


@triton.jit
def _fused_mamba_state_scatter_with_mask_kernel(
src_ptr,
Expand Down Expand Up @@ -288,9 +490,13 @@ def fused_mamba_state_scatter_with_mask(
dst_layer_stride = dst.stride(0)
dst_req_stride = dst.stride(1)

# Ensure indices are int32 and contiguous
dst_indices_raw = dst_indices_raw.to(torch.int32).contiguous()
step_indices_raw = step_indices_raw.to(torch.int32).contiguous()
# Ensure index buffers are contiguous.
if dst_indices_raw.dtype not in (torch.int32, torch.int64):
dst_indices_raw = dst_indices_raw.to(torch.int32)
if step_indices_raw.dtype not in (torch.int32, torch.int64):
step_indices_raw = step_indices_raw.to(torch.int32)
dst_indices_raw = dst_indices_raw.contiguous()
step_indices_raw = step_indices_raw.contiguous()

_require_entry_contiguous_dst(dst, 2, "fused_mamba_state_scatter_with_mask")
if not src.is_contiguous():
Expand Down
97 changes: 86 additions & 11 deletions python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,9 @@
)
from sglang.kernels.ops.mamba.mamba_state_scatter_triton import (
fused_conv_window_scatter_with_mask,
fused_mamba_state_scatter_multi,
fused_mamba_state_scatter_with_mask,
prepare_mamba_state_scatter_multi,
scatter_mamba_states_after_mtp_verify,
track_mamba_states_all_layers,
track_mamba_states_if_needed,
Expand Down Expand Up @@ -1498,6 +1500,19 @@ def update_mamba_state_after_mtp_verify(
)
return

prepared = self._prepare_verify_state_scatter(mamba_caches)
if prepared is not None:
state_pairs, metadata = prepared
fused_mamba_state_scatter_multi(
state_pairs,
state_indices_tensor,
last_correct_step_indices,
mamba_track_indices,
mamba_steps_to_track,
_metadata=metadata,
)
return

scatter_mamba_states_after_mtp_verify(
mamba_caches,
state_indices_tensor,
Expand Down Expand Up @@ -1546,18 +1561,50 @@ def _scatter_speculative_state_with_mask(
:, src_indices[valid_indices], steps[valid_indices]
]

def _update_ple_state_after_mtp_verify(
self,
state_indices_tensor: torch.Tensor,
last_correct_step_indices: torch.Tensor,
mamba_track_indices: Optional[torch.Tensor],
mamba_steps_to_track: Optional[torch.Tensor],
):
"""Roll the accepted per-step PLE side states into their main slots."""
req_to_token_pool = self.linear_attn_backend.req_to_token_pool
if mamba_track_indices is not None:
assert mamba_steps_to_track is not None
def _prepare_verify_state_scatter(self, mamba_caches):
pool = self.linear_attn_backend.req_to_token_pool
tensors = (
mamba_caches.temporal,
mamba_caches.intermediate_ssm,
*mamba_caches.conv,
*mamba_caches.intermediate_conv_window,
pool.short_conv_pool.conv_state,
pool.short_conv_pool.intermediate_conv_state,
pool.ngram_pool.context,
pool.ngram_pool.intermediate_context,
)
key = tuple(
(
(id(t), t.data_ptr(), t.shape, t.stride(), t.device, t.dtype)
if t is not None
else None
)
for t in tensors
)
if getattr(self, "_verify_scatter_key", None) == key:
return self._verify_scatter_prepared
state_pairs = self._ple_state_pairs()
prepared = None
if state_pairs:
state_pairs = (
list(zip(mamba_caches.conv, mamba_caches.intermediate_conv_window))
+ state_pairs
)
if mamba_caches.temporal.numel() > 0:
state_pairs.insert(
0, (mamba_caches.temporal, mamba_caches.intermediate_ssm)
)
if all(
dst.is_cuda and src.is_cuda and (src.is_contiguous() or src.ndim == 5)
for dst, src in state_pairs
):
prepared = (state_pairs, prepare_mamba_state_scatter_multi(state_pairs))
self._verify_scatter_key = key
self._verify_scatter_prepared = prepared
return prepared

def _ple_state_pairs(self):
req_to_token_pool = self.linear_attn_backend.req_to_token_pool
state_pairs = []
short_conv_pool = req_to_token_pool.short_conv_pool
if (
Expand All @@ -1583,6 +1630,34 @@ def _update_ple_state_after_mtp_verify(
)
)

return state_pairs

def _update_ple_state_after_mtp_verify(
self,
state_indices_tensor: torch.Tensor,
last_correct_step_indices: torch.Tensor,
mamba_track_indices: Optional[torch.Tensor],
mamba_steps_to_track: Optional[torch.Tensor],
):
"""Roll the accepted per-step PLE side states into their main slots."""
if mamba_track_indices is not None:
assert mamba_steps_to_track is not None

state_pairs = self._ple_state_pairs()

if state_pairs and all(
state.is_cuda and intermediate.is_cuda
for state, intermediate in state_pairs
):
fused_mamba_state_scatter_multi(
state_pairs, state_indices_tensor, last_correct_step_indices
)
if mamba_track_indices is not None:
fused_mamba_state_scatter_multi(
state_pairs, mamba_track_indices, mamba_steps_to_track
)
return

for state, intermediate_state in state_pairs:
self._scatter_speculative_state_with_mask(
state,
Expand Down
Loading
Loading