4747from sglang .srt .managers .expert_distribution import ExpertDistributionRecorder
4848from sglang .srt .model_executor .forward_batch_info import ForwardBatch
4949from sglang .srt .model_loader .weight_utils import default_weight_loader
50- from sglang .srt .utils import add_prefix
50+ from sglang .srt .utils import add_prefix , make_layers
5151
5252expert_distribution_recorder = ExpertDistributionRecorder ()
5353
@@ -333,6 +333,7 @@ def __init__(
333333 config : PretrainedConfig ,
334334 quant_config : Optional [QuantizationConfig ] = None ,
335335 prefix : str = "" ,
336+ decoder_layer_type : type [nn .Module ] = Qwen2MoeDecoderLayer ,
336337 ) -> None :
337338 super ().__init__ ()
338339 self .padding_idx = config .pad_token_id
@@ -343,16 +344,17 @@ def __init__(
343344 config .hidden_size ,
344345 prefix = add_prefix ("embed_tokens" , prefix ),
345346 )
346- self .layers = nn .ModuleList (
347- [
348- Qwen2MoeDecoderLayer (
349- config ,
350- layer_id ,
351- quant_config = quant_config ,
352- prefix = add_prefix (f"layers.{ layer_id } " , prefix ),
353- )
354- for layer_id in range (config .num_hidden_layers )
355- ]
347+ # Use the provided decoder layer type or default to Qwen2MoeDecoderLayer
348+ decoder_layer_type = decoder_layer_type or Qwen2MoeDecoderLayer
349+ self .layers = make_layers (
350+ config .num_hidden_layers ,
351+ lambda idx , prefix : decoder_layer_type (
352+ layer_id = idx ,
353+ config = config ,
354+ quant_config = quant_config ,
355+ prefix = prefix ,
356+ ),
357+ prefix = add_prefix ("layers" , prefix ),
356358 )
357359 self .norm = RMSNorm (config .hidden_size , eps = config .rms_norm_eps )
358360
0 commit comments