From 5d299c25c04a9e28170659b7588afcc0b6f84599 Mon Sep 17 00:00:00 2001 From: chenxu214 Date: Thu, 22 Jan 2026 22:02:05 +0800 Subject: [PATCH] [NPU] bugfix with Kimi-k2 and bge-reranker-v2 model (#17478) Co-authored-by: amote-i <49533125+amote-i@users.noreply.github.com> Co-authored-by: cy --- .../npu/attention/ascend_backend.py | 16 +-- .../compressed_tensors_moe.py | 2 +- python/sglang/srt/layers/rotary_embedding.py | 3 +- python/sglang/srt/models/bert.py | 8 +- .../test_ascend_bge_large_en_v1_5.py | 112 ++++++++++++++++++ 5 files changed, 130 insertions(+), 11 deletions(-) create mode 100644 test/registered/ascend/embedding_models/test_ascend_bge_large_en_v1_5.py diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py index 45aa241e5..1be763233 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py @@ -675,7 +675,14 @@ class AscendAttnBackend(AttentionBackend): ) else: - if layer.qk_head_dim <= 128: + causal = True + if ( + layer.is_cross_attention + or layer.attn_type == AttentionType.ENCODER_ONLY + ): + causal = False + + if layer.qk_head_dim <= 128 and causal: query = q.reshape(-1, layer.tp_q_head_num * layer.qk_head_dim) attn_output = torch.empty( (query.shape[0], layer.tp_q_head_num * layer.v_head_dim), @@ -709,13 +716,6 @@ class AscendAttnBackend(AttentionBackend): q_ = q.view(-1, layer.tp_q_head_num, layer.qk_head_dim) o_ = attn_output.view(-1, layer.tp_q_head_num, layer.v_head_dim) - causal = True - if ( - layer.is_cross_attention - or layer.attn_type == AttentionType.ENCODER_ONLY - ): - causal = False - self.native_attn._run_sdpa_forward_extend( q_, o_, diff --git a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors_moe.py b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors_moe.py index 4bfb7abf3..9e10dd837 100644 --- a/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors_moe.py +++ b/python/sglang/srt/layers/quantization/compressed_tensors/compressed_tensors_moe.py @@ -1648,7 +1648,7 @@ class NPUCompressedTensorsW4A16Int4DynamicMoEMethod(CompressedTensorsMoEMethod): self.num_experts = num_experts if ( extra_weight_attrs.get( - "intermediate_size_full", intermediate_size_per_partition + "moe_intermediate_size", intermediate_size_per_partition ) // intermediate_size_per_partition > 1 diff --git a/python/sglang/srt/layers/rotary_embedding.py b/python/sglang/srt/layers/rotary_embedding.py index 8fbdf3160..b50f2374a 100644 --- a/python/sglang/srt/layers/rotary_embedding.py +++ b/python/sglang/srt/layers/rotary_embedding.py @@ -115,9 +115,10 @@ class RotaryEmbedding(MultiPlatformOp): cache = cache.to(dtype) if ( - (not (_is_cuda or _is_npu) or self.head_size not in [64, 128, 256, 512]) + (not (_is_cuda) or self.head_size not in [64, 128, 256, 512]) and not (_is_cpu) and not (_is_xpu) + and not (_is_npu) ): if _is_cuda or _is_hip: from sgl_kernel import rotary_embedding diff --git a/python/sglang/srt/models/bert.py b/python/sglang/srt/models/bert.py index 45494423f..976b69ab8 100644 --- a/python/sglang/srt/models/bert.py +++ b/python/sglang/srt/models/bert.py @@ -17,6 +17,7 @@ from sglang.srt.layers.radix_attention import AttentionType, RadixAttention from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_loader.weight_utils import default_weight_loader +from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import add_prefix BertConfig = None @@ -365,10 +366,15 @@ class BertModel(nn.Module): quant_config=quant_config, prefix=add_prefix("encoder", prefix), ) + pooling_type = ( + PoolingType.CLS + if get_global_server_args().is_embedding + else PoolingType.LAST + ) self.pooler = ( BertPooler(config) if self.use_bert_pooler - else Pooler(pooling_type=PoolingType.LAST, normalize=True) + else Pooler(pooling_type=pooling_type, normalize=True) ) @torch.no_grad() diff --git a/test/registered/ascend/embedding_models/test_ascend_bge_large_en_v1_5.py b/test/registered/ascend/embedding_models/test_ascend_bge_large_en_v1_5.py new file mode 100644 index 000000000..8e2869d70 --- /dev/null +++ b/test/registered/ascend/embedding_models/test_ascend_bge_large_en_v1_5.py @@ -0,0 +1,112 @@ +import multiprocessing as mp +import unittest +from typing import Optional + +import torch +from transformers import AutoConfig, AutoTokenizer + +from sglang.test.ci.ci_register import register_npu_ci +from sglang.test.runners import HFRunner, SRTRunner +from sglang.test.test_utils import CustomTestCase, get_similarities + +register_npu_ci( + est_time=400, + suite="nightly-1-npu-a3", + nightly=True, + disabled="embeddings are not all close", +) + +DEFAULT_PROMPTS = [ + "The capital of the United Kingdom is", + "Today is a sunny day and I like", + "AI is a field of computer science focused on", +] + +MODELS = [ + ("/root/.cache/modelscope/hub/models/bge-large-en-v1.5", 1, 1e-5), +] +TORCH_DTYPES = [torch.float16] + + +class TestEmbeddingModels(CustomTestCase): + + @classmethod + def setUpClass(cls): + mp.set_start_method("spawn", force=True) + + def _truncate_prompts(self, prompts, model_path): + config = AutoConfig.from_pretrained(model_path) + max_length = getattr(config, "max_position_embeddings", 2048) + + tokenizer = AutoTokenizer.from_pretrained(model_path) + + truncated_prompts = [] + for prompt in prompts: + tokens = tokenizer(prompt, return_tensors="pt", truncation=False) + if len(tokens.input_ids[0]) > max_length: + truncated_text = tokenizer.decode( + tokens.input_ids[0][: max_length - 1], skip_special_tokens=True + ) + truncated_prompts.append(truncated_text) + else: + truncated_prompts.append(prompt) + return truncated_prompts + + def assert_close_prefill_logits( + self, + prompts, + model_path, + tp_size, + torch_dtype, + prefill_tolerance, + matryoshka_dim: Optional[int] = None, + ) -> None: + truncated_prompts = self._truncate_prompts(prompts, model_path) + + with HFRunner( + model_path, + torch_dtype=torch_dtype, + model_type="embedding", + matryoshka_dim=matryoshka_dim, + ) as hf_runner: + hf_outputs = hf_runner.forward(truncated_prompts) + + attention_backend = "ascend" + with SRTRunner( + model_path, + tp_size=tp_size, + torch_dtype=torch_dtype, + model_type="embedding", + attention_backend=attention_backend, + json_model_override_args=( + {"matryoshka_dimensions": [matryoshka_dim]} if matryoshka_dim else None + ), + ) as srt_runner: + srt_outputs = srt_runner.forward( + truncated_prompts, dimensions=matryoshka_dim + ) + + for i in range(len(prompts)): + hf_logits = torch.Tensor(hf_outputs.embed_logits[i]) + srt_logits = torch.Tensor(srt_outputs.embed_logits[i]) + + similarity = torch.tensor(get_similarities(hf_logits, srt_logits)) + print("similarity diff", abs(similarity - 1)) + + if len(prompts[i]) <= 1000: + assert torch.all( + abs(similarity - 1) < prefill_tolerance + ), "embeddings are not all close" + + def test_prefill_logits(self): + models_to_test = MODELS + + for model, tp_size, prefill_tolerance in models_to_test: + for torch_dtype in TORCH_DTYPES: + self.assert_close_prefill_logits( + DEFAULT_PROMPTS, model, tp_size, torch_dtype, prefill_tolerance + ) + + +if __name__ == "__main__": + unittest.main()