[NPU] support model skywork-reward-gemma2-2-27B-v0.2 (#16947)

Co-authored-by: cy <chenyang08056032@163.com>
This commit is contained in:
McZyWu
2026-02-11 15:34:53 +08:00
committed by GitHub
parent 72c1526657
commit 4f7422f7ba
7 changed files with 135 additions and 10 deletions

View File

@@ -1163,6 +1163,7 @@ def is_generation_model(model_architectures: List[str], is_embedding: bool = Fal
or "BertForSequenceClassification" in model_architectures
or "XLMRobertaModel" in model_architectures
or "XLMRobertaForSequenceClassification" in model_architectures
or "Gemma2ForSequenceClassification" in model_architectures
):
return False
else:

View File

@@ -294,6 +294,10 @@ class Envs:
SGLANG_NPU_DISABLE_ACL_FORMAT_WEIGHT = EnvBool(False)
SGLANG_NPU_USE_MULTI_STREAM = EnvBool(False)
SGLANG_NPU_USE_MLAPO = EnvBool(False)
# Forward native implementation for activation gelu tanh for model Skywork-Reward-Gemma-2-27B-v0.2
SGLANG_NPU_FORWARD_NATIVE_GELUTANH = EnvBool(False)
# Forward native implementation for gemma rms norm for model Skywork-Reward-Gemma-2-27B-v0.2
SGLANG_NPU_FORWARD_NATIVE_GEMMA_RMS_NORM = EnvBool(False)
# Quantization
SGLANG_INT4_WEIGHT = EnvBool(False)

View File

@@ -225,6 +225,11 @@ class AscendAttnBackend(AttentionBackend):
self.q_head_dim = self.qk_rope_head_dim + self.qk_nope_head_dim
else:
self.use_alibi = getattr(model_runner.model_config, "use_alibi", False)
if (
"Gemma2ForSequenceClassification"
in model_runner.model_config.hf_config.architectures
):
self.use_native_sdpa = True
self.native_attn = AscendTorchNativeAttnBackend()
self.graph_metadata = {}
self.max_context_len = model_runner.model_config.context_len
@@ -821,10 +826,12 @@ class AscendAttnBackend(AttentionBackend):
# there are some accuracy issues in cross attention scene to use torch_npu._npu_flash_attention_qlens
# forward_batch.encoder_lens is not None in cross attention scend, we add native attn to solve accuracy issues
# Model skywork-reward-gemma2-2-27B also suffers from precision anomalies, thus the torch native backend becomes beneficial approach.
if (
layer.qk_head_dim <= 128
and causal
and forward_batch.encoder_lens is None
and not getattr(self, "use_native_sdpa", False)
):
if not self.use_alibi:
query = q.reshape(-1, layer.tp_q_head_num * layer.qk_head_dim)

View File

@@ -27,6 +27,7 @@ from sglang.srt.distributed import (
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from sglang.srt.environ import envs
from sglang.srt.layers.quantization.base_config import QuantizationConfig
from sglang.srt.layers.utils import MultiPlatformOp
from sglang.srt.server_args import get_global_server_args
@@ -131,6 +132,8 @@ class GeluAndMul(MultiPlatformOp):
return self._forward_impl(x)
def forward_npu(self, x: torch.Tensor) -> torch.Tensor:
if envs.SGLANG_NPU_FORWARD_NATIVE_GELUTANH.get():
return self.forward_native(x)
y_npu, gelu_npu = torch_npu.npu_geglu(
x,
dim=-1,

View File

@@ -24,6 +24,7 @@ from sglang.srt.batch_invariant_ops import (
is_batch_invariant_mode_enabled,
rms_norm_batch_invariant,
)
from sglang.srt.environ import envs
from sglang.srt.layers.utils import MultiPlatformOp
from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
@@ -468,6 +469,8 @@ class GemmaRMSNorm(MultiPlatformOp):
residual: Optional[torch.Tensor] = None,
post_residual_addition: Optional[torch.Tensor] = None,
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
if envs.SGLANG_NPU_FORWARD_NATIVE_GEMMA_RMS_NORM.get():
return self.forward_native(x, residual)
if residual is not None:
if post_residual_addition is not None:
residual = residual + post_residual_addition

View File

@@ -39,7 +39,9 @@ from sglang.srt.model_loader.weight_utils import (
default_weight_loader,
maybe_remap_kv_scale_name,
)
from sglang.srt.utils import add_prefix, make_layers
from sglang.srt.utils import add_prefix, is_npu, make_layers
_is_npu = is_npu()
# Aligned with HF's implementation, using sliding window inclusive with the last token
@@ -142,13 +144,28 @@ class Gemma2Attention(nn.Module):
quant_config=quant_config,
prefix=add_prefix("o_proj", prefix),
)
self.rotary_emb = get_rope(
self.head_dim,
rotary_dim=self.head_dim,
max_position=max_position_embeddings,
base=self.rope_theta,
is_neox_style=True,
)
if (
not _is_npu
or "Gemma2ForSequenceClassification" not in self.config.architectures
):
self.rotary_emb = get_rope(
self.head_dim,
rotary_dim=self.head_dim,
max_position=max_position_embeddings,
base=self.rope_theta,
is_neox_style=True,
)
logit_cap = self.config.attn_logit_softcapping
else:
self.rotary_emb = get_rope(
self.head_dim,
rotary_dim=self.head_dim,
max_position=max_position_embeddings,
base=self.rope_theta,
is_neox_style=True,
dtype=torch.float32,
)
logit_cap = 0.0
use_sliding_window = layer_id % 2 == 0 and hasattr(config, "sliding_window")
self.attn = RadixAttention(
@@ -157,7 +174,7 @@ class Gemma2Attention(nn.Module):
self.scaling,
num_kv_heads=self.num_kv_heads,
layer_id=layer_id,
logit_cap=self.config.attn_logit_softcapping,
logit_cap=logit_cap,
sliding_window_size=(
get_attention_sliding_window_size(config)
if use_sliding_window
@@ -294,7 +311,9 @@ class Gemma2Model(nn.Module):
hidden_states = self.embed_tokens(input_ids)
else:
hidden_states = input_embeds
normalizer = torch.tensor(self.config.hidden_size**0.5, dtype=torch.float16)
normalizer = torch.tensor(
self.config.hidden_size**0.5, dtype=hidden_states.dtype
)
hidden_states *= normalizer
residual = None