4545 report_current_memory_info ,
4646)
4747from megatron .training import get_args , get_model , initialize_megatron
48+ from megatron .training .arguments import parse_and_validate_args
4849from utils import get_hf_tokenizer
4950from megatron .training .checkpointing import save_checkpoint
5051from megatron .training .utils import print_rank_0 , unwrap_model
@@ -188,13 +189,10 @@ def get_first_layers_disabled_config(config, num_layers: int = 1, num_layers_to_
188189 """
189190 config = copy .deepcopy (config )
190191 quant_cfg = config .get ("quant_cfg" , {})
191- quant_cfg .update (
192- {
193- functools .partial (
194- _is_first_layers , num_layers = num_layers , num_layers_to_disable = num_layers_to_disable
195- ): {"enable" : False }
196- }
192+ predicate = functools .partial (
193+ _is_first_layers , num_layers = num_layers , num_layers_to_disable = num_layers_to_disable
197194 )
195+ quant_cfg .append ({"quantizer_name" : predicate , "enable" : False })
198196 config ["quant_cfg" ] = quant_cfg
199197 return config
200198
@@ -206,13 +204,10 @@ def get_last_layers_disabled_config(config, num_layers: int = 1, num_layers_to_d
206204 """
207205 config = copy .deepcopy (config )
208206 quant_cfg = config .get ("quant_cfg" , {})
209- quant_cfg .update (
210- {
211- functools .partial (
212- _is_last_layers , num_layers = num_layers , num_layers_to_disable = num_layers_to_disable
213- ): {"enable" : False }
214- }
207+ predicate = functools .partial (
208+ _is_last_layers , num_layers = num_layers , num_layers_to_disable = num_layers_to_disable
215209 )
210+ quant_cfg .append ({"quantizer_name" : predicate , "enable" : False })
216211 config ["quant_cfg" ] = quant_cfg
217212 return config
218213
@@ -224,28 +219,34 @@ def get_modelopt_torch_quantization_config():
224219 raise ValueError (f"Unsupported quantization config { args .export_quant_cfg } ." )
225220 mtq_config = QUANT_CFG_CHOICES [args .export_quant_cfg ]
226221
227- fp8_config = {"enable" : True , "num_bits" : (4 , 3 ), "axis" : None }
222+ if isinstance (mtq_config ["quant_cfg" ], dict ):
223+ # Normalize old dict format to new list format
224+ mtq_config ["quant_cfg" ] = mtq .normalize_quant_cfg_list (mtq_config ["quant_cfg" ])
225+
226+ fp8_config = {"enable" : True , "cfg" : {"num_bits" : (4 , 3 ), "axis" : None }}
228227 fp4_config = {
229- "num_bits" : (2 , 1 ),
230- "block_sizes" : {- 1 : 16 , "type" : "dynamic" , "scale_bits" : (4 , 3 )},
231- "axis" : None ,
232228 "enable" : True ,
229+ "cfg" : {
230+ "num_bits" : (2 , 1 ),
231+ "block_sizes" : {- 1 : 16 , "type" : "dynamic" , "scale_bits" : (4 , 3 )},
232+ "axis" : None ,
233+ },
233234 }
234235 if args .export_quant_cfg == "FP8_DEFAULT_CFG" :
235236 # Enable Medusa heads and kv-cache quantization
236- mtq_config ["quant_cfg" ][ " *medusa_heads**"] = fp8_config
237+ mtq_config ["quant_cfg" ]. append ({ "quantizer_name" : " *medusa_heads**", ** fp8_config })
237238 if "FP4" in args .export_quant_cfg :
238239 # Enable Medusa heads and kv-cache quantization
239- mtq_config ["quant_cfg" ][ " *medusa_heads**"] = fp4_config
240+ mtq_config ["quant_cfg" ]. append ({ "quantizer_name" : " *medusa_heads**", ** fp4_config })
240241 if "AWQ" in args .export_quant_cfg :
241- weight_quantizer = mtq_config [ "quant_cfg" ][ "*weight_quantizer" ] # type: ignore
242- if isinstance ( weight_quantizer , list ):
243- weight_quantizer = weight_quantizer [ 0 ]
244- weight_quantizer [ "block_sizes" ][ - 1 ] = 128
245-
242+ try :
243+ weight_quantizer = mtq . find_quant_cfg_entry_by_path ( mtq_config [ "quant_cfg" ], "*weight_quantizer" )
244+ weight_quantizer [ "block_sizes" ][ - 1 ] = 128
245+ except KeyError :
246+ weight_quantizer = None
246247 # Customization
247248 if args .disable_qkv_quant :
248- mtq_config ["quant_cfg" ][ " *self_attention*"] = { "enable" : False }
249+ mtq_config ["quant_cfg" ]. append ({ "quantizer_name" : " *self_attention*", "enable" : False })
249250
250251 # KV Cache Quantization
251252 enable_quant_kv_cache = args .export_kv_cache_quant != "none"
@@ -257,7 +258,7 @@ def get_modelopt_torch_quantization_config():
257258
258259 # Weight Only Quantization
259260 if args .weight_only :
260- mtq_config ["quant_cfg" ][ " *input_quantizer"] = { "enable" : False }
261+ mtq_config ["quant_cfg" ]. append ({ "quantizer_name" : " *input_quantizer", "enable" : False })
261262 if args .num_first_layers_to_skip_quant is not None :
262263 mtq_config = get_first_layers_disabled_config (
263264 mtq_config ,
@@ -346,14 +347,12 @@ def get_calib_dataloader(
346347
347348
348349if __name__ == "__main__" :
349- initialize_megatron (
350- extra_args_provider = add_text_generate_ptq_args ,
351- args_defaults = {
350+ parse_and_validate_args (extra_args_provider = add_text_generate_ptq_args , args_defaults = {
352351 "tokenizer_type" : "HuggingFaceTokenizer" ,
353352 "no_load_rng" : True ,
354353 "no_load_optim" : True ,
355- },
356- )
354+ })
355+ initialize_megatron ( )
357356
358357 check_arguments ()
359358
0 commit comments