【NPU】【bugfix】fix server error when mtp unquant (#26389)

Co-authored-by: cen121212 <luochen23@huawei.com>
Co-authored-by: Even Zhou <even.y.zhou@outlook.com>
This commit is contained in:
cen121212
2026-05-30 15:01:19 +03:00
committed by GitHub
co-authored by cen121212 Even Zhou
parent 282c46133f
commit b421e60eed
5 changed files with 148 additions and 98 deletions
@@ -1,7 +1,6 @@
from __future__ import annotations from __future__ import annotations
import logging import logging
import os
from contextlib import nullcontext from contextlib import nullcontext
from dataclasses import dataclass from dataclasses import dataclass
from typing import TYPE_CHECKING, List, NamedTuple, Optional, Tuple, Union from typing import TYPE_CHECKING, List, NamedTuple, Optional, Tuple, Union
@@ -437,11 +436,8 @@ class _DeepEPDispatcherImplBase:
# NVFP4 is supported on GPU, no adjustment needed # NVFP4 is supported on GPU, no adjustment needed
def _update_int8_quant_env(self) -> None: def _update_int8_quant_env(self) -> None:
"""Update the DEEP_NORMAL_MODE_USE_INT8_QUANT environment variable.""" """TODO adapt different quantization schemes for base model and draft model on NPU"""
if self.use_fp8: pass
os.environ["DEEP_NORMAL_MODE_USE_INT8_QUANT"] = "1"
else:
os.environ["DEEP_NORMAL_MODE_USE_INT8_QUANT"] = "0"
def set_overlap_args( def set_overlap_args(
self, combine_overlap_args: CombineOverlapArgs, meta_overlap_args: dict self, combine_overlap_args: CombineOverlapArgs, meta_overlap_args: dict
+19 -1
View File
@@ -16,6 +16,7 @@
import logging import logging
import os import os
from contextlib import ExitStack
from typing import Iterable, Optional, Tuple from typing import Iterable, Optional, Tuple
import torch import torch
@@ -169,11 +170,26 @@ class DeepseekModelNextN(nn.Module):
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
input_embeds: torch.Tensor = None, input_embeds: torch.Tensor = None,
) -> torch.Tensor: ) -> torch.Tensor:
exit_stack = ExitStack()
if (
_is_npu
and self.quant_config is None
and get_global_server_args().quantization is not None
):
# ascend mtp unquant
exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True))
exit_stack.enter_context(
envs.DEEP_NORMAL_MODE_USE_INT8_QUANT.override(False)
)
try:
zero_allocator = BumpAllocator( zero_allocator = BumpAllocator(
buffer_size=2, buffer_size=2,
dtype=torch.float32, dtype=torch.float32,
device=( device=(
input_embeds.device if input_embeds is not None else input_ids.device input_embeds.device
if input_embeds is not None
else input_ids.device
), ),
) )
@@ -232,6 +248,8 @@ class DeepseekModelNextN(nn.Module):
forward_batch, forward_batch,
torch.cuda.current_stream(), torch.cuda.current_stream(),
) )
finally:
exit_stack.close()
return hidden_states return hidden_states
+20 -1
View File
@@ -16,6 +16,7 @@
import copy import copy
import logging import logging
from contextlib import ExitStack
from typing import Iterable, Optional, Tuple from typing import Iterable, Optional, Tuple
import torch import torch
@@ -23,6 +24,7 @@ from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size
from sglang.srt.environ import envs
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
from sglang.srt.layers.layernorm import GemmaRMSNorm from sglang.srt.layers.layernorm import GemmaRMSNorm
@@ -140,7 +142,19 @@ class Qwen3_5ForCausalLMMTP(nn.Module):
input_embeds: Optional[torch.Tensor] = None, input_embeds: Optional[torch.Tensor] = None,
**kwargs, **kwargs,
): ):
exit_stack = ExitStack()
if (
is_npu()
and self.quant_config is None
and get_global_server_args().quantization is not None
):
# ascend mtp unquant
exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True))
exit_stack.enter_context(
envs.DEEP_NORMAL_MODE_USE_INT8_QUANT.override(False)
)
try:
assert input_embeds is None assert input_embeds is None
input_embeds = forward_batch.mm_input_embeds input_embeds = forward_batch.mm_input_embeds
if ( if (
@@ -150,7 +164,10 @@ class Qwen3_5ForCausalLMMTP(nn.Module):
): ):
assert input_embeds is not None assert input_embeds is not None
input_embeds = torch.cat( input_embeds = torch.cat(
[input_embeds[:-1], self.model.embed_tokens(input_ids[-1].unsqueeze(0))] [
input_embeds[:-1],
self.model.embed_tokens(input_ids[-1].unsqueeze(0)),
]
) )
if input_embeds is None: if input_embeds is None:
@@ -172,6 +189,8 @@ class Qwen3_5ForCausalLMMTP(nn.Module):
forward_batch, forward_batch,
hidden_states, hidden_states,
) )
finally:
exit_stack.close()
return self.logits_processor( return self.logits_processor(
input_ids, hidden_states, self.lm_head, forward_batch input_ids, hidden_states, self.lm_head, forward_batch
@@ -16,6 +16,7 @@
import copy import copy
import logging import logging
from contextlib import ExitStack
from typing import Iterable, Optional, Tuple from typing import Iterable, Optional, Tuple
import torch import torch
@@ -23,6 +24,7 @@ from torch import nn
from transformers import PretrainedConfig from transformers import PretrainedConfig
from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size
from sglang.srt.environ import envs
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.layers.layernorm import GemmaRMSNorm from sglang.srt.layers.layernorm import GemmaRMSNorm
from sglang.srt.layers.logits_processor import LogitsProcessor from sglang.srt.layers.logits_processor import LogitsProcessor
@@ -94,7 +96,19 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM):
input_embeds: Optional[torch.Tensor] = None, input_embeds: Optional[torch.Tensor] = None,
**kwargs, **kwargs,
): ):
exit_stack = ExitStack()
if (
is_npu()
and self.quant_config is None
and get_global_server_args().quantization is not None
):
# ascend mtp unquant
exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True))
exit_stack.enter_context(
envs.DEEP_NORMAL_MODE_USE_INT8_QUANT.override(False)
)
try:
if input_embeds is None: if input_embeds is None:
input_embeds = self.model.embed_tokens(input_ids) input_embeds = self.model.embed_tokens(input_ids)
@@ -112,6 +126,8 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM):
forward_batch, forward_batch,
hidden_states, hidden_states,
) )
finally:
exit_stack.close()
return self.logits_processor( return self.logits_processor(
input_ids, hidden_states, self.lm_head, forward_batch input_ids, hidden_states, self.lm_head, forward_batch
@@ -60,6 +60,7 @@ class TestAscendDeepEP(CustomTestCase):
"SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "32", "SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "32",
"SGLANG_NPU_USE_MLAPO": "1", "SGLANG_NPU_USE_MLAPO": "1",
"TRANSFORMERS_VERBOSITY": "error", "TRANSFORMERS_VERBOSITY": "error",
"DEEP_NORMAL_MODE_USE_INT8_QUANT": "1",
} }
os.environ.update(cls.extra_envs) os.environ.update(cls.extra_envs)