@@ -46,7 +46,6 @@ def _load_checkpoint(queue, args):
4646 from megatron .training .checkpointing import load_args_from_checkpoint , load_checkpoint
4747 from megatron .legacy .model import module
4848 from megatron .core import mpu
49- from megatron .legacy import fused_kernels
5049 from megatron .core .tensor_parallel .random import (
5150 get_cuda_rng_tracker , _DATA_PARALLEL_RNG_TRACKER_NAME ,
5251 _EXPERT_PARALLEL_RNG_TRACKER_NAME , _MODEL_PARALLEL_RNG_TRACKER_NAME
@@ -78,7 +77,6 @@ def queue_put(name, msg):
7877 '--no-masked-softmax-fusion' ,
7978 '--no-bias-gelu-fusion' ,
8079 '--no-bias-dropout-fusion' ,
81- '--no-async-tensor-model-parallel-allreduce' ,
8280 '--use-cpu-initialization' ,
8381 '--micro-batch-size' , '1' ,
8482 '--no-load-optim' ,
@@ -125,20 +123,27 @@ def _set_arg(arg_name):
125123 _set_arg ("hetero_pipeline_layer_split" )
126124
127125 # for engram
128- _set_arg ("use_engram" )
129- _set_arg ("engram_layer_ids" )
130- _set_arg ("engram_hc_mult" )
131- _set_arg ("engram_kernel_size" )
132- _set_arg ("engram_pad_id" )
133- _set_arg ("engram_seed" )
134- _set_arg ("engram_vocab_size" )
135- _set_arg ("engram_tokenizer_name_or_path" )
136- _set_arg ("max_ngram_size" )
137- _set_arg ("n_embed_per_ngram" )
138- _set_arg ("n_head_per_ngram" )
139- setattr (margs , "vocab_size" , args .true_vocab_size )
140- engram_tokenizer_path_ckpt = getattr (checkpoint_args , "engram_tokenizer_name_or_path" , None )
141- setattr (margs , "engram_tokenizer_name_or_path" , os .path .join (root_path , engram_tokenizer_path_ckpt ))
126+ if getattr (checkpoint_args , "use_engram" , False ):
127+ _set_arg ("use_engram" )
128+ _set_arg ("engram_layer_ids" )
129+ _set_arg ("engram_hc_mult" )
130+ _set_arg ("engram_kernel_size" )
131+ _set_arg ("engram_pad_id" )
132+ _set_arg ("engram_seed" )
133+ _set_arg ("engram_vocab_size" )
134+ _set_arg ("engram_tokenizer_name_or_path" )
135+ _set_arg ("max_ngram_size" )
136+ _set_arg ("n_embed_per_ngram" )
137+ _set_arg ("n_head_per_ngram" )
138+ if args .true_vocab_size is not None :
139+ setattr (margs , "vocab_size" , args .true_vocab_size )
140+ engram_tokenizer_path_ckpt = getattr (checkpoint_args , "engram_tokenizer_name_or_path" , None )
141+ if engram_tokenizer_path_ckpt and not os .path .isabs (engram_tokenizer_path_ckpt ):
142+ engram_tokenizer_path_ckpt = os .path .join (root_path , engram_tokenizer_path_ckpt )
143+ setattr (margs , "engram_tokenizer_name_or_path" , engram_tokenizer_path_ckpt )
144+ else :
145+ setattr (margs , "use_engram" , False )
146+ setattr (margs , "engram_layer_ids" , [])
142147
143148 # for hetero
144149 if margs .hetero_process_meshes is not None :
@@ -259,9 +264,6 @@ def check_for_arg(arg_name, default=None):
259264 mpu ._INTRA_PARTIAL_DATA_PARALLEL_GROUP_WITH_CP = fake_dp_group
260265 mpu ._LAST_RANK_WHEN_USING_PIPELINE = pp_size - 1
261266
262- # fused kernel
263- fused_kernels .load (margs )
264-
265267 # random
266268 CUDA_RNG_STATE_TRACKER = get_cuda_rng_tracker ()
267269 torch .cuda .manual_seed (42 )
@@ -387,7 +389,9 @@ def get_models(count, dtype):
387389 margs .total_layer_num = total_layer_num
388390
389391 engram_layer_id = total_layer_num # get_global_layer_id
390- if margs .use_engram and engram_layer_id in margs .engram_layer_ids :
392+ if getattr (margs , "use_engram" , False ) and engram_layer_id in getattr (
393+ margs , "engram_layer_ids" , []
394+ ):
391395 ckpt_plugin .get_engram_ckpt (message , models , engram_layer_id , margs )
392396
393397 ckpt_plugin .get_attn_ckpt (message , models , layer_id , margs )
@@ -406,9 +410,11 @@ def get_models(count, dtype):
406410 ckpt_plugin .get_output_layer_ckpt (message , models , margs )
407411 queue_put ("output layer" , message )
408412
409- message = dict ()
410- if margs .mtp_num_layers :
411- for mtp_layer_id in range (margs .mtp_num_layers ):
413+ mtp_num_layers = getattr (margs , "mtp_num_layers" , 0 )
414+ if getattr (args , "skip_mtp" , False ):
415+ mtp_num_layers = 0
416+ if mtp_num_layers :
417+ for mtp_layer_id in range (mtp_num_layers ):
412418 message = dict ()
413419 ckpt_plugin .get_mtp_ckpt (message , models , mtp_layer_id , margs )
414420 queue_put (f"mtp module { mtp_layer_id } " , message )
0 commit comments