[NPU]Support GPT-OSS for NPU (#14197)
This commit is contained in:
@@ -5,6 +5,10 @@ from typing import TYPE_CHECKING, List, Optional
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch_npu
|
import torch_npu
|
||||||
|
from sgl_kernel_npu.attention.sinks_attention import (
|
||||||
|
attention_sinks_prefill_triton,
|
||||||
|
attention_sinks_triton,
|
||||||
|
)
|
||||||
|
|
||||||
from sglang.srt.configs.model_config import AttentionArch
|
from sglang.srt.configs.model_config import AttentionArch
|
||||||
from sglang.srt.hardware_backend.npu.attention.mla_preprocess import (
|
from sglang.srt.hardware_backend.npu.attention.mla_preprocess import (
|
||||||
@@ -260,9 +264,17 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
// self.page_size
|
// self.page_size
|
||||||
)
|
)
|
||||||
if forward_batch.extend_seq_lens is not None:
|
if forward_batch.extend_seq_lens is not None:
|
||||||
|
self.forward_metadata.extend_seq_lens = forward_batch.extend_seq_lens
|
||||||
self.forward_metadata.extend_seq_lens_cpu_int = (
|
self.forward_metadata.extend_seq_lens_cpu_int = (
|
||||||
forward_batch.extend_seq_lens.cpu().int()
|
forward_batch.extend_seq_lens.cpu().int()
|
||||||
)
|
)
|
||||||
|
if forward_batch.seq_lens is not None:
|
||||||
|
self.forward_metadata.seq_lens = forward_batch.seq_lens.int()
|
||||||
|
else:
|
||||||
|
self.forward_metadata.seq_lens = forward_batch.seq_lens_cpu.to(
|
||||||
|
self.device
|
||||||
|
).int()
|
||||||
|
|
||||||
self.forward_metadata.seq_lens_cpu_int = forward_batch.seq_lens_cpu.int()
|
self.forward_metadata.seq_lens_cpu_int = forward_batch.seq_lens_cpu.int()
|
||||||
if (
|
if (
|
||||||
not forward_batch.forward_mode.is_draft_extend_v2()
|
not forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
@@ -576,6 +588,7 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
q_rope: Optional[torch.Tensor] = None,
|
q_rope: Optional[torch.Tensor] = None,
|
||||||
k_rope: Optional[torch.Tensor] = None,
|
k_rope: Optional[torch.Tensor] = None,
|
||||||
topk_indices: Optional[torch.Tensor] = None,
|
topk_indices: Optional[torch.Tensor] = None,
|
||||||
|
sinks: Optional[torch.Tensor] = None,
|
||||||
):
|
):
|
||||||
if topk_indices is not None:
|
if topk_indices is not None:
|
||||||
return self.forward_sparse(
|
return self.forward_sparse(
|
||||||
@@ -617,6 +630,22 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||||
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||||
|
|
||||||
|
if sinks is not None:
|
||||||
|
attn_out = attention_sinks_prefill_triton(
|
||||||
|
q,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
sinks,
|
||||||
|
self.forward_metadata.extend_seq_lens,
|
||||||
|
self.forward_metadata.block_tables,
|
||||||
|
self.forward_metadata.seq_lens,
|
||||||
|
layer.scaling,
|
||||||
|
layer.sliding_window_size,
|
||||||
|
layer.tp_q_head_num,
|
||||||
|
layer.tp_k_head_num,
|
||||||
|
)
|
||||||
|
return attn_out
|
||||||
|
|
||||||
if self.use_fia:
|
if self.use_fia:
|
||||||
"""FIA will support multi-bs in the later version of CANN"""
|
"""FIA will support multi-bs in the later version of CANN"""
|
||||||
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
q = q.reshape(-1, layer.tp_q_head_num, layer.qk_head_dim)
|
||||||
@@ -1036,6 +1065,7 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
save_kv_cache: bool = True,
|
save_kv_cache: bool = True,
|
||||||
q_rope: Optional[torch.Tensor] = None,
|
q_rope: Optional[torch.Tensor] = None,
|
||||||
k_rope: Optional[torch.Tensor] = None,
|
k_rope: Optional[torch.Tensor] = None,
|
||||||
|
sinks: Optional[torch.Tensor] = None,
|
||||||
):
|
):
|
||||||
if save_kv_cache:
|
if save_kv_cache:
|
||||||
if self.use_mla:
|
if self.use_mla:
|
||||||
@@ -1049,6 +1079,24 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
layer, forward_batch.out_cache_loc, k, v
|
layer, forward_batch.out_cache_loc, k, v
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if sinks is not None:
|
||||||
|
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||||
|
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||||
|
|
||||||
|
attn_out = attention_sinks_triton(
|
||||||
|
q,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
sinks,
|
||||||
|
self.forward_metadata.block_tables,
|
||||||
|
self.forward_metadata.seq_lens,
|
||||||
|
layer.scaling,
|
||||||
|
layer.sliding_window_size,
|
||||||
|
layer.tp_q_head_num,
|
||||||
|
layer.tp_k_head_num,
|
||||||
|
)
|
||||||
|
return attn_out
|
||||||
|
|
||||||
if not self.use_mla:
|
if not self.use_mla:
|
||||||
num_tokens = q.shape[0]
|
num_tokens = q.shape[0]
|
||||||
"""PA will support bs<tp in the later version of CANN"""
|
"""PA will support bs<tp in the later version of CANN"""
|
||||||
@@ -1217,6 +1265,7 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
q_rope: Optional[torch.Tensor] = None,
|
q_rope: Optional[torch.Tensor] = None,
|
||||||
k_rope: Optional[torch.Tensor] = None,
|
k_rope: Optional[torch.Tensor] = None,
|
||||||
topk_indices: Optional[torch.Tensor] = None,
|
topk_indices: Optional[torch.Tensor] = None,
|
||||||
|
sinks: Optional[torch.Tensor] = None,
|
||||||
):
|
):
|
||||||
if is_mla_preprocess_enabled():
|
if is_mla_preprocess_enabled():
|
||||||
# MLAPO does saving kv_cache
|
# MLAPO does saving kv_cache
|
||||||
@@ -1244,6 +1293,7 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
save_kv_cache,
|
save_kv_cache,
|
||||||
q_rope=q_rope,
|
q_rope=q_rope,
|
||||||
k_rope=k_rope,
|
k_rope=k_rope,
|
||||||
|
sinks=sinks,
|
||||||
)
|
)
|
||||||
|
|
||||||
if not self.use_mla:
|
if not self.use_mla:
|
||||||
@@ -1254,6 +1304,22 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
num_tokens = q.shape[0]
|
num_tokens = q.shape[0]
|
||||||
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||||
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||||
|
|
||||||
|
if sinks is not None:
|
||||||
|
attn_out = attention_sinks_triton(
|
||||||
|
q,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
sinks,
|
||||||
|
self.forward_metadata.block_tables,
|
||||||
|
self.forward_metadata.seq_lens,
|
||||||
|
layer.scaling,
|
||||||
|
layer.sliding_window_size,
|
||||||
|
layer.tp_q_head_num,
|
||||||
|
layer.tp_k_head_num,
|
||||||
|
)
|
||||||
|
return attn_out
|
||||||
|
|
||||||
if self.use_fia:
|
if self.use_fia:
|
||||||
if self.forward_metadata.seq_lens_cpu_int is None:
|
if self.forward_metadata.seq_lens_cpu_int is None:
|
||||||
actual_seq_len_kv = self.forward_metadata.seq_lens_cpu_list
|
actual_seq_len_kv = self.forward_metadata.seq_lens_cpu_list
|
||||||
|
|||||||
@@ -492,6 +492,8 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
|
|||||||
)
|
)
|
||||||
|
|
||||||
expert_tokens = expert_tokens.to(torch.int64)
|
expert_tokens = expert_tokens.to(torch.int64)
|
||||||
|
w13_bias = [layer.w13_weight_bias] if self.with_bias else None
|
||||||
|
w2_bias = [layer.w2_weight_bias] if self.with_bias else None
|
||||||
if layer.w13_weight.shape[-1] == layer.hidden_size:
|
if layer.w13_weight.shape[-1] == layer.hidden_size:
|
||||||
w13 = layer.w13_weight.transpose(1, 2)
|
w13 = layer.w13_weight.transpose(1, 2)
|
||||||
w2 = layer.w2_weight.transpose(1, 2)
|
w2 = layer.w2_weight.transpose(1, 2)
|
||||||
@@ -500,6 +502,7 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
|
|||||||
hidden_states = torch_npu.npu_grouped_matmul(
|
hidden_states = torch_npu.npu_grouped_matmul(
|
||||||
x=[hidden_states],
|
x=[hidden_states],
|
||||||
weight=[w13],
|
weight=[w13],
|
||||||
|
bias=w13_bias,
|
||||||
split_item=2,
|
split_item=2,
|
||||||
group_list_type=0,
|
group_list_type=0,
|
||||||
group_type=0,
|
group_type=0,
|
||||||
@@ -508,7 +511,11 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
|
|||||||
)[0]
|
)[0]
|
||||||
|
|
||||||
# act_fn:
|
# act_fn:
|
||||||
if self.moe_runner_config.activation == "silu":
|
if self.moe_runner_config.activation == "npu_swiglu_oai":
|
||||||
|
from sgl_kernel_npu.activation.swiglu_oai import swiglu_oai
|
||||||
|
|
||||||
|
hidden_states = swiglu_oai(layer, hidden_states)
|
||||||
|
elif self.moe_runner_config.activation == "silu":
|
||||||
hidden_states = torch_npu.npu_swiglu(hidden_states)
|
hidden_states = torch_npu.npu_swiglu(hidden_states)
|
||||||
else:
|
else:
|
||||||
from sglang.srt.layers.activation import GeluAndMul
|
from sglang.srt.layers.activation import GeluAndMul
|
||||||
@@ -519,6 +526,7 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, MultiPlatformOp):
|
|||||||
hidden_states = torch_npu.npu_grouped_matmul(
|
hidden_states = torch_npu.npu_grouped_matmul(
|
||||||
x=[hidden_states],
|
x=[hidden_states],
|
||||||
weight=[w2],
|
weight=[w2],
|
||||||
|
bias=w2_bias,
|
||||||
split_item=2,
|
split_item=2,
|
||||||
group_list_type=0,
|
group_list_type=0,
|
||||||
group_type=0,
|
group_type=0,
|
||||||
|
|||||||
@@ -71,9 +71,10 @@ from sglang.srt.models.utils import (
|
|||||||
enable_fused_set_kv_buffer,
|
enable_fused_set_kv_buffer,
|
||||||
)
|
)
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
|
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, is_npu, make_layers
|
||||||
|
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
|
_is_npu = is_npu()
|
||||||
|
|
||||||
|
|
||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
@@ -129,6 +130,7 @@ class GptOssSparseMoeBlock(nn.Module):
|
|||||||
"use_weight_loader_fused": quant_config_name
|
"use_weight_loader_fused": quant_config_name
|
||||||
!= "mxfp4"
|
!= "mxfp4"
|
||||||
}
|
}
|
||||||
|
|
||||||
self.experts = experts_type(
|
self.experts = experts_type(
|
||||||
num_experts=config.num_local_experts
|
num_experts=config.num_local_experts
|
||||||
+ get_global_server_args().ep_num_redundant_experts,
|
+ get_global_server_args().ep_num_redundant_experts,
|
||||||
@@ -305,11 +307,10 @@ class GptOssAttention(nn.Module):
|
|||||||
qkv, _ = self.qkv_proj(hidden_states)
|
qkv, _ = self.qkv_proj(hidden_states)
|
||||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||||
|
|
||||||
q, k = self.rotary_emb(
|
extra_args = {}
|
||||||
positions,
|
if not _is_npu:
|
||||||
q,
|
extra_args = {
|
||||||
k,
|
"fused_set_kv_buffer_arg": (
|
||||||
fused_set_kv_buffer_arg=(
|
|
||||||
create_fused_set_kv_buffer_arg(
|
create_fused_set_kv_buffer_arg(
|
||||||
value=v,
|
value=v,
|
||||||
layer=self.attn,
|
layer=self.attn,
|
||||||
@@ -318,7 +319,8 @@ class GptOssAttention(nn.Module):
|
|||||||
if enable_fused_set_kv_buffer(forward_batch)
|
if enable_fused_set_kv_buffer(forward_batch)
|
||||||
else None
|
else None
|
||||||
),
|
),
|
||||||
)
|
}
|
||||||
|
q, k = self.rotary_emb(positions, q, k, **extra_args)
|
||||||
inner_state = q, k, v, forward_batch
|
inner_state = q, k, v, forward_batch
|
||||||
return None, forward_batch, inner_state
|
return None, forward_batch, inner_state
|
||||||
|
|
||||||
@@ -490,6 +492,9 @@ class GptOssModel(nn.Module):
|
|||||||
self.vocab_size = config.vocab_size
|
self.vocab_size = config.vocab_size
|
||||||
self.pp_group = get_pp_group()
|
self.pp_group = get_pp_group()
|
||||||
|
|
||||||
|
if is_npu:
|
||||||
|
config.hidden_act = "npu_swiglu_oai"
|
||||||
|
|
||||||
if self.pp_group.is_first_rank:
|
if self.pp_group.is_first_rank:
|
||||||
self.embed_tokens = VocabParallelEmbedding(
|
self.embed_tokens = VocabParallelEmbedding(
|
||||||
config.vocab_size,
|
config.vocab_size,
|
||||||
|
|||||||
@@ -1263,7 +1263,7 @@ class ServerArgs:
|
|||||||
else:
|
else:
|
||||||
self.attention_backend = "triton"
|
self.attention_backend = "triton"
|
||||||
|
|
||||||
supported_backends = ["triton", "trtllm_mha", "fa3", "fa4"]
|
supported_backends = ["triton", "trtllm_mha", "fa3", "fa4", "ascend"]
|
||||||
prefill_attn_backend, decode_attn_backend = self.get_attention_backends()
|
prefill_attn_backend, decode_attn_backend = self.get_attention_backends()
|
||||||
assert (
|
assert (
|
||||||
prefill_attn_backend in supported_backends
|
prefill_attn_backend in supported_backends
|
||||||
|
|||||||
Reference in New Issue
Block a user