[NPU] Determine the topk norm_type through scoring_func (#31393)
This commit is contained in:
@@ -98,7 +98,7 @@ def fused_topk_npu(
|
|||||||
group_select_mode=(1 if use_grouped_topk else 0),
|
group_select_mode=(1 if use_grouped_topk else 0),
|
||||||
renorm=0,
|
renorm=0,
|
||||||
# 1 for sigmoid, 0 for softmax
|
# 1 for sigmoid, 0 for softmax
|
||||||
norm_type=1,
|
norm_type=(0 if topk_config.scoring_func == "softmax" else 1),
|
||||||
routed_scaling_factor=(
|
routed_scaling_factor=(
|
||||||
topk_config.routed_scaling_factor
|
topk_config.routed_scaling_factor
|
||||||
if topk_config.apply_routed_scaling_factor_on_output
|
if topk_config.apply_routed_scaling_factor_on_output
|
||||||
|
|||||||
@@ -85,12 +85,14 @@ from sglang.srt.utils import (
|
|||||||
LazyValue,
|
LazyValue,
|
||||||
add_prefix,
|
add_prefix,
|
||||||
is_non_idle_and_non_empty,
|
is_non_idle_and_non_empty,
|
||||||
|
is_npu,
|
||||||
log_info_on_rank0,
|
log_info_on_rank0,
|
||||||
make_layers,
|
make_layers,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.hf_transformers_utils import get_rope_config
|
from sglang.srt.utils.hf_transformers_utils import get_rope_config
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
_is_npu = is_npu()
|
||||||
|
|
||||||
|
|
||||||
class Glm4MoeLiteMLP(nn.Module):
|
class Glm4MoeLiteMLP(nn.Module):
|
||||||
@@ -242,6 +244,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module):
|
|||||||
# Some Fp4 MoE backends require the output format to be bypassed but the MTP layers are unquantized
|
# Some Fp4 MoE backends require the output format to be bypassed but the MTP layers are unquantized
|
||||||
# and requires the output format to be standard. We use quant_config to determine the output format.
|
# and requires the output format to be standard. We use quant_config to determine the output format.
|
||||||
output_format=TopKOutputFormat.STANDARD if quant_config is None else None,
|
output_format=TopKOutputFormat.STANDARD if quant_config is None else None,
|
||||||
|
**({"scoring_func": "sigmoid"} if _is_npu else {}),
|
||||||
)
|
)
|
||||||
|
|
||||||
self.shared_experts_is_int8 = False
|
self.shared_experts_is_int8 = False
|
||||||
|
|||||||
Reference in New Issue
Block a user