Update llama.py
This commit is contained in:
parent
0d680d87f4
commit
cbc3401649
1 changed files with 38 additions and 54 deletions
|
|
@ -2510,18 +2510,24 @@ class FastLlamaModel:
|
|||
# return
|
||||
# pass
|
||||
|
||||
internal_model = model
|
||||
internal_model.gradient_checkpointing = False
|
||||
internal_model.training = False
|
||||
|
||||
while hasattr(internal_model, "model"):
|
||||
internal_model = internal_model.model
|
||||
internal_model.gradient_checkpointing = False
|
||||
internal_model.training = False
|
||||
pass
|
||||
if hasattr(internal_model, "training"):
|
||||
internal_model.training = False
|
||||
m = model
|
||||
while hasattr(m, "model"):
|
||||
if hasattr(m, "gradient_checkpointing"):
|
||||
m.gradient_checkpointing = False
|
||||
if hasattr(m, "training"):
|
||||
m.training = False
|
||||
# Pad tokenizer to the left
|
||||
if hasattr(m, "_saved_temp_tokenizer"):
|
||||
m._saved_temp_tokenizer.padding_side = "left"
|
||||
m = m.model
|
||||
pass
|
||||
if hasattr(m, "gradient_checkpointing"):
|
||||
m.gradient_checkpointing = False
|
||||
if hasattr(m, "training"):
|
||||
m.training = False
|
||||
# Pad tokenizer to the left
|
||||
if hasattr(m, "_saved_temp_tokenizer"):
|
||||
m._saved_temp_tokenizer.padding_side = "left"
|
||||
|
||||
# Also check if lm_head / embeddings are trained
|
||||
internal_model = model
|
||||
|
|
@ -2530,30 +2536,13 @@ class FastLlamaModel:
|
|||
pass
|
||||
lm_head = internal_model.lm_head.weight
|
||||
device_type = lm_head.device.type
|
||||
dtype = model.config.torch_dtype
|
||||
|
||||
if type(dtype) is str:
|
||||
if dtype == "float16": dtype = torch.float16
|
||||
elif dtype == "bfloat16": dtype = torch.bfloat16
|
||||
pass
|
||||
dtype = _get_dtype(model.config.torch_dtype)
|
||||
|
||||
# Wrap model.generate
|
||||
if model.generate.__name__ != "_fast_generate":
|
||||
model._unwrapped_old_generate = model.generate
|
||||
model.generate = _wrap_fast_inference(model.generate, device_type, dtype, model)
|
||||
pass
|
||||
|
||||
# Patch tokenizer to pad to the left
|
||||
internal_model = model
|
||||
while hasattr(internal_model, "model"):
|
||||
if hasattr(internal_model, "_saved_temp_tokenizer"):
|
||||
internal_model._saved_temp_tokenizer.padding_side = "left"
|
||||
pass
|
||||
internal_model = internal_model.model
|
||||
pass
|
||||
if hasattr(internal_model, "_saved_temp_tokenizer"):
|
||||
internal_model._saved_temp_tokenizer.padding_side = "left"
|
||||
pass
|
||||
|
||||
# Also disable training for embeddings for NEFTune
|
||||
if hasattr(model, "get_input_embeddings"):
|
||||
|
|
@ -2571,9 +2560,6 @@ class FastLlamaModel:
|
|||
|
||||
@staticmethod
|
||||
def for_training(model, use_gradient_checkpointing = True):
|
||||
internal_model = model
|
||||
internal_model.gradient_checkpointing = use_gradient_checkpointing
|
||||
internal_model.training = True
|
||||
|
||||
# Delete all fast inference loras
|
||||
for param in model.parameters():
|
||||
|
|
@ -2581,14 +2567,24 @@ class FastLlamaModel:
|
|||
del param._fast_lora
|
||||
pass
|
||||
|
||||
while hasattr(internal_model, "model"):
|
||||
internal_model = internal_model.model
|
||||
internal_model.gradient_checkpointing = use_gradient_checkpointing
|
||||
internal_model.training = True
|
||||
pass
|
||||
if hasattr(internal_model, "training"):
|
||||
internal_model.training = True
|
||||
m = model
|
||||
while hasattr(m, "model"):
|
||||
if hasattr(m, "gradient_checkpointing"):
|
||||
m.gradient_checkpointing = use_gradient_checkpointing
|
||||
if hasattr(m, "training"):
|
||||
m.training = True
|
||||
# Pad tokenizer to the right
|
||||
if hasattr(m, "_saved_temp_tokenizer"):
|
||||
m._saved_temp_tokenizer.padding_side = "right"
|
||||
m = m.model
|
||||
pass
|
||||
if hasattr(m, "gradient_checkpointing"):
|
||||
m.gradient_checkpointing = use_gradient_checkpointing
|
||||
if hasattr(m, "training"):
|
||||
m.training = True
|
||||
# Pad tokenizer to the right
|
||||
if hasattr(m, "_saved_temp_tokenizer"):
|
||||
m._saved_temp_tokenizer.padding_side = "right"
|
||||
|
||||
# Also revert model.generate
|
||||
if hasattr(model, "_unwrapped_old_generate"):
|
||||
|
|
@ -2596,18 +2592,6 @@ class FastLlamaModel:
|
|||
del model._unwrapped_old_generate
|
||||
pass
|
||||
|
||||
# Patch tokenizer to pad to the right
|
||||
internal_model = model
|
||||
while hasattr(internal_model, "model"):
|
||||
if hasattr(internal_model, "_saved_temp_tokenizer"):
|
||||
internal_model._saved_temp_tokenizer.padding_side = "right"
|
||||
pass
|
||||
internal_model = internal_model.model
|
||||
pass
|
||||
if hasattr(internal_model, "_saved_temp_tokenizer"):
|
||||
internal_model._saved_temp_tokenizer.padding_side = "right"
|
||||
pass
|
||||
|
||||
# Also re-enable training for embeddings for NEFTune
|
||||
if hasattr(model, "get_input_embeddings"):
|
||||
embeddings = model.get_input_embeddings()
|
||||
|
|
@ -2617,7 +2601,7 @@ class FastLlamaModel:
|
|||
embeddings = model.get_output_embeddings()
|
||||
if hasattr(embeddings, "training"): embeddings.training = True
|
||||
pass
|
||||
|
||||
|
||||
return model
|
||||
pass
|
||||
pass
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue