[NPU] add causal conv1d for ascend kda backend (#35021)

This commit is contained in:
zhaozx-cn
2026-08-29 11:39:45 +08:00
committed by GitHub
parent 8a4c517a60
commit ca8cc101b8
2 changed files with 104 additions and 105 deletions
@@ -13,13 +13,6 @@ from sgl_kernel_npu.fla.kda_prefill import (
from sgl_kernel_npu.fla.kda_target_verify import kda_target_verify_npu from sgl_kernel_npu.fla.kda_target_verify import kda_target_verify_npu
from sgl_kernel_npu.fla.solve_tril import solve_tril_npu from sgl_kernel_npu.fla.solve_tril import solve_tril_npu
from sgl_kernel_npu.fla.utils import prepare_chunk_indices from sgl_kernel_npu.fla.utils import prepare_chunk_indices
from sgl_kernel_npu.mamba.causal_conv1d import (
causal_conv1d_fn_npu,
causal_conv1d_update_npu,
)
from sgl_kernel_npu.mamba.causal_conv1d_verify import (
causal_conv1d_linear_verify_npu,
)
from sglang.kernels.ops.attention.fla.cumsum import chunk_local_cumsum from sglang.kernels.ops.attention.fla.cumsum import chunk_local_cumsum
from sglang.kernels.ops.attention.fla.kda import chunk_kda_scaled_dot_kkt_fwd from sglang.kernels.ops.attention.fla.kda import chunk_kda_scaled_dot_kkt_fwd
@@ -125,18 +118,52 @@ class AscendKDAAttnBackend(KDAAttnBackend):
The model, scheduler, metadata, and non-operator control flow stay in the The model, scheduler, metadata, and non-operator control flow stay in the
shared KDA backend. This class contains only the layout and operator shared KDA backend. This class contains only the layout and operator
differences required by Ascend. differences required by Ascend.
Conv states use the GDN-style [layers, pool, window, channels] layout
(transposed from the shared KDA backend's [channels, window]). The
speculative window is extended by draft_tokens - 1 so that verify
writes all draft token conv states directly into conv_states; after
verify, conv_state_rollback reverts unaccepted tokens. This replaces
the previous intermediate_conv_window snapshot + scatter scheme.
""" """
supports_speculative_conv_state_snapshots: bool = True supports_speculative_conv_state_snapshots: bool = False
def __init__(self, model_runner): def __init__(self, model_runner):
super().__init__(model_runner) super().__init__(model_runner)
# The NPU pool is allocated directly as [pool, channels, window]. # The NPU pool is allocated as [layers, pool, window, channels]
self.conv_states_shape = ( # (transposed from the shared KDA [channels, window]). Expose the
model_runner.req_to_token_pool.mamba_pool.mamba_cache.conv[0].shape # transposed shape so _init_track_conv_indices reads
# conv_states_shape[-1] as the conv window length.
conv_pool_shape = model_runner.req_to_token_pool.mamba_pool.mamba_cache.conv[
0
].shape
self.conv_states_shape = torch.Size(
(
*conv_pool_shape[:-2],
conv_pool_shape[-1],
conv_pool_shape[-2],
)
) )
self.kernel_dispatcher.extend_kernel = _AscendKDAExtendKernel() self.kernel_dispatcher.extend_kernel = _AscendKDAExtendKernel()
def _get_conv_weights_t(
self, layer: RadixLinearAttention, dtype: torch.dtype
) -> torch.Tensor:
"""Transposed conv weights [width, dim], cached on the layer.
The NPU causal_conv1d CANN op expects weight as [width, dim]
(transposed from layer.conv_weights [dim, width]) and requires
weight dtype to match the input. KDA keeps conv_weights in FP32
while inputs/conv_states are BF16, so the cached FP32 transpose is
cast to the caller's dtype here.
"""
w = getattr(layer, "_conv_weights_t", None)
if w is None:
w = layer.conv_weights.transpose(0, 1).contiguous().to(dtype)
layer._conv_weights_t = w
return w
def forward_decode( def forward_decode(
self, self,
layer: RadixLinearAttention, layer: RadixLinearAttention,
@@ -154,13 +181,17 @@ class AscendKDAAttnBackend(KDAAttnBackend):
query_start_loc = self.forward_metadata.query_start_loc query_start_loc = self.forward_metadata.query_start_loc
cache_indices = self.forward_metadata.mamba_cache_indices cache_indices = self.forward_metadata.mamba_cache_indices
qkv = causal_conv1d_update_npu( # setting activation_mode to 1 means using SiLU activation after conv.
mixed_qkv, qkv = torch.ops.npu.causal_conv1d(
conv_states, mixed_qkv.contiguous(),
layer.conv_weights, self._get_conv_weights_t(layer, mixed_qkv.dtype),
layer.bias, conv_states=conv_states,
activation="silu", bias=layer.bias,
conv_state_indices=cache_indices, query_start_loc=query_start_loc,
cache_indices=cache_indices,
activation_mode=1,
pad_slot_id=-1,
run_mode=1,
) )
if self.kernel_dispatcher.supports_packed_decode: if self.kernel_dispatcher.supports_packed_decode:
@@ -245,37 +276,27 @@ class AscendKDAAttnBackend(KDAAttnBackend):
has_initial_state = forward_batch.extend_prefix_lens > 0 has_initial_state = forward_batch.extend_prefix_lens > 0
if self.forward_metadata.has_mamba_track_mask: if self.forward_metadata.has_mamba_track_mask:
conv_states[self.forward_metadata.conv_states_mask_indices] = mixed_qkv[ mixed_qkv_to_track = mixed_qkv[self.forward_metadata.track_conv_indices]
self.forward_metadata.track_conv_indices conv_states[self.forward_metadata.conv_states_mask_indices] = (
].transpose(-1, -2) mixed_qkv_to_track
splits = [layer.q_dim, layer.k_dim, layer.v_dim]
q, k, v = mixed_qkv.transpose(0, 1).split(splits, dim=0)
q_conv_weight, k_conv_weight, v_conv_weight = layer.conv_weights.split(
splits, dim=0
) )
q_conv_state, k_conv_state, v_conv_state = conv_states.split(splits, dim=-2)
if layer.bias is not None:
q_bias, k_bias, v_bias = layer.bias.split(splits, dim=0)
else:
q_bias, k_bias, v_bias = None, None, None
conv_kwargs = dict( kernel_size = layer.conv_weights.shape[-1]
has_initial_state=has_initial_state, conv_states_for_prefill = conv_states[:, -(kernel_size - 1) :, :].contiguous()
cache_indices=cache_indices, mixed_qkv = torch.ops.npu.causal_conv1d(
mixed_qkv.contiguous(),
self._get_conv_weights_t(layer, mixed_qkv.dtype),
conv_states=conv_states_for_prefill,
bias=layer.bias,
query_start_loc=query_start_loc, query_start_loc=query_start_loc,
seq_lens_cpu=forward_batch.extend_seq_lens_cpu, cache_indices=cache_indices,
has_initial_state=has_initial_state,
activation_mode=1,
pad_slot_id=-1,
run_mode=0,
) )
q = self._causal_conv1d_extend( conv_states[:, -(kernel_size - 1) :, :] = conv_states_for_prefill
q, q_conv_weight, q_bias, q_conv_state, **conv_kwargs q, k, v = mixed_qkv.split([layer.q_dim, layer.k_dim, layer.v_dim], dim=-1)
)
k = self._causal_conv1d_extend(
k, k_conv_weight, k_bias, k_conv_state, **conv_kwargs
)
v = self._causal_conv1d_extend(
v, v_conv_weight, v_bias, v_conv_state, **conv_kwargs
)
q = q.unflatten(-1, (-1, layer.head_q_dim)).unsqueeze(0) q = q.unflatten(-1, (-1, layer.head_q_dim)).unsqueeze(0)
k = k.unflatten(-1, (-1, layer.head_k_dim)).unsqueeze(0) k = k.unflatten(-1, (-1, layer.head_k_dim)).unsqueeze(0)
v = v.unflatten(-1, (-1, layer.head_v_dim)).unsqueeze(0) v = v.unflatten(-1, (-1, layer.head_v_dim)).unsqueeze(0)
@@ -309,42 +330,6 @@ class AscendKDAAttnBackend(KDAAttnBackend):
) )
return core_attn_out return core_attn_out
def _causal_conv1d_extend(
self,
x: torch.Tensor,
weight: torch.Tensor,
bias: Optional[torch.Tensor],
state: torch.Tensor,
*,
has_initial_state: torch.Tensor,
cache_indices: torch.Tensor,
query_start_loc: torch.Tensor,
seq_lens_cpu: torch.Tensor,
) -> torch.Tensor:
# The Ascend varlen kernel pads in the weight dtype. K3 keeps its
# weights in FP32 and its persistent convolution cache in BF16, so use
# a compact FP32 working set for the active rows and cast it back.
local_indices = torch.arange(
cache_indices.shape[0],
device=cache_indices.device,
dtype=cache_indices.dtype,
)
state_work = state.index_select(0, cache_indices.to(torch.int64))
state_work = state_work.to(weight.dtype).contiguous()
out = causal_conv1d_fn_npu(
x.to(weight.dtype),
weight,
bias,
activation="silu",
conv_states=state_work,
has_initial_state=has_initial_state,
cache_indices=local_indices,
query_start_loc=query_start_loc,
seq_lens_cpu=seq_lens_cpu,
)
state.index_copy_(0, cache_indices.to(torch.int64), state_work.to(state.dtype))
return out.to(x.dtype).transpose(0, 1)
def _prepare_extend_gate_inputs( def _prepare_extend_gate_inputs(
self, self,
layer: RadixLinearAttention, layer: RadixLinearAttention,
@@ -418,18 +403,32 @@ class AscendKDAAttnBackend(KDAAttnBackend):
) )
intermediate_indices = self.verify_intermediate_state_indices[:batch_size] intermediate_indices = self.verify_intermediate_state_indices[:batch_size]
processed_qkv = causal_conv1d_linear_verify_npu( conv_states = cache.conv[0]
dense_qkv.transpose(1, 2).contiguous(), num_accepted_tokens = torch.full(
cache.conv[0], (batch_size,),
layer.conv_weights, draft_token_num,
layer.bias, dtype=torch.int32,
cache_indices[:batch_size], device=mixed_qkv.device,
cache.intermediate_conv_window[0], )
intermediate_indices, dense_query_start_loc = torch.arange(
activation="silu", 0,
update_persistent_state=False, num_dense_tokens + 1,
step=draft_token_num,
dtype=torch.int32,
device=mixed_qkv.device,
)
processed_qkv = torch.ops.npu.causal_conv1d(
dense_qkv.reshape(num_dense_tokens, -1).contiguous(),
self._get_conv_weights_t(layer, mixed_qkv.dtype),
conv_states=conv_states,
bias=layer.bias,
query_start_loc=dense_query_start_loc,
cache_indices=cache_indices[:batch_size],
num_accepted_tokens=num_accepted_tokens,
activation_mode=1,
pad_slot_id=-1,
run_mode=1,
) )
processed_qkv = processed_qkv.transpose(1, 2).reshape(num_dense_tokens, -1)
q, k, v = processed_qkv.split([layer.q_dim, layer.k_dim, layer.v_dim], dim=-1) q, k, v = processed_qkv.split([layer.q_dim, layer.k_dim, layer.v_dim], dim=-1)
q = q.unflatten(-1, (-1, layer.head_q_dim)).unsqueeze(0) q = q.unflatten(-1, (-1, layer.head_q_dim)).unsqueeze(0)
k = k.unflatten(-1, (-1, layer.head_k_dim)).unsqueeze(0) k = k.unflatten(-1, (-1, layer.head_k_dim)).unsqueeze(0)
@@ -603,11 +602,10 @@ class AscendKDAHybridLinearAttnBackend:
) )
else: else:
track_mask = mamba_steps_to_track >= 0 track_mask = mamba_steps_to_track >= 0
track_indices = mamba_track_indices[track_mask] src_slots = torch.where(
if track_indices.numel() > 0: track_mask, dst_indices_tensor, mamba_track_indices
conv_states[:, track_indices] = conv_states[ )
:, dst_indices_tensor[track_mask] conv_states[:, mamba_track_indices] = conv_states[:, src_slots]
]
if not has_conv_snapshots: if not has_conv_snapshots:
if dst_indices_tensor.numel() > 0: if dst_indices_tensor.numel() > 0:
@@ -32,18 +32,19 @@ def _init_npu_conv_state(
if speculative_num_draft_tokens is not None: if speculative_num_draft_tokens is not None:
extra_conv_len = speculative_num_draft_tokens - 1 extra_conv_len = speculative_num_draft_tokens - 1
# Mamba shapes are (channels, window), while KDA shapes are # Both KDA and Mamba/GDN NPU conv states use the unified
# (window, channels). NPU kernels consume KDA state as # [layers, pool, window, channels] layout. KDA shapes arrive as
# [layers, pool, channels, window] and other Mamba state as # (window, channels) while Mamba/GDN shapes arrive as (channels, window);
# [layers, pool, window, channels]. KDA keeps the base window fixed; # resolve the correct axis ordering and extend the window by
# speculative per-step windows live in the intermediate cache. # speculative_num_draft_tokens - 1 so that verify can write all draft
# token conv states directly into conv_states (GDN rollback scheme).
conv_state = [ conv_state = [
torch.zeros( torch.zeros(
size=( size=(
conv_state_in.shape[0], conv_state_in.shape[0],
conv_state_in.shape[1], conv_state_in.shape[1],
conv_shape[1] if is_kda else conv_shape[1] + extra_conv_len, (conv_shape[0] if is_kda else conv_shape[1]) + extra_conv_len,
conv_shape[0], conv_shape[1] if is_kda else conv_shape[0],
), ),
dtype=conv_state_in.dtype, dtype=conv_state_in.dtype,
device=conv_state_in.device, device=conv_state_in.device,