diff --git a/python/sglang/srt/model_loader/loader.py b/python/sglang/srt/model_loader/loader.py index 26aade419..8bfd57d33 100644 --- a/python/sglang/srt/model_loader/loader.py +++ b/python/sglang/srt/model_loader/loader.py @@ -78,7 +78,6 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import ( ) from sglang.srt.model_loader.utils import ( get_model_architecture, - post_load_weights, set_default_torch_dtype, ) @@ -286,6 +285,15 @@ def _initialize_model( return model_class(**kwargs) +def _post_load_weights(model: nn.Module) -> None: + # Loaders that bypass `model.load_weights()` (dummy / sharded state / remote instance / + # remote fs) must trigger the model's post-load fixup explicitly; `model.load_weights()` + # would normally do it internally. NextN subclasses override the method to fill in + # `is_nextn=True`, so the loader doesn't need to know. + if hasattr(model, "post_load_weights"): + model.post_load_weights() + + class BaseModelLoader(ABC): """Base class for model loaders.""" @@ -1322,7 +1330,7 @@ class DummyModelLoader(BaseModelLoader): # random values to the weights. initialize_dummy_weights(model) - post_load_weights(model, model_config) + _post_load_weights(model) return model.eval() @@ -1465,7 +1473,7 @@ class ShardedStateLoader(BaseModelLoader): if state_dict: raise ValueError(f"Missing keys {tuple(state_dict)} in loaded state!") - post_load_weights(model, model_config) + _post_load_weights(model) return model.eval() @@ -2228,8 +2236,7 @@ class RemoteInstanceModelLoader(BaseModelLoader): ) torch.cuda.synchronize() - if hasattr(model, "post_load_weights"): - model.post_load_weights() + _post_load_weights(model) end_get_weights_tic = time.time() logger.debug( f"finish getting all weights from remote instance, time used: {(end_get_weights_tic - start_get_weights_tic):.4f}s" @@ -2292,8 +2299,7 @@ class RemoteInstanceModelLoader(BaseModelLoader): logger.error(f"batch transfer failed, error: {ret}") return False - if hasattr(model, "post_load_weights"): - model.post_load_weights() + _post_load_weights(model) return True @@ -2414,8 +2420,7 @@ class RemoteInstanceModelLoader(BaseModelLoader): model, transfer_engine, source_worker, tp_rank ) - if hasattr(model, "post_load_weights"): - model.post_load_weights() + _post_load_weights(model) logger.info("ModelExpress: weight transfer complete for tp_rank=%d", tp_rank) @@ -2645,7 +2650,7 @@ class RemoteModelLoader(BaseModelLoader): if state_dict: raise ValueError(f"Missing keys {tuple(state_dict)} in loaded state!") - post_load_weights(model, model_config) + _post_load_weights(model) def _load_model_from_remote_fs( self, model, client, model_config: ModelConfig, device_config: DeviceConfig diff --git a/python/sglang/srt/model_loader/utils.py b/python/sglang/srt/model_loader/utils.py index 832ff1424..df708f342 100644 --- a/python/sglang/srt/model_loader/utils.py +++ b/python/sglang/srt/model_loader/utils.py @@ -247,18 +247,6 @@ def get_architecture_class_name(model_config: ModelConfig) -> str: return get_model_architecture(model_config)[1] -def post_load_weights(model: nn.Module, model_config: ModelConfig): - # Model weight loading consists of two stages: - # 1. Initial weight loading. - # 2. Post-processing of weights, including assigning specific member variables. - # For `dummy_init`, only the second stage is required. - if hasattr(model, "post_load_weights"): - if model_config.hf_config.architectures[0] == "DeepseekV3ForCausalLMNextN": - model.post_load_weights(is_nextn=True) - else: - model.post_load_weights() - - def should_deepgemm_weight_requant_ue8m0(weight_block_size): """Should we requant fp8 weights into UE8M0 format when loading the model""" return ( diff --git a/python/sglang/srt/models/bailing_moe_nextn.py b/python/sglang/srt/models/bailing_moe_nextn.py index fcbae6bba..2c392a5ae 100644 --- a/python/sglang/srt/models/bailing_moe_nextn.py +++ b/python/sglang/srt/models/bailing_moe_nextn.py @@ -217,7 +217,8 @@ class BailingMoeForCausalLMNextN(nn.Module): self.post_load_weights_func = BailingMoeV2_5ForCausalLM.post_load_weights else: self.base_load_weights_func = BailingMoEForCausalLM.load_weights - self.post_load_weights_func = BailingMoEForCausalLM.post_load_weights + # V1 BailingMoeAttention is standard QKV (no kv_b_proj), no fixup needed. + self.post_load_weights_func = None @torch.no_grad() def forward( @@ -243,8 +244,13 @@ class BailingMoeForCausalLMNextN(nn.Module): def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): self.base_load_weights_func(self, weights, is_nextn=True) - def post_load_weights(self, is_nextn=False, weight_names=None): - self.post_load_weights_func(self, is_nextn=is_nextn, weight_names=weight_names) + def post_load_weights(self, is_nextn=True, weight_names=None): + # `is_nextn` is pinned to True for the NextN subclass; the parameter is kept + # only because the underlying `load_weights` flow calls `self.post_load_weights` + # with `is_nextn=...` as a kwarg. + if self.post_load_weights_func is None: + return + self.post_load_weights_func(self, is_nextn=True, weight_names=weight_names) EntryClass = [BailingMoeForCausalLMNextN] diff --git a/python/sglang/srt/models/deepseek_nextn.py b/python/sglang/srt/models/deepseek_nextn.py index aa7f02274..2196a9b26 100644 --- a/python/sglang/srt/models/deepseek_nextn.py +++ b/python/sglang/srt/models/deepseek_nextn.py @@ -313,5 +313,11 @@ class DeepseekV3ForCausalLMNextN(DeepseekV3ForCausalLM): def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): super().load_weights(weights, is_nextn=True) + def post_load_weights(self, is_nextn=True, weight_names=None): + # `is_nextn` is pinned to True for the NextN subclass; the parameter is kept + # only because the mixin's `do_load_weights` calls `self.post_load_weights` + # with `is_nextn=...` as a kwarg. + super().post_load_weights(is_nextn=True, weight_names=weight_names) + EntryClass = [DeepseekV3ForCausalLMNextN]