【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:
co-authored by
cen121212
Even Zhou
parent
282c46133f
commit
b421e60eed
@@ -1,7 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
from contextlib import nullcontext
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, List, NamedTuple, Optional, Tuple, Union
|
||||
@@ -437,11 +436,8 @@ class _DeepEPDispatcherImplBase:
|
||||
# NVFP4 is supported on GPU, no adjustment needed
|
||||
|
||||
def _update_int8_quant_env(self) -> None:
|
||||
"""Update the DEEP_NORMAL_MODE_USE_INT8_QUANT environment variable."""
|
||||
if self.use_fp8:
|
||||
os.environ["DEEP_NORMAL_MODE_USE_INT8_QUANT"] = "1"
|
||||
else:
|
||||
os.environ["DEEP_NORMAL_MODE_USE_INT8_QUANT"] = "0"
|
||||
"""TODO adapt different quantization schemes for base model and draft model on NPU"""
|
||||
pass
|
||||
|
||||
def set_overlap_args(
|
||||
self, combine_overlap_args: CombineOverlapArgs, meta_overlap_args: dict
|
||||
|
||||
@@ -16,6 +16,7 @@
|
||||
|
||||
import logging
|
||||
import os
|
||||
from contextlib import ExitStack
|
||||
from typing import Iterable, Optional, Tuple
|
||||
|
||||
import torch
|
||||
@@ -169,11 +170,26 @@ class DeepseekModelNextN(nn.Module):
|
||||
forward_batch: ForwardBatch,
|
||||
input_embeds: torch.Tensor = None,
|
||||
) -> 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(
|
||||
buffer_size=2,
|
||||
dtype=torch.float32,
|
||||
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,
|
||||
torch.cuda.current_stream(),
|
||||
)
|
||||
finally:
|
||||
exit_stack.close()
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
@@ -16,6 +16,7 @@
|
||||
|
||||
import copy
|
||||
import logging
|
||||
from contextlib import ExitStack
|
||||
from typing import Iterable, Optional, Tuple
|
||||
|
||||
import torch
|
||||
@@ -23,6 +24,7 @@ from torch import nn
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
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_location import ModelConfigForExpertLocation
|
||||
from sglang.srt.layers.layernorm import GemmaRMSNorm
|
||||
@@ -140,7 +142,19 @@ class Qwen3_5ForCausalLMMTP(nn.Module):
|
||||
input_embeds: Optional[torch.Tensor] = None,
|
||||
**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
|
||||
input_embeds = forward_batch.mm_input_embeds
|
||||
if (
|
||||
@@ -150,7 +164,10 @@ class Qwen3_5ForCausalLMMTP(nn.Module):
|
||||
):
|
||||
assert input_embeds is not None
|
||||
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:
|
||||
@@ -172,6 +189,8 @@ class Qwen3_5ForCausalLMMTP(nn.Module):
|
||||
forward_batch,
|
||||
hidden_states,
|
||||
)
|
||||
finally:
|
||||
exit_stack.close()
|
||||
|
||||
return self.logits_processor(
|
||||
input_ids, hidden_states, self.lm_head, forward_batch
|
||||
|
||||
@@ -16,6 +16,7 @@
|
||||
|
||||
import copy
|
||||
import logging
|
||||
from contextlib import ExitStack
|
||||
from typing import Iterable, Optional, Tuple
|
||||
|
||||
import torch
|
||||
@@ -23,6 +24,7 @@ from torch import nn
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
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.layers.layernorm import GemmaRMSNorm
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||
@@ -94,7 +96,19 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM):
|
||||
input_embeds: Optional[torch.Tensor] = None,
|
||||
**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:
|
||||
input_embeds = self.model.embed_tokens(input_ids)
|
||||
|
||||
@@ -112,6 +126,8 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM):
|
||||
forward_batch,
|
||||
hidden_states,
|
||||
)
|
||||
finally:
|
||||
exit_stack.close()
|
||||
|
||||
return self.logits_processor(
|
||||
input_ids, hidden_states, self.lm_head, forward_batch
|
||||
|
||||
+1
@@ -60,6 +60,7 @@ class TestAscendDeepEP(CustomTestCase):
|
||||
"SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK": "32",
|
||||
"SGLANG_NPU_USE_MLAPO": "1",
|
||||
"TRANSFORMERS_VERBOSITY": "error",
|
||||
"DEEP_NORMAL_MODE_USE_INT8_QUANT": "1",
|
||||
}
|
||||
os.environ.update(cls.extra_envs)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user