Update llama.py

This commit is contained in:
Daniel Han 2025-02-02 02:06:13 -08:00
commit cbc3401649

View file

@ -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