diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index c0ce9c629..8ffc8323d 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -68,7 +68,6 @@ from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.layers.linear import ColumnParallelLinear, RowParallelLinear from sglang.srt.layers.logits_processor import LogitsProcessor -from sglang.srt.layers.mhc import mhc_fused_post_pre, npu_hc_pre from sglang.srt.layers.moe import get_moe_a2a_backend, should_use_dp_reduce_scatterv from sglang.srt.layers.moe.fused_moe_triton import FusedMoE from sglang.srt.layers.quantization.fp8_kernel import sglang_per_token_group_quant_fp8 @@ -113,9 +112,12 @@ from sglang.srt.models.deepseek_common.amd.deepseek_v4_fused_mhc import ( try_fused_hc_post_pre, ) from sglang.srt.models.deepseek_common.utils import _use_aiter_bpreshuffle_gfx95 -from sglang.srt.models.deepseek_v2 import ParallelLMHead, _is_cuda, _is_hip, _is_npu -from sglang.srt.models.triton_ops.deepseek_v4 import ( - rms_normalize_triton as rms_normalize_triton, +from sglang.srt.models.deepseek_v2 import ( + ParallelLMHead, + _is_cuda, + _is_hip, + _is_npu, + _is_xpu, ) from sglang.srt.runtime_context import get_parallel @@ -124,6 +126,11 @@ if not _is_hip: prepare_context_parallel_metadata, ) +if _is_xpu: + from sgl_kernel import hc_split_sinkhorn +else: + from sglang.srt.layers.mhc import hc_split_sinkhorn, mhc_fused_post_pre, npu_hc_pre + from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import ( LazyValue, @@ -1391,8 +1398,6 @@ class DeepseekV4DecoderLayer(nn.Module): else: x_flat, mixes = hc_pre_torch_impl(x, hc_fn) - from sglang.srt.layers.mhc import hc_split_sinkhorn - pre, post, comb = hc_split_sinkhorn( mixes, hc_scale,