diff --git a/src/peft/__init__.py b/src/peft/__init__.py index df42cef..b0784a4 100644 --- a/src/peft/__init__.py +++ b/src/peft/__init__.py @@ -47,6 +47,7 @@ from .utils import ( TaskType, bloom_model_postprocess_past_key_value, get_peft_model_state_dict, + prepare_model_for_training, set_peft_model_state_dict, shift_tokens_right, ) diff --git a/src/peft/peft_model.py b/src/peft/peft_model.py index fc91a25..e2932d9 100644 --- a/src/peft/peft_model.py +++ b/src/peft/peft_model.py @@ -288,36 +288,6 @@ class PeftModel(PushToHubMixin, torch.nn.Module): else: return self.base_model.model(*args, **kwargs) - def prepare_model_for_training(self): - r""" - This method wrapps the entire protocol for preparing a model before running a training. This includes: - 1- Cast the layernorm in fp32 2- making output embedding layer require grads - """ - loaded_in_8bit = getattr(self.base_model, "is_loaded_in_8bit", False) - - for param in self.base_model.parameters(): - # freeze base model's layers - param.requires_grad = False - - if loaded_in_8bit: - # cast layer norm in fp32 for stability for 8bit models - if param.ndim == 1: - param.data = param.data.to(torch.float32) - - # For backward compatibility - if hasattr(self.base_model, "enable_input_require_grads"): - self.base_model.enable_input_require_grads() - else: - - def make_inputs_require_grad(module, input, output): - output.requires_grad_(True) - - self.base_model.get_input_embeddings().register_forward_hook(make_inputs_require_grad) - - if loaded_in_8bit: - # enable gradient checkpointing for memory efficiency - self.base_model.model.gradient_checkpointing_enable() - class PeftModelForSequenceClassification(PeftModel): """ diff --git a/src/peft/utils/__init__.py b/src/peft/utils/__init__.py index 55cdce7..f466824 100644 --- a/src/peft/utils/__init__.py +++ b/src/peft/utils/__init__.py @@ -23,6 +23,7 @@ from .other import ( TRANSFORMERS_MODELS_TO_PREFIX_TUNING_POSTPROCESS_MAPPING, _set_trainable, bloom_model_postprocess_past_key_value, + prepare_model_for_training, shift_tokens_right, transpose, ) diff --git a/src/peft/utils/other.py b/src/peft/utils/other.py index 3f8627e..c062534 100644 --- a/src/peft/utils/other.py +++ b/src/peft/utils/other.py @@ -30,6 +30,59 @@ def bloom_model_postprocess_past_key_value(past_key_values): return tuple(zip(keys, values)) +def prepare_model_for_training(model): + r""" + This method wrapps the entire protocol for preparing a model before running a training. This includes: + 1- Cast the layernorm in fp32 2- making output embedding layer require grads 3- Add the upcasting of the lm + head to fp32 + + Args: + model, (`transformers.PreTrainedModel`): + The loaded model from `transformers` + """ + loaded_in_8bit = getattr(model, "is_loaded_in_8bit", False) + + for param in model.parameters(): + # freeze base model's layers + param.requires_grad = False + + if loaded_in_8bit: + # cast layer norm in fp32 for stability for 8bit models + if param.ndim == 1: + param.data = param.data.to(torch.float32) + + # For backward compatibility + if hasattr(model, "enable_input_require_grads"): + model.enable_input_require_grads() + else: + + def make_inputs_require_grad(module, input, output): + output.requires_grad_(True) + + model.get_input_embeddings().register_forward_hook(make_inputs_require_grad) + + if loaded_in_8bit: + # enable gradient checkpointing for memory efficiency + model.gradient_checkpointing_enable() + + if hasattr(model, "lm_head"): + input_dtype = model.lm_head.weight.dtype + + class CastOutputToFloat(torch.nn.Sequential): + r""" + Manually cast to the expected dtype of the lm_head as sometimes there is a final layer norm that is casted + in fp32 + + """ + + def forward(self, x): + return super().forward(x.to(input_dtype)).to(torch.float32) + + model.lm_head = CastOutputToFloat(model.lm_head) + + return model + + TRANSFORMERS_MODELS_TO_PREFIX_TUNING_POSTPROCESS_MAPPING = { "bloom": bloom_model_postprocess_past_key_value, } diff --git a/tests/test_save_and_load.py b/tests/test_peft_model.py similarity index 92% rename from tests/test_save_and_load.py rename to tests/test_peft_model.py index b5d775e..5ceda03 100644 --- a/tests/test_save_and_load.py +++ b/tests/test_peft_model.py @@ -27,6 +27,7 @@ from peft import ( PromptTuningConfig, get_peft_model, get_peft_model_state_dict, + prepare_model_for_training, ) @@ -103,11 +104,19 @@ class PeftModelTester(unittest.TestCase, PeftTestMixin): self.assertTrue(not dummy_output.requires_grad) - model.prepare_model_for_training() + # load with `prepare_model_for_training` + model = AutoModelForCausalLM.from_pretrained(model_id) + model = prepare_model_for_training(model) - for param in model.base_model.parameters(): + for param in model.parameters(): self.assertTrue(not param.requires_grad) + config = config_cls( + base_model_name_or_path=model_id, + **self.config_kwargs[i], + ) + model = get_peft_model(model, config) + dummy_input = torch.LongTensor([[1, 1, 1]]) dummy_output = model.get_input_embeddings()(dummy_input)