init v0.23.0

Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
This commit is contained in:
2026-08-27 15:11:51 +08:00
parent b582a8e7d1
commit 7f8a1b1f7a
2849 changed files with 712887 additions and 22001 deletions

View File

View File

@@ -0,0 +1,588 @@
# adapted from vllm/model_executor/layers/mamba/ops/causal_conv1d.py
# Adapted from https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/mamba/ops/causal_conv1d.py
# SPDX-License-Identifier: Apache-2.0
# Copyright (c) 2024, Tri Dao.
# Adapted from https://github.com/Dao-AILab/causal-conv1d/blob/main/causal_conv1d/causal_conv1d_interface.py
# and https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/mamba/ops/causal_conv1d.py
# mypy: ignore-errors
import torch
from vllm.triton_utils import HAS_TRITON, tl, triton
from vllm.v1.attention.backends.utils import PAD_SLOT_ID # type: ignore
from vllm_ascend.ops.triton.triton_utils import get_vectorcore_num
if not HAS_TRITON:
from vllm_ascend._310p.ops.causal_conv1d import (
causal_conv1d_update as _pytorch_update,
)
else:
_pytorch_update = None
def extract_last_width(x, start_loc, width):
end_loc = start_loc[1:]
offsets = torch.arange(width, device=x.device)
indices = end_loc.unsqueeze(1) - width + offsets.unsqueeze(0) # (num_seqs, width)
return x[:, indices].permute(1, 0, 2)
@triton.jit(
do_not_specialize=[
"batch",
"state_len",
"num_cache_lines",
"stride_x_seq",
"stride_x_token",
"stride_conv_state_seq",
"stride_state_indices",
"stride_o_seq",
"stride_o_token",
]
)
def _causal_conv1d_update_kernel_npu_tiled(
# Pointers
x_ptr, # (batch, dim, seqlen) OR (num_tokens, dim) for varlen
w_ptr, # (dim, width)
bias_ptr,
conv_state_ptr, # (num_cache_lines, dim, state_len)
conv_state_indices_ptr,
num_accepted_tokens_ptr,
query_start_loc_ptr, # (batch + 1)
block_idx_last_scheduled_token, # (batch,)
initial_state_idx, # (batch,)
o_ptr, # same shape as x_ptr
batch: tl.int32,
dim: tl.constexpr,
seqlen: tl.constexpr, # max seqlen for varlen, or exact seqlen
state_len, # effective state_len computed in wrapper
num_cache_lines,
# Strides
stride_x_seq,
stride_x_dim: tl.constexpr,
stride_x_token,
stride_w_dim: tl.constexpr,
stride_w_width: tl.constexpr,
stride_conv_state_seq,
stride_conv_state_dim: tl.constexpr,
stride_conv_state_tok: tl.constexpr,
stride_state_indices,
stride_o_seq,
stride_o_dim: tl.constexpr,
stride_o_token,
# others
pad_slot_id: tl.constexpr,
# Meta
HAS_BIAS: tl.constexpr,
KERNEL_WIDTH: tl.constexpr, # <= 6
SILU_ACTIVATION: tl.constexpr,
IS_VARLEN: tl.constexpr,
IS_APC_ENABLED: tl.constexpr,
IS_SPEC_DECODING: tl.constexpr,
NP2_STATELEN: tl.constexpr,
USE_PAD_SLOT: tl.constexpr,
# tiling
BLOCK_N: tl.constexpr, # channel tile (C_TILE)
B_TILE: tl.constexpr, # batch tile
T_CHUNK: tl.constexpr, # token chunk for state update
):
# program ids
pid_b = tl.program_id(0) # batch-tile id
pid_c = tl.program_id(1) # channel-tile id
# channel indices for this program
idx_feats = pid_c * BLOCK_N + tl.arange(0, BLOCK_N) # [BLOCK_N]
mask_w = idx_feats < dim
# preload weights once per program (shared by B_TILE sequences)
w_base = w_ptr + idx_feats * stride_w_dim
# define to avoid "undefined" in branches
w_col0 = tl.zeros((BLOCK_N,), dtype=tl.float32)
w_col1 = tl.zeros((BLOCK_N,), dtype=tl.float32)
w_col2 = tl.zeros((BLOCK_N,), dtype=tl.float32)
w_col3 = tl.zeros((BLOCK_N,), dtype=tl.float32)
w_col4 = tl.zeros((BLOCK_N,), dtype=tl.float32)
w_col5 = tl.zeros((BLOCK_N,), dtype=tl.float32)
if KERNEL_WIDTH >= 1:
w_col0 = tl.load(w_base + 0 * stride_w_width, mask=mask_w, other=0.0).to(tl.float32)
if KERNEL_WIDTH >= 2:
w_col1 = tl.load(w_base + 1 * stride_w_width, mask=mask_w, other=0.0).to(tl.float32)
if KERNEL_WIDTH >= 3:
w_col2 = tl.load(w_base + 2 * stride_w_width, mask=mask_w, other=0.0).to(tl.float32)
if KERNEL_WIDTH >= 4:
w_col3 = tl.load(w_base + 3 * stride_w_width, mask=mask_w, other=0.0).to(tl.float32)
if KERNEL_WIDTH >= 5:
w_col4 = tl.load(w_base + 4 * stride_w_width, mask=mask_w, other=0.0).to(tl.float32)
if KERNEL_WIDTH >= 6:
w_col5 = tl.load(w_base + 5 * stride_w_width, mask=mask_w, other=0.0).to(tl.float32)
# bias vector once per program
if HAS_BIAS:
acc_bias = tl.load(bias_ptr + idx_feats, mask=mask_w, other=0.0).to(tl.float32)
else:
acc_bias = tl.zeros((BLOCK_N,), dtype=tl.float32)
# token index vector for chunked copy
tok_vec = tl.arange(0, T_CHUNK) # [T_CHUNK]
# process B_TILE sequences inside the same program instance
for bi in tl.static_range(0, B_TILE):
b = pid_b * B_TILE + bi # scalar tl.int32
lane_active = b < batch # scalar predicate
# -------------------------
# APC mapping (optional)
# -------------------------
if IS_APC_ENABLED:
conv_state_init = tl.load(initial_state_idx + b, mask=lane_active, other=0).to(tl.int32)
current_last_index = tl.load(block_idx_last_scheduled_token + b, mask=lane_active, other=0).to(tl.int32)
else:
conv_state_init = tl.full((), 0, tl.int32)
current_last_index = tl.full((), 0, tl.int32)
# input cache line
conv_states_input_coord = tl.load(
conv_state_indices_ptr + b * stride_state_indices + conv_state_init, mask=lane_active, other=0
).to(tl.int64)
if USE_PAD_SLOT:
lane_active = lane_active & (conv_states_input_coord != pad_slot_id)
# -------------------------
# varlen (optional): revise seqlen_run and state_len_run like original kernel does
# -------------------------
if IS_VARLEN:
qs = tl.load(query_start_loc_ptr + b, mask=lane_active, other=0).to(tl.int64)
qe = tl.load(query_start_loc_ptr + (b + 1), mask=lane_active, other=0).to(tl.int64)
seqlen_run = (qe - qs).to(tl.int32)
# revise effective state_len for shorter sequences (same formula as original)
state_len_run = (state_len - (seqlen - seqlen_run)).to(tl.int32)
x_offset = (qs * stride_x_token).to(tl.int64)
o_offset = (qs * stride_o_token).to(tl.int64)
else:
seqlen_run = tl.full((), seqlen, tl.int32)
state_len_run = tl.full((), state_len, tl.int32)
x_offset = (b * stride_x_seq).to(tl.int64)
o_offset = (b * stride_o_seq).to(tl.int64)
# empty sequence -> skip (avoid early return because other lanes in tile)
lane_active = lane_active & (seqlen_run > 0)
# -------------------------
# spec decoding offset (optional)
# -------------------------
if IS_SPEC_DECODING:
conv_state_token_offset = tl.load(num_accepted_tokens_ptr + b, mask=lane_active, other=1).to(tl.int64) - 1
shift = tl.full((), 1, tl.int32) # sliding by 1 in spec mode
else:
conv_state_token_offset = tl.full((), 0, tl.int64)
shift = seqlen_run # normal mode shift by seqlen
# -------------------------
# STEP 1: read initial history cols BEFORE state update (out==x safe)
# -------------------------
conv_states_base = (
conv_state_ptr + conv_states_input_coord * stride_conv_state_seq + idx_feats * stride_conv_state_dim
)
prior_tokens = conv_states_base + conv_state_token_offset * stride_conv_state_tok
# define history vectors as zeros then load conditionally
col0 = tl.zeros((BLOCK_N,), dtype=tl.float16)
col1 = tl.zeros((BLOCK_N,), dtype=tl.float16)
col2 = tl.zeros((BLOCK_N,), dtype=tl.float16)
col3 = tl.zeros((BLOCK_N,), dtype=tl.float16)
col4 = tl.zeros((BLOCK_N,), dtype=tl.float16)
if KERNEL_WIDTH >= 2:
col0 = tl.load(prior_tokens + 0 * stride_conv_state_tok, mask=lane_active & mask_w, other=0.0).to(
tl.float16
)
if KERNEL_WIDTH >= 3:
col1 = tl.load(prior_tokens + 1 * stride_conv_state_tok, mask=lane_active & mask_w, other=0.0).to(
tl.float16
)
if KERNEL_WIDTH >= 4:
col2 = tl.load(prior_tokens + 2 * stride_conv_state_tok, mask=lane_active & mask_w, other=0.0).to(
tl.float16
)
if KERNEL_WIDTH >= 5:
col3 = tl.load(prior_tokens + 3 * stride_conv_state_tok, mask=lane_active & mask_w, other=0.0).to(
tl.float16
)
if KERNEL_WIDTH >= 6:
col4 = tl.load(prior_tokens + 4 * stride_conv_state_tok, mask=lane_active & mask_w, other=0.0).to(
tl.float16
)
# -------------------------
# STEP 2: chunked state update (replaces original NP2_STATELEN x BLOCK_N big block)
# Semantics: conv_state <- concat(old_state, x)[-state_len_run:].
# - If seqlen_run >= state_len_run: dst[:] = x[seqlen_run - state_len_run : seqlen_run]
# - Else: keep = state_len_run - seqlen_run,
# dst[0:keep] = src[shift : shift+keep], dst[keep:keep+seqlen_run] = x[0:seqlen_run]
# -------------------------
# output cache line
conv_states_offset = tl.load(
conv_state_indices_ptr + b * stride_state_indices + current_last_index, mask=lane_active, other=0
).to(tl.int64)
use_shift = seqlen_run < state_len_run
use_tail = seqlen_run >= state_len_run
zero_i32 = tl.full((), 0, tl.int32)
keep_shift = tl.where(use_shift, (state_len_run - seqlen_run), zero_i32).to(tl.int32)
tail_start = tl.where(use_tail, (seqlen_run - state_len_run), zero_i32).to(tl.int32)
# base pointers
state_src_base = (
conv_state_ptr
+ conv_states_input_coord * stride_conv_state_seq
+ conv_state_token_offset * stride_conv_state_tok
+ idx_feats * stride_conv_state_dim
)
state_dst_base = conv_state_ptr + conv_states_offset * stride_conv_state_seq + idx_feats * stride_conv_state_dim
x_base = x_ptr + x_offset + idx_feats * stride_x_dim
# A) shift old state into dst[0:keep_shift) (only when seqlen_run < state_len_run)
for t0 in tl.static_range(0, NP2_STATELEN, T_CHUNK):
dst_tok = (t0 + tok_vec).to(tl.int32) # [T_CHUNK]
src_tok = (dst_tok + shift).to(tl.int32) # [T_CHUNK]
m_tok = use_shift & (dst_tok < keep_shift) & (src_tok < state_len_run) & (dst_tok < state_len_run)
m = (
(lane_active & m_tok)[:, None]
& mask_w[None, :]
& (conv_states_input_coord < num_cache_lines)
& (conv_states_offset < num_cache_lines)
)
src_ptrs = state_src_base[None, :] + src_tok[:, None] * stride_conv_state_tok
dst_ptrs = state_dst_base[None, :] + dst_tok[:, None] * stride_conv_state_tok
vals = tl.load(src_ptrs, mask=m, other=0.0)
tl.store(dst_ptrs, vals, mask=m)
# B) append x into dst[keep_shift : keep_shift+seqlen_run) (only when seqlen_run < state_len_run)
for t0 in tl.static_range(0, seqlen, T_CHUNK):
x_tok = (t0 + tok_vec).to(tl.int32) # [T_CHUNK]
dst_tok = (keep_shift + x_tok).to(tl.int32) # [T_CHUNK]
m_tok = use_shift & (x_tok < seqlen_run) & (dst_tok < state_len_run)
m = (lane_active & m_tok)[:, None] & mask_w[None, :] & (conv_states_offset < num_cache_lines)
x_ptrs = x_base[None, :] + x_tok[:, None] * stride_x_token
dst_ptrs = state_dst_base[None, :] + dst_tok[:, None] * stride_conv_state_tok
x_vals = tl.load(x_ptrs, mask=m, other=0.0)
tl.store(dst_ptrs, x_vals, mask=m)
# C) if seqlen_run >= state_len_run, overwrite dst with the tail of x
for t0 in tl.static_range(0, NP2_STATELEN, T_CHUNK):
dst_tok = (t0 + tok_vec).to(tl.int32) # [T_CHUNK]
x_tok = (tail_start + dst_tok).to(tl.int32) # [T_CHUNK]
m_tok = use_tail & (dst_tok < state_len_run) & (x_tok < seqlen_run)
m = (lane_active & m_tok)[:, None] & mask_w[None, :] & (conv_states_offset < num_cache_lines)
x_ptrs = x_base[None, :] + x_tok[:, None] * stride_x_token
dst_ptrs = state_dst_base[None, :] + dst_tok[:, None] * stride_conv_state_tok
x_vals = tl.load(x_ptrs, mask=m, other=0.0)
tl.store(dst_ptrs, x_vals, mask=m)
# -------------------------
# STEP 3/4/5: causal conv1d (+ optional SiLU) and store output
# This is original STEP3~5, but per-lane and without debug_barrier.
# -------------------------
x_base_1d = x_base
o_base_1d = o_ptr + o_offset + idx_feats * stride_o_dim
# accumulator preload (bias)
acc_preload = acc_bias
# compute each token; keep tl.range so varlen can use seqlen_run as runtime trip count (like original)
for idx_token in tl.range(seqlen_run):
acc = acc_preload
# same selection logic as original (unrolled by KERNEL_WIDTH)
matrix_w = w_col0
matrix_x = col0
for j in tl.static_range(KERNEL_WIDTH):
if KERNEL_WIDTH == 1:
# only x[t] * w0
x_ptrs_1d = x_base_1d + idx_token * stride_x_token
matrix_x = tl.load(x_ptrs_1d, mask=lane_active & mask_w, other=0.0).to(tl.float16)
matrix_w = w_col0
elif KERNEL_WIDTH == 2:
if j == 1:
matrix_w = w_col1
x_ptrs_1d = x_base_1d + idx_token * stride_x_token
matrix_x = tl.load(x_ptrs_1d, mask=lane_active & mask_w, other=0.0).to(tl.float16)
elif KERNEL_WIDTH == 3:
if j == 1:
matrix_w = w_col1
matrix_x = col1
elif j == 2:
matrix_w = w_col2
x_ptrs_1d = x_base_1d + idx_token * stride_x_token
matrix_x = tl.load(x_ptrs_1d, mask=lane_active & mask_w, other=0.0).to(tl.float16)
elif KERNEL_WIDTH == 4:
if j == 1:
matrix_w = w_col1
matrix_x = col1
elif j == 2:
matrix_w = w_col2
matrix_x = col2
elif j == 3:
matrix_w = w_col3
x_ptrs_1d = x_base_1d + idx_token * stride_x_token
matrix_x = tl.load(x_ptrs_1d, mask=lane_active & mask_w, other=0.0).to(tl.float16)
elif KERNEL_WIDTH == 5:
if j == 1:
matrix_w = w_col1
matrix_x = col1
elif j == 2:
matrix_w = w_col2
matrix_x = col2
elif j == 3:
matrix_w = w_col3
matrix_x = col3
elif j == 4:
matrix_w = w_col4
x_ptrs_1d = x_base_1d + idx_token * stride_x_token
matrix_x = tl.load(x_ptrs_1d, mask=lane_active & mask_w, other=0.0).to(tl.float16)
elif KERNEL_WIDTH == 6:
if j == 1:
matrix_w = w_col1
matrix_x = col1
elif j == 2:
matrix_w = w_col2
matrix_x = col2
elif j == 3:
matrix_w = w_col3
matrix_x = col3
elif j == 4:
matrix_w = w_col4
matrix_x = col4
elif j == 5:
matrix_w = w_col5
x_ptrs_1d = x_base_1d + idx_token * stride_x_token
matrix_x = tl.load(x_ptrs_1d, mask=lane_active & mask_w, other=0.0).to(tl.float16)
acc += matrix_x.to(tl.float32) * matrix_w # [BLOCK_N]
# roll history window
if KERNEL_WIDTH == 2:
col0 = matrix_x
elif KERNEL_WIDTH == 3:
col0 = col1
col1 = matrix_x
elif KERNEL_WIDTH == 4:
col0 = col1
col1 = col2
col2 = matrix_x
elif KERNEL_WIDTH == 5:
col0 = col1
col1 = col2
col2 = col3
col3 = matrix_x
elif KERNEL_WIDTH == 6:
col0 = col1
col1 = col2
col2 = col3
col3 = col4
col4 = matrix_x
if SILU_ACTIVATION:
acc = acc / (1.0 + tl.exp(-acc))
# store output
o_ptrs = o_base_1d + idx_token * stride_o_token
tl.store(o_ptrs, acc, mask=lane_active & mask_w)
def causal_conv1d_update_npu(
x: torch.Tensor,
conv_state: torch.Tensor,
weight: torch.Tensor,
bias: torch.Tensor | None = None,
activation: bool | str | None = None,
conv_state_indices: torch.Tensor | None = None,
num_accepted_tokens: torch.Tensor | None = None,
query_start_loc: torch.Tensor | None = None,
max_query_len: int = -1,
pad_slot_id: int = PAD_SLOT_ID,
block_idx_last_scheduled_token: torch.Tensor | None = None,
initial_state_idx: torch.Tensor | None = None,
validate_data=False,
):
"""
x: Input tensor which can take the following shapes:
- `[batch, dim]` - single token prediction
- `[batch, dim, seqlen]` - single or multiple tokens prediction
- `[num_tokens, dim]` - continuous batching, where num_tokens is
the total tokens of all sequences in that batch
conv_state: (..., dim, state_len), where state_len >= width - 1
weight: (dim, width)
bias: (dim,)
conv_state_indices: (batch,), dtype int32
If not None, the conv_state is a larger tensor along the batch dim,
and we are selecting the batch coords specified by conv_state_indices.
Useful for a continuous batching scenario.
block_idx_last_scheduled_token: (batch,), dtype int32
The pointer into conv_state_indices, where the last cache block to be filled is located.
initial_state_idx: (batch,), dtype int32
The pointer into conv_state_indices, where the cache block containing the initial state is located.
num_accepted_tokens: (batch,), dtype int32
If not None, it indicates the number of accepted tokens for each
sequence in the batch.
This is used in speculative decoding, where the conv_state is updated
in a sliding window manner.
query_start_loc: (batch + 1,) int32
If not None, the inputs is given in a varlen fashion and this indicates
the starting index of each sequence in the batch.
max_query_len: int
If query_start_loc is not None, this indicates the maximum query
length in the batch.
pad_slot_id: int
if conv_state_indices is passed, lets the kernel identify padded
entries that will not be processed,
for example: conv_state_indices = [pad_slot_id, 1 ,20 ,pad_slot_id]
in this case, the kernel will not process entries at
indices 0 and 3
out: (batch, dim) or (batch, dim, seqlen) or (num_tokens, dim), same shape as `x`
"""
if not HAS_TRITON:
return _pytorch_update(
x,
conv_state,
weight,
bias,
activation,
conv_state_indices=conv_state_indices,
num_accepted_tokens=num_accepted_tokens,
query_start_loc=query_start_loc,
pad_slot_id=pad_slot_id,
)
weight = weight.transpose(0, 1).contiguous()
conv_state = conv_state.transpose(1, 2).contiguous()
if validate_data:
assert pad_slot_id is not None
assert x.stride(1) == 1
if isinstance(activation, bool):
activation = "silu" if activation is True else None
elif activation is not None:
assert activation in ["silu", "swish"]
original_x_dtype = x.dtype
x = x.to(conv_state.dtype)
unsqueeze = query_start_loc is None and x.dim() == 2
if unsqueeze:
# make it (batch, dim, seqlen) with seqlen == 1
x = x.unsqueeze(1)
if query_start_loc is None:
batch, seqlen, dim = x.shape
else:
assert conv_state_indices is not None
batch = conv_state_indices.size(0)
dim = x.size(1)
seqlen = max_query_len
width, _ = weight.shape
num_cache_lines, state_len_total, _ = conv_state.size()
# overwrite-on-x strategy same as original
out = x
stride_w_width, stride_w_dim = weight.stride()
if query_start_loc is None:
stride_x_seq, stride_x_token, stride_x_dim = x.stride()
stride_o_seq, stride_o_token, stride_o_dim = out.stride()
else:
stride_x_token, stride_x_dim = x.stride()
stride_x_seq = 0
stride_o_token, stride_o_dim = out.stride()
stride_o_seq = 0
stride_istate_seq, stride_istate_token, stride_istate_dim = conv_state.stride()
stride_state_indices = conv_state_indices.stride(0) if conv_state_indices is not None else 0
# effective state_len exactly as original
if num_accepted_tokens is not None:
eff_state_len = width - 1 + (seqlen - 1)
else:
eff_state_len = width - 1
np2_statelen = triton.next_power_of_2(eff_state_len)
# -------- tiling heuristic--------
# keep program count around ~[80..160]
CORE_HINT = get_vectorcore_num()
# channel tile: 512 when dim large (reduce tasks), else 256
block_n = 512 if dim >= 512 else 256
g = triton.cdiv(dim, block_n)
target = 2 * CORE_HINT # ~80
b_tile_raw = max(1, (batch * g + target - 1) // target)
# clamp to small set
if b_tile_raw <= 1:
b_tile = 1
elif b_tile_raw <= 2:
b_tile = 2
elif b_tile_raw <= 4:
b_tile = 4
else:
b_tile = 8
# token chunk based on block_n (32KB UB idea); conservative
t_chunk = 1 if block_n == 512 else 48
def grid(META):
return (
triton.cdiv(batch, META["B_TILE"]),
triton.cdiv(dim, META["BLOCK_N"]),
)
_causal_conv1d_update_kernel_npu_tiled[grid](
x,
weight,
bias,
conv_state,
conv_state_indices,
num_accepted_tokens,
query_start_loc,
block_idx_last_scheduled_token,
initial_state_idx,
out,
batch,
dim,
seqlen,
eff_state_len,
num_cache_lines,
stride_x_seq,
stride_x_dim,
stride_x_token,
stride_w_dim,
stride_w_width,
stride_istate_seq,
stride_istate_dim,
stride_istate_token,
stride_state_indices,
stride_o_seq,
stride_o_dim,
stride_o_token,
pad_slot_id,
HAS_BIAS=bias is not None,
KERNEL_WIDTH=width,
SILU_ACTIVATION=activation in ["silu", "swish"],
IS_VARLEN=query_start_loc is not None,
IS_APC_ENABLED=block_idx_last_scheduled_token is not None,
IS_SPEC_DECODING=num_accepted_tokens is not None,
NP2_STATELEN=np2_statelen,
USE_PAD_SLOT=pad_slot_id is not None,
BLOCK_N=block_n,
B_TILE=b_tile,
T_CHUNK=t_chunk,
)
if unsqueeze:
out = out.squeeze(1)
return out.to(original_x_dtype)

View File

@@ -0,0 +1,627 @@
#
# Copyright (c) 2026 Huawei Technologies Co., Ltd. All Rights Reserved.
# This file is a part of the vllm-ascend project.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
"""NPU-compatible linear attention operators for BailingMoE.
This module provides NPU-compatible replacements for GPU-only Triton kernels
used in BailingMoELinearAttention:
- ``linear_decode_forward_npu``: replaces ``linear_decode_forward_triton``
- ``LightningAttentionKernelNPU``: replaces ``MiniMaxText01LinearKernel``
"""
import torch
from einops import rearrange
from vllm.triton_utils import tl, triton
@triton.jit
def _fwd_diag_kernel(
Q,
K,
V,
Out,
S,
b: tl.constexpr,
h: tl.constexpr,
n: tl.constexpr,
d: tl.constexpr,
e: tl.constexpr,
BLOCK: tl.constexpr,
CBLOCK: tl.constexpr,
NUM_BLOCK: tl.constexpr,
):
# This kernel computes the diagonal blocks of the attention matrix
# Each diagonal block represents attention
# where queries attend to keys in the same block
off = tl.program_id(0)
off_bh = off // NUM_BLOCK # batch-head index
off_block = off % NUM_BLOCK # block index within the sequence
off_cblock = tl.program_id(1) # sub-block index within a block
off_h = off_bh % h # head index
# Calculate base offsets for the current batch and head
qk_offset = off_bh * n * d
v_offset = off_bh * n * e
o_offset = off_bh * n * e
# Calculate offsets for the current block
block_offset = off_block * BLOCK
qk_block_offset = block_offset * d
v_block_offset = block_offset * e
o_block_offset = block_offset * e
# Calculate offsets for the current sub-block
cblock_offset = off_cblock * CBLOCK
q_cblock_offset = cblock_offset * d
o_cblock_offset = cblock_offset * e
# Calculate pointers to the query, key, value, and output tensors
Q_block_ptr = (
Q + qk_offset + qk_block_offset + q_cblock_offset + tl.arange(0, CBLOCK)[:, None] * d + tl.arange(0, d)[None, :]
)
K_block_ptr = K + qk_offset + qk_block_offset + tl.arange(0, CBLOCK)[:, None] * d + tl.arange(0, d)[None, :]
V_block_ptr = V + v_offset + v_block_offset + tl.arange(0, CBLOCK)[:, None] * e + tl.arange(0, e)[None, :]
O_block_ptr = (
Out + o_offset + o_block_offset + o_cblock_offset + tl.arange(0, CBLOCK)[:, None] * e + tl.arange(0, e)[None, :]
)
# Load the decay rate for the current head
S_block_ptr = S + off_h
s = tl.load(S_block_ptr)
i = off_cblock
q_index = tl.arange(0, CBLOCK) + i * CBLOCK
# Load query values
q = tl.load(Q_block_ptr, mask=block_offset + q_index[:, None] < n, other=0.0).to(tl.float32)
# Re-apply mask to zero out padding elements in the last block.
# On Ascend, tl.load(..., other=0.0) may not reliably clear out-of-bound data
# due to hardware-specific vector-to-cube loading behavior. If the sequence length
# is not a multiple of BLOCK_SIZE, the trailing block may contain garbage values.
# These "dirty" elements can cause NaNs during dot-product computation, leading
# to corrupted attention outputs and model instability. Explicitly masking here
# ensures numerical safety.
q = tl.where(block_offset + q_index[:, None] < n, q, 0.0)
# Initialize output accumulator
qkv = tl.zeros([CBLOCK, e], dtype=tl.float32)
# Process all sub-blocks up to and
# including the current one (causal attention)
for j in range(i + 1):
kv_index = tl.arange(0, CBLOCK) + j * CBLOCK
diff = q_index[:, None] - kv_index[None, :]
s_index = s * diff
# Apply causal mask: only attend to positions before the current one
s_index = tl.where(diff >= 0, -s_index, float("-inf"))
decay = tl.exp(s_index)
# Load key and value
k = tl.load(
K_block_ptr,
mask=block_offset + kv_index[:, None] < n,
other=0.0,
).to(tl.float32)
# Same masking required for k to prevent garbage values in dot product (see above).
k = tl.where(block_offset + kv_index[:, None] < n, k, 0.0)
v = tl.load(
V_block_ptr,
mask=block_offset + kv_index[:, None] < n,
other=0.0,
).to(tl.float32)
# Compute attention scores and apply decay
qk = tl.dot(q, k.trans()) * decay
# Compute weighted values and accumulate
qkv += tl.dot(qk, v)
# Move to the next sub-block
K_block_ptr += CBLOCK * d
V_block_ptr += CBLOCK * e
tl.store(
O_block_ptr,
qkv.to(O_block_ptr.dtype.element_ty),
mask=block_offset + q_index[:, None] < n,
)
@triton.jit
def _fwd_kv_parallel(
K,
V,
K_decay,
KV,
b: tl.constexpr,
h: tl.constexpr,
n: tl.constexpr,
d: tl.constexpr,
e: tl.constexpr,
BLOCK: tl.constexpr,
NUM_BLOCK: tl.constexpr,
D_FBLOCK: tl.constexpr,
E_FBLOCK: tl.constexpr,
NUM_FBLOCK: tl.constexpr,
CBLOCK: tl.constexpr,
NUM_CBLOCK: tl.constexpr,
):
# This kernel computes the key-value outer
# products for each block in parallel
off_bh = tl.program_id(0) # batch-head index
off_block = tl.program_id(1) # block index within the sequence
off_e = tl.program_id(2) # e-dimension tile index for UB overflow prevention
off_h = off_bh % h # head index
block_offset = off_block * BLOCK
# e-dimension tile offset: each program handles E_FBLOCK columns of e
e_offset = off_e * E_FBLOCK
# Calculate offsets for the current block
k_block_offset = block_offset * d
v_block_offset = block_offset * e
kv_block_offset = off_block * d * e
# Calculate base offsets for the current batch and head
k_offset = off_bh * n * d
v_offset = off_bh * n * e
kv_offset = off_bh * NUM_BLOCK * d * e
# Calculate pointers to the key, value, and key-value tensors
# K does not depend on e_offset (K is [n, d], not [n, e])
K_block_ptr = K + k_offset + k_block_offset + tl.arange(0, CBLOCK)[:, None] * d + tl.arange(0, D_FBLOCK)[None, :]
# V is offset by e_offset to select the current E_FBLOCK columns
V_block_ptr = (
V + v_offset + v_block_offset + tl.arange(0, CBLOCK)[:, None] * e + e_offset + tl.arange(0, E_FBLOCK)[None, :]
)
# KV is offset by e_offset to write into the correct E_FBLOCK columns
KV_block_ptr = (
KV
+ kv_offset
+ kv_block_offset
+ tl.arange(0, D_FBLOCK)[:, None] * e
+ e_offset
+ tl.arange(0, E_FBLOCK)[None, :]
)
# Load the decay factors for the current head and block
k_decay_ptr = K_decay + off_h * BLOCK + tl.arange(0, CBLOCK)
kv_index = tl.arange(0, CBLOCK)
# Initialize the key-value outer product accumulator
kv = tl.zeros([D_FBLOCK, E_FBLOCK], dtype=tl.float32)
# Handle the last block which might be smaller than BLOCK
split_n = n - (NUM_BLOCK - 1) * BLOCK if off_block == NUM_BLOCK - 1 else BLOCK
left_shift = tl.cdiv(split_n, CBLOCK) * CBLOCK - split_n
num_blocks = min(tl.cdiv(split_n, CBLOCK), NUM_CBLOCK)
k_decay_ptr += (NUM_CBLOCK - num_blocks) * CBLOCK
# Process all sub-blocks in the current block
for j in range(num_blocks):
left_bound = (1 - j) * left_shift
# Load key and value, handling boundary conditions
k_block = tl.load(
K_block_ptr - left_shift * d,
mask=(kv_index[:, None] >= left_bound),
other=0.0,
).to(tl.float32)
v = tl.load(
V_block_ptr - left_shift * e,
mask=(kv_index[:, None] >= left_bound),
other=0.0,
).to(tl.float32)
# Load decay factor and compute weighted key-value outer product
k_decay = tl.load(k_decay_ptr)
k_trans = tl.trans(k_block)
# NOTE: Need to add the extra dim here due to AMD MLIR lowering error.
# Please don't move it back until issue is resolved.
# Issue: https://github.com/ROCm/triton/issues/907
k_decay = k_decay[None, :]
kv += tl.dot(k_trans * k_decay, v)
# Move to the next sub-block
K_block_ptr += CBLOCK * d
V_block_ptr += CBLOCK * e
k_decay_ptr += CBLOCK
# Store the result
tl.store(KV_block_ptr, kv.to(KV_block_ptr.dtype.element_ty))
@triton.jit
def _fwd_kv_reduce(
S,
KV,
KV_HISTORY,
b: tl.constexpr,
h: tl.constexpr,
n,
d: tl.constexpr,
e: tl.constexpr,
BLOCK: tl.constexpr,
NUM_BLOCK,
D_FBLOCK: tl.constexpr,
E_FBLOCK: tl.constexpr,
):
# This kernel reduces the key-value outer products
# across blocks and updates the KV history
off_bh = tl.program_id(0) # batch-head index
off_e = tl.program_id(1) # e-dimension tile index for UB overflow prevention
off_h = off_bh % h # head index
# e-dimension tile offset: each program handles E_FBLOCK columns of e
e_offset = off_e * E_FBLOCK
kv_offset = off_bh * NUM_BLOCK * d * e
# Calculate pointer to the key-value tensor, offset by e_offset
KV_block_ptr = KV + kv_offset + tl.arange(0, D_FBLOCK)[:, None] * e + e_offset + tl.arange(0, E_FBLOCK)[None, :]
# Load the decay rate for the current head
s_ptrs = S + off_h
s = tl.load(s_ptrs)
# Calculate pointer to the key-value history tensor, offset by e_offset
kv_history_offset = off_bh * d * e
KV_HISTORY_block_ptr = (
KV_HISTORY
+ kv_history_offset
+ tl.arange(0, D_FBLOCK)[:, None] * e
+ e_offset
+ tl.arange(0, E_FBLOCK)[None, :]
)
# Load the previous key-value history
kv_pre = tl.load(KV_HISTORY_block_ptr).to(tl.float32)
# Process all blocks in reverse order to compute the prefix sum
for i in range(NUM_BLOCK):
block_size = min(n - i * BLOCK, BLOCK)
# Compute decay factor for the current block
block_decay = tl.exp(-s.to(tl.float32) * block_size)
# Load the current key-value outer product
kv_cur = tl.load(KV_block_ptr).to(tl.float32)
# Store the previous key-value history to the current block
tl.store(KV_block_ptr, kv_pre.to(KV_block_ptr.dtype.element_ty))
# Update the key-value history with the current block
kv_pre = block_decay * kv_pre + kv_cur
KV_block_ptr += d * e
# Store the updated key-value history
tl.store(KV_HISTORY_block_ptr, kv_pre)
@triton.jit
def _fwd_none_diag_kernel(
Q,
Out,
S,
KV,
b: tl.constexpr,
h: tl.constexpr,
n,
d: tl.constexpr,
e: tl.constexpr,
BLOCK: tl.constexpr,
NUM_BLOCK,
E_FBLOCK: tl.constexpr,
CBLOCK: tl.constexpr,
NUM_CBLOCK: tl.constexpr,
):
# This kernel computes the non-diagonal blocks of the attention matrix
# Each non-diagonal block represents attention
# where queries attend to keys in different blocks
off_bh = tl.program_id(0) # batch-head index
off_h = off_bh % h # head index
off_nc = tl.program_id(1)
off_n = off_nc // NUM_CBLOCK # block index
off_c = off_nc % NUM_CBLOCK # sub-block index
off_e = tl.program_id(2) # output feature block index
n_offset = off_n * BLOCK
c_offset = off_c * CBLOCK
e_offset = off_e * E_FBLOCK
block_offset = n_offset + c_offset
# Calculate offsets for the current batch, head, and block
q_offset = off_bh * n * d + (n_offset + c_offset) * d
o_offset = off_bh * n * e + (n_offset + c_offset) * e + e_offset
kv_offset = off_bh * NUM_BLOCK * d * e + off_n * d * e + e_offset
# Calculate pointers to the query, output, and key-value tensors
Q_block_ptr = Q + q_offset + tl.arange(0, CBLOCK)[:, None] * d + tl.arange(0, d)[None, :]
O_block_ptr = Out + o_offset + tl.arange(0, CBLOCK)[:, None] * e + tl.arange(0, E_FBLOCK)[None, :]
KV_block_ptr = KV + kv_offset + tl.arange(0, d)[:, None] * e + tl.arange(0, E_FBLOCK)[None, :]
# Load the decay rate for the current head
S_block_ptr = S + off_h
s = tl.load(S_block_ptr)
c_array = tl.arange(0, CBLOCK)
# Load the key-value outer product for the current block
kv = tl.load(KV_block_ptr).to(tl.float32)
q_index = block_offset + tl.arange(0, CBLOCK)
# Load query values
q = tl.load(Q_block_ptr, mask=q_index[:, None] < n, other=0.0).to(tl.float32)
# Compute decay factors for the current sub-block
q_decay = tl.exp(-s.to(tl.float32) * (off_c * CBLOCK + c_array[:, None]))
# Compute non-diagonal attention output
qkv_none_diag = tl.dot(q, kv) * q_decay
# Load diagonal attention output (computed by _fwd_diag_kernel)
qkv_diag = tl.load(O_block_ptr, mask=q_index[:, None] < n, other=0.0).to(tl.float32)
# Combine diagonal and non-diagonal attention outputs
qkv = qkv_diag + qkv_none_diag
# Store the result
tl.store(O_block_ptr, qkv.to(O_block_ptr.dtype.element_ty), mask=q_index[:, None] < n)
class _attention(torch.autograd.Function):
@staticmethod
def forward(ctx, q, k, v, s, kv_history):
# Forward pass of the lightning attention algorithm
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
s = s.contiguous()
# Get input dimensions
b, h, n, d = q.shape
e = v.shape[-1]
# Initialize output tensor
o = torch.empty((b, h, n, e), dtype=q.dtype, device=q.device)
# ------------------------------------------------------------------ #
# Tiling parameters (NPU UB-safe, all kernels share the same BLOCK) #
# #
# BLOCK = 256 : sequence tile size, unified across all kernels #
# to keep diag / non-diag semantics consistent. #
# CBLOCK_D = 32 : sub-tile for _fwd_diag_kernel #
# UB ≈ 72 KB (< 192 KB limit) #
# CBLOCK_KV = 64 : sub-tile for _fwd_kv_parallel / #
# _fwd_none_diag_kernel #
# UB ≈ 112 KB (< 192 KB limit) #
# E_FBLOCK = e//2 : split the e-dimension into two tiles so that #
# _fwd_kv_parallel UB stays within limits. #
# The full e range is covered via grid dim-2 #
# (NUM_EFBLOCK tiles), NOT by truncation. #
# ------------------------------------------------------------------ #
BLOCK = 256
NUM_BLOCK = triton.cdiv(n, BLOCK)
# Step 1: Compute diagonal blocks of attention
# Each program handles CBLOCK_D rows of Q within one BLOCK-sized tile.
# UB breakdown (fp32): q[32,d] + k[32,d] + v[32,e] +
# qk[32,32] + decay[32,32] + qkv[32,e]
# = 4*(2*32*128 + 2*32*128 + 2*32*32) ≈ 72 KB
CBLOCK_D = 32
NUM_CBLOCK_D = BLOCK // CBLOCK_D
assert BLOCK % CBLOCK_D == 0, "BLOCK must be a multiple of CBLOCK_D"
grid_diag = (b * h * NUM_BLOCK, NUM_CBLOCK_D)
_fwd_diag_kernel[grid_diag](
q,
k,
v,
o,
s,
b,
h,
n,
d,
e,
BLOCK=BLOCK,
CBLOCK=CBLOCK_D,
NUM_BLOCK=NUM_BLOCK,
multibuffer=True,
limit_auto_multi_buffer_only_for_local_buffer=False,
set_workspace_multibuffer=4,
tile_mix_vector_loop=2,
tile_mix_cube_loop=2,
)
# Compute decay factors for keys (shape: [h, BLOCK])
array = torch.arange(0, BLOCK, device=q.device) + 1
k_decay = torch.exp(-s * (BLOCK - array.reshape(1, -1)))
# Feature-dimension tiling:
# D_FBLOCK covers the full d dimension (no split needed for d).
# E_FBLOCK splits e into NUM_EFBLOCK tiles; each tile is processed
# by a separate program (grid dim-2) so the full e is always covered.
D_FBLOCK = d # process all d columns in one shot
E_FBLOCK = e // 2 # half of e per program instance
NUM_EFBLOCK = e // E_FBLOCK # = 2 tiles to cover full e
assert e % E_FBLOCK == 0, "e must be divisible by E_FBLOCK"
CBLOCK_KV = 64
NUM_CBLOCK_KV = BLOCK // CBLOCK_KV
assert BLOCK % CBLOCK_KV == 0, "BLOCK must be a multiple of CBLOCK_KV"
# Step 2: Compute key-value outer products for each block in parallel.
# Grid dim-2 (NUM_EFBLOCK) ensures the full e dimension is covered
# without UB overflow.
# UB breakdown (fp32): kv[d,E_FBLOCK] + k[CBLOCK_KV,d] +
# v[CBLOCK_KV,E_FBLOCK] + k_trans*decay[d,CBLOCK_KV]
# = 4*(128*64 + 64*128 + 64*64 + 128*64) ≈ 112 KB
kv = torch.empty((b, h, NUM_BLOCK, d, e), dtype=torch.float32, device=q.device)
grid_kv = (b * h, NUM_BLOCK, NUM_EFBLOCK)
_fwd_kv_parallel[grid_kv](
k,
v,
k_decay,
kv,
b,
h,
n,
d,
e,
BLOCK=BLOCK,
NUM_BLOCK=NUM_BLOCK,
D_FBLOCK=D_FBLOCK,
E_FBLOCK=E_FBLOCK,
NUM_FBLOCK=NUM_EFBLOCK,
CBLOCK=CBLOCK_KV,
NUM_CBLOCK=NUM_CBLOCK_KV,
)
# Step 3: Reduce key-value outer products across blocks and update
# KV history. Grid dim-1 (NUM_EFBLOCK) covers the full e dimension.
# UB breakdown (fp32): kv_pre[d,E_FBLOCK] + kv_cur[d,E_FBLOCK]
# = 2*4*128*64 = 64 KB
grid_reduce = (b * h, NUM_EFBLOCK)
_fwd_kv_reduce[grid_reduce](
s,
kv,
kv_history,
b,
h,
n,
d,
e,
BLOCK=BLOCK,
NUM_BLOCK=NUM_BLOCK,
D_FBLOCK=D_FBLOCK,
E_FBLOCK=E_FBLOCK,
)
# Step 4: Compute non-diagonal blocks of attention.
# Grid dim-2 (NUM_EFBLOCK) covers the full e dimension.
# UB breakdown (fp32): kv[d,E_FBLOCK] + q[CBLOCK_KV,d] +
# qkv_none[CBLOCK_KV,E_FBLOCK] +
# qkv_diag[CBLOCK_KV,E_FBLOCK] + q_decay[CBLOCK_KV,1]
# = 4*(128*64 + 64*128 + 64*64 + 64*64 + 64) ≈ 96 KB
grid_none_diag = (b * h, NUM_BLOCK * NUM_CBLOCK_KV, NUM_EFBLOCK)
_fwd_none_diag_kernel[grid_none_diag](
q,
o,
s,
kv,
b,
h,
n,
d,
e,
BLOCK=BLOCK,
NUM_BLOCK=NUM_BLOCK,
E_FBLOCK=E_FBLOCK,
CBLOCK=CBLOCK_KV,
NUM_CBLOCK=NUM_CBLOCK_KV,
)
# Save tensors for backward pass
ctx.save_for_backward(q, k, v, s, kv)
ctx.BLOCK = BLOCK
return o, torch.cat([kv, kv_history.unsqueeze(2)], dim=2)
lightning_attention_npu_ = _attention.apply
def lightning_attention_npu(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
ed: torch.Tensor,
block_size: int,
kv_history: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
"""lightning attention forward pass (NPU-friendly)."""
d = q.shape[-1]
e = v.shape[-1]
if ed.dim() == 1:
ed = ed.view(1, -1, 1, 1)
# Split the computation into chunks for better parallelism
m = 128 if d >= 128 else 64
arr = [m * i for i in range(d // m + 1)]
if arr[-1] != d:
arr.append(d)
n = len(arr)
output = 0
# Initialize or clone key-value history
if kv_history is None:
kv_history = torch.zeros((q.shape[0], q.shape[1], d, e), dtype=torch.float32, device=q.device)
else:
kv_history = kv_history.clone().contiguous()
# Process each chunk and accumulate results
for i in range(n - 1):
s = arr[i]
e = arr[i + 1]
q1 = q[..., s:e]
k1 = k[..., s:e]
o, kv = lightning_attention_npu_(q1, k1, v, ed, kv_history)
output = output + o
return output, kv
class AscendLightningAttentionKernel:
"""NPU-friendly lightning attention kernel for BailingMoE prefill.
Replaces ``MiniMaxText01LinearKernel`` by providing an NPU-friendly
implementation of the prefill forward pass
"""
@staticmethod
def jit_linear_forward_prefix(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
kv_caches: torch.Tensor,
slope_rate: torch.Tensor,
block_size: int,
layer_idx: int | None = None,
**kwargs,
) -> torch.Tensor:
slope_rate = slope_rate.to(torch.float32)
should_squeeze = q.dim() == 3
if should_squeeze:
q = q.unsqueeze(0)
k = k.unsqueeze(0)
v = v.unsqueeze(0)
b, h, n, d = q.shape
e = v.shape[-1]
kv_history = kv_caches.reshape(1, h, d, e).contiguous()
output, kv_history = lightning_attention_npu(
q,
k,
v,
slope_rate,
block_size=block_size,
kv_history=kv_history,
)
kv_caches.copy_(kv_history[:, :, -1, :, :].reshape(h, d, e))
assert output.shape[0] == 1, "batch size must be 1"
return rearrange(output.squeeze(0), "h n d -> n (h d)")

View File

@@ -0,0 +1,157 @@
# Adapt from https://github.com/vllm-project/vllm/blob/main/vllm/v1/worker/mamba_utils.py
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from vllm.triton_utils import tl, triton
@triton.jit
def postprocess_mamba_fused_kernel(
# Decision inputs (per-request)
num_accepted_tokens_ptr,
mamba_state_idx_ptr,
num_scheduled_tokens_ptr,
num_computed_tokens_ptr,
num_draft_tokens_ptr,
# Per-group block table base addresses: int64[num_groups]. Each entry is
# the data_ptr of that group's persistent [max_reqs, max_blocks] int32
# block table.
block_table_ptrs_ptr,
block_table_stride_req: tl.int64, # stride between requests (in elements)
# Mamba state metadata (per-layer, per-state-type)
# These are 1D arrays indexed by (layer_idx * num_state_types + state_type_idx)
state_base_addrs_ptr, # base address of each state tensor
state_block_strides_ptr, # bytes per block for each state
state_elem_sizes_ptr, # element size for each state
state_inner_sizes_ptr, # number of elements in inner dimensions
state_conv_widths_ptr, # conv width for conv states (0 for temporal)
state_group_indices_ptr, # maps state_idx to group index in block table
# Output: num_accepted_tokens update (for src==dst case)
num_accepted_tokens_out_ptr,
# Runtime parameter (varies per batch - NOT constexpr to avoid recompilation)
num_reqs,
# Compile-time constants (fixed after model initialization)
# block_size: determined by model config, constant for all invocations
block_size: tl.constexpr,
# COPY_BLOCK_SIZE: fixed tuning parameter for memory copy loop
COPY_BLOCK_SIZE: tl.constexpr,
):
"""
Fused GPU kernel for postprocess_mamba that computes decisions AND performs
mamba state copies without any CPU-GPU synchronization.
Grid: (num_reqs, num_layers * num_state_types)
- program_id(0) = request index
- program_id(1) = state_idx (flattened index into layer/state_type metadata)
Note: num_layers and num_state_types are not passed as kernel parameters
because the kernel indexes directly into pre-flattened metadata arrays
using program_id(1). The grid dimensions encode the total state count.
"""
req_idx = tl.program_id(0)
state_idx = tl.program_id(1)
# Bounds check
if req_idx >= num_reqs:
return
# Compute decision logic (mirrors postprocess_mamba Python reference)
num_accepted = tl.load(num_accepted_tokens_ptr + req_idx)
src_block_idx = tl.load(mamba_state_idx_ptr + req_idx)
num_scheduled = tl.load(num_scheduled_tokens_ptr + req_idx)
num_computed = tl.load(num_computed_tokens_ptr + req_idx)
num_draft = tl.load(num_draft_tokens_ptr + req_idx)
num_tokens_running_state = num_computed + num_scheduled - num_draft
new_num_computed = num_tokens_running_state + num_accepted - 1
aligned_new_computed = (new_num_computed // block_size) * block_size
needs_copy = aligned_new_computed >= num_tokens_running_state
if not needs_copy:
return
# Compute copy parameters
accept_token_bias = aligned_new_computed - num_tokens_running_state
dest_block_idx = aligned_new_computed // block_size - 1
# Load state metadata for this layer/state_type
state_base_addr = tl.load(state_base_addrs_ptr + state_idx)
state_block_stride = tl.load(state_block_strides_ptr + state_idx)
state_elem_size = tl.load(state_elem_sizes_ptr + state_idx)
state_inner_size = tl.load(state_inner_sizes_ptr + state_idx)
conv_width = tl.load(state_conv_widths_ptr + state_idx)
# Load the group index for this state, then index into the correct
# group's block table. Each mamba group has independently allocated
# physical blocks.
group_idx = tl.load(state_group_indices_ptr + state_idx).to(tl.int64)
# block_table_ptrs_ptr holds one pointer per group (each group owns its own
# block table). Reinterpret as int32* since block ids are int32.
group_base_addr = tl.load(block_table_ptrs_ptr + group_idx)
block_table_typed = group_base_addr.to(tl.pointer_type(tl.int32))
block_table_base = block_table_typed + req_idx * block_table_stride_req
# Widen block ids to int64 before they reach `block_id * state_block_stride`
# below: state_block_stride can exceed 2**31 bytes for large mamba caches,
# and Triton would otherwise do the multiply in int32 and wrap.
src_block_id = tl.load(block_table_base + src_block_idx).to(tl.int64)
dest_block_id = tl.load(block_table_base + dest_block_idx).to(tl.int64)
# Compute source and destination addresses based on state type
# conv_width > 0 means this is a conv state (get_conv_copy_spec logic)
# conv_width == 0 means this is a temporal state (get_temporal_copy_spec logic)
is_conv_state = conv_width > 0
if is_conv_state:
# Conv state: copy
# state[block_table[req_idx, src_block_idx], accept_token_bias:]
# to
# state[block_table[req_idx, dest_block_idx], :conv_width - accept_token_bias]
src_offset = accept_token_bias.to(tl.int64) * state_inner_size * state_elem_size
src_addr = state_base_addr + src_block_id * state_block_stride + src_offset
dst_addr = state_base_addr + dest_block_id * state_block_stride
# Number of elements to copy:
# (conv_width - accept_token_bias) * inner_size
num_elems_to_copy = (conv_width - accept_token_bias).to(tl.int64) * state_inner_size
copy_size = num_elems_to_copy * state_elem_size
else:
# Temporal state: copy
# state[block_table[req_idx, src_block_idx + accept_token_bias]]
# to
# state[block_table[req_idx, dest_block_idx]]
actual_src_block_idx = src_block_idx + accept_token_bias
actual_src_block_id = tl.load(block_table_base + actual_src_block_idx).to(tl.int64)
src_addr = state_base_addr + actual_src_block_id * state_block_stride
dst_addr = state_base_addr + dest_block_id * state_block_stride
# Use natural block data size (inner_size * elem_size), NOT
# state_block_stride which is the page stride and can exceed the
# actual data when the state tensor uses as_strided page padding.
copy_size = state_inner_size * state_elem_size
# Mirror postprocess_mamba's trailing
# if src_block_idx == dest_block_idx: num_accepted_tokens_cpu[i] = 1
# This runs whether or not the copy below is skipped (it's per-request, so
# only state_idx == 0 writes).
if src_block_idx == dest_block_idx and state_idx == 0:
tl.store(num_accepted_tokens_out_ptr + req_idx, 1)
# Mirror collect_mamba_copy_meta's early return: src==dst with no token
# bias means source and destination ranges coincide, so the copy is a
# no-op.
if src_block_idx == dest_block_idx and accept_token_bias == 0:
return
# Hoist the pointer-type cast out of the copy loop. triton-ascend's
# PtrOffsetInfo::AxisInfo analysis aborts on `(addr + i + offsets).to(...)`
# inside the loop (SmallVector assertion `idx < size()`); casting once
# here and doing plain pointer arithmetic inside the loop is the same fix
# vllm-ascend applies to batch_memcpy_kernel.
src_ptr = src_addr.to(tl.pointer_type(tl.uint8))
dst_ptr = dst_addr.to(tl.pointer_type(tl.uint8))
offsets = tl.arange(0, COPY_BLOCK_SIZE)
for i in range(0, copy_size, COPY_BLOCK_SIZE):
mask = (i + offsets) < copy_size
data = tl.load(src_ptr + i + offsets, mask=mask)
tl.store(dst_ptr + i + offsets, data, mask=mask)