0
vllm_ascend/ops/triton/mamba/__init__.py
Normal file
0
vllm_ascend/ops/triton/mamba/__init__.py
Normal file
588
vllm_ascend/ops/triton/mamba/causal_conv1d.py
Normal file
588
vllm_ascend/ops/triton/mamba/causal_conv1d.py
Normal 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)
|
||||
627
vllm_ascend/ops/triton/mamba/lightning_attn.py
Normal file
627
vllm_ascend/ops/triton/mamba/lightning_attn.py
Normal 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)")
|
||||
157
vllm_ascend/ops/triton/mamba/postprocess.py
Normal file
157
vllm_ascend/ops/triton/mamba/postprocess.py
Normal 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)
|
||||
Reference in New Issue
Block a user