From f2b84b90ac86b19d6ea366b9ae15ce0e0779f875 Mon Sep 17 00:00:00 2001 From: ranjiewen Date: Mon, 27 Apr 2026 18:14:33 +0800 Subject: [PATCH] [npu]fix: qwen3-next w8a8 precision bugs (#21698) --- python/sglang/srt/models/qwen3_next.py | 27 +++++++++++++++++++++----- 1 file changed, 22 insertions(+), 5 deletions(-) diff --git a/python/sglang/srt/models/qwen3_next.py b/python/sglang/srt/models/qwen3_next.py index d550be33e..ccde50d53 100644 --- a/python/sglang/srt/models/qwen3_next.py +++ b/python/sglang/srt/models/qwen3_next.py @@ -65,6 +65,14 @@ _is_cpu = is_cpu() _is_amx_available = cpu_has_amx_support() +if _is_npu: + from sgl_kernel_npu.fla.utils import ( + fused_qkvzba_split_reshape_cat as fused_qkvzba_split_reshape_cat_npu, + ) + + fused_qkvzba_split_reshape_cat = fused_qkvzba_split_reshape_cat_npu + + class Qwen3GatedDeltaNet(nn.Module): def __init__( self, @@ -223,11 +231,20 @@ class Qwen3GatedDeltaNet(nn.Module): ModelWeightParameter exposes weight_loader as a read-only property backed by _weight_loader, while plain parameters store it as a regular attribute. This helper handles both cases.""" - param = module.weight - if hasattr(param, "_weight_loader"): - param._weight_loader = new_loader - else: - param.weight_loader = new_loader + for attr_name in ( + "weight", + "weight_scale_inv", + "weight_scale", + "input_scale", + "weight_offset", + ): + param = getattr(module, attr_name, None) + if param is None: + continue + if hasattr(param, "_weight_loader"): + param._weight_loader = new_loader + else: + param.weight_loader = new_loader @staticmethod def _make_packed_weight_loader(module):