Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
from sgl_kernel_npu.mamba.causal_conv1d import (
causal_conv1d_fn_npu,
causal_conv1d_update_npu,
causal_conv1d_update_v2,
)

from sglang.srt.hardware_backend.npu.attention.ascend_hybrid_linear_attn_backend import (
Expand Down Expand Up @@ -224,9 +225,7 @@ def forward_extend(
else:
has_initial_states = forward_batch.extend_prefix_lens > 0
if is_target_verify:
draft_token_num = forward_batch.spec_info.draft_token_num
num_token_padding = mixed_qkv.shape[0]
batch_size = cache_indices.shape[0]
if (
not self.graph_mode
and forward_batch.num_token_non_padded_cpu != num_token_padding
Expand All @@ -236,23 +235,24 @@ def forward_extend(
b = b[: forward_batch.num_token_non_padded_cpu]
seq_len = forward_batch.num_token_non_padded_cpu

mixed_qkv_reshaped = mixed_qkv.view(batch_size, draft_token_num, -1)
batch_size = cache_indices.shape[0]
draft_token_num = forward_batch.spec_info.draft_token_num
num_accepted_tokens = torch.full(
(batch_size,),
draft_token_num,
dtype=torch.int32,
device=mixed_qkv.device,
)
mixed_qkv = torch.ops.npu.causal_conv1d_update(
mixed_qkv_reshaped,
layer.conv_weights.transpose(0, 1).contiguous(),
conv_states,
cache_indices,
layer.bias,
num_accepted_tokens,
None,
layer.activation == "silu",
self.pad_slot_id,
mixed_qkv = causal_conv1d_update_v2(
x=mixed_qkv.view(batch_size, draft_token_num, -1).contiguous(),
conv_state=conv_states.contiguous(),
Comment thread
iridiumine marked this conversation as resolved.
weight=layer.conv_weights.transpose(0, 1).contiguous(),
bias=layer.bias,
activation=layer.activation,
conv_state_indices=cache_indices,
num_accepted_tokens=num_accepted_tokens,
pad_slot_id=-1,
Comment thread
iridiumine marked this conversation as resolved.
validate_data=False,
).view(seq_len, -1)
Comment thread
iridiumine marked this conversation as resolved.
else:
mixed_qkv = mixed_qkv.transpose(0, 1)
Expand Down
5 changes: 4 additions & 1 deletion python/sglang/srt/layers/attention/triton_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,6 @@
import torch
import triton
import triton.language as tl
from sgl_kernel.utils import is_arch_support_pdl

from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
Expand All @@ -20,9 +19,13 @@
get_bool_env_var,
get_device_core_count,
get_int_env_var,
is_npu,
next_power_of_2,
)

if not is_npu():
from sgl_kernel.utils import is_arch_support_pdl

if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.model_runner import ModelRunner
Expand Down
Loading