apply suggestions

This commit is contained in:
younesbelkada
2023-02-14 11:50:55 +00:00
parent 0e80648010
commit 36c7e3b441
5 changed files with 66 additions and 32 deletions
+1
View File
@@ -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,
)
-30
View File
@@ -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):
"""
+1
View File
@@ -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,
)
+53
View File
@@ -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,
}
@@ -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)