Forward hook

This commit is contained in:
Daniel Han 2024-11-05 01:36:11 -08:00
commit 44b3006678
2 changed files with 8 additions and 7 deletions

View file

@ -467,13 +467,13 @@ def prepare_model_for_kbit_training(
pass
# If use_reentrant = True which is the Pytorch default, we just make the input requires_grad.
if use_reentrant:
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 use_reentrant:
# 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)
return model
pass

View file

@ -606,6 +606,7 @@ def LlamaModel_fast_forward(
# Embed positions
if inputs_embeds is None:
inputs_embeds = self.embed_tokens(input_ids)
inputs_embeds.requires_grad_(True)
# inputs_embeds = inputs_embeds.to(self.config.torch_dtype)
torch_dtype = __DTYPE_MAP.get(self.config.torch_dtype, None)