diff --git a/configs/ds_config_zero3.json b/configs/ds_config_zero3.json index 877a2e93a..2fedca377 100644 --- a/configs/ds_config_zero3.json +++ b/configs/ds_config_zero3.json @@ -12,16 +12,6 @@ "enabled": "auto" }, - "optimizer": { - "type": "AdamW", - "params": { - "lr": "auto", - "betas": "auto", - "eps": "auto", - "weight_decay": "auto" - } - }, - "zero_optimization": { "stage": 3, "offload_optimizer": { diff --git a/src/lmflow/models/hf_decoder_model.py b/src/lmflow/models/hf_decoder_model.py index 09fe3a294..8b316570c 100644 --- a/src/lmflow/models/hf_decoder_model.py +++ b/src/lmflow/models/hf_decoder_model.py @@ -276,18 +276,32 @@ def __init__( bnb_4bit_use_double_quant=model_args.double_quant, bnb_4bit_quant_type=model_args.quant_type, ) - model = AutoModelForCausalLM.from_pretrained( - model_args.model_name_or_path, - from_tf=bool(".ckpt" in model_args.model_name_or_path), - config=config, - quantization_config=quant_config if model_args.use_qlora else None, - cache_dir=model_args.cache_dir, - revision=model_args.model_revision, - use_auth_token=True if model_args.use_auth_token else None, - torch_dtype=torch_dtype, - device_map=device_map, - trust_remote_code = model_args.trust_remote_code, - ) + try: + model = AutoModelForCausalLM.from_pretrained( + model_args.model_name_or_path, + from_tf=bool(".ckpt" in model_args.model_name_or_path), + config=config, + quantization_config=quant_config if model_args.use_qlora else None, + cache_dir=model_args.cache_dir, + revision=model_args.model_revision, + use_auth_token=True if model_args.use_auth_token else None, + torch_dtype=torch_dtype, + device_map=device_map, + trust_remote_code = model_args.trust_remote_code, + ) + #for deepspeed zero3, we don't need to specify device_map + except: + model = AutoModelForCausalLM.from_pretrained( + model_args.model_name_or_path, + from_tf=bool(".ckpt" in model_args.model_name_or_path), + config=config, + quantization_config=quant_config if model_args.use_qlora else None, + cache_dir=model_args.cache_dir, + revision=model_args.model_revision, + use_auth_token=True if model_args.use_auth_token else None, + torch_dtype=torch_dtype, + trust_remote_code = model_args.trust_remote_code, + ) if model_args.use_qlora: model.gradient_checkpointing_enable() model = prepare_model_for_kbit_training(model)