Forward hook
This commit is contained in:
parent
34209e9160
commit
44b3006678
2 changed files with 8 additions and 7 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue