diff --git a/unsloth/models/gemma.py b/unsloth/models/gemma.py index 095c8cdd6b..1ec116b2ea 100644 --- a/unsloth/models/gemma.py +++ b/unsloth/models/gemma.py @@ -339,18 +339,7 @@ class FastGemmaModel(FastLlamaModel): @staticmethod - def post_patch(model, tokenizer, max_seq_length): - # Add max_seq_length to all modules - extra_ignored_labels = torch.full((max_seq_length, 1), -100, device = "cuda:0") - internal_model = model - while hasattr(internal_model, "model"): - internal_model.max_seq_length = max_seq_length - internal_model.extra_ignored_labels = extra_ignored_labels - internal_model = internal_model.model - pass - internal_model.max_seq_length = max_seq_length - internal_model.extra_ignored_labels = extra_ignored_labels - + def post_patch(model, tokenizer): # Torch.compile fails on embedding matrix?? # Workaround randomnly fixes it for torch versions < 2.2 model.model.embed_tokens = torch.nn.Embedding.from_pretrained(model.model.embed_tokens.weight) diff --git a/unsloth/models/gemma2.py b/unsloth/models/gemma2.py index 231d8f2661..54d8f628cb 100644 --- a/unsloth/models/gemma2.py +++ b/unsloth/models/gemma2.py @@ -490,18 +490,7 @@ class FastGemma2Model(FastLlamaModel): @staticmethod - def post_patch(model, tokenizer, max_seq_length): - # Add max_seq_length to all modules - extra_ignored_labels = torch.full((max_seq_length, 1), -100, device = "cuda:0") - internal_model = model - while hasattr(internal_model, "model"): - internal_model.max_seq_length = max_seq_length - internal_model.extra_ignored_labels = extra_ignored_labels - internal_model = internal_model.model - pass - internal_model.max_seq_length = max_seq_length - internal_model.extra_ignored_labels = extra_ignored_labels - + def post_patch(model, tokenizer): # Torch.compile fails on embedding matrix?? # Workaround randomnly fixes it for torch versions < 2.2 model.model.embed_tokens = torch.nn.Embedding.from_pretrained(model.model.embed_tokens.weight) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 4712d9ca05..65e2d773e8 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -1621,7 +1621,7 @@ class FastLlamaModel: ) model, tokenizer = patch_tokenizer(model, tokenizer) - model, tokenizer = model_patcher.post_patch(model, tokenizer, max_position_embeddings) + model, tokenizer = model_patcher.post_patch(model, tokenizer) # Patch up QKV / O and MLP for idx, layer in enumerate(model.model.layers): @@ -1827,18 +1827,7 @@ class FastLlamaModel: @staticmethod - def post_patch(model, tokenizer, max_seq_length): - # Add max_seq_length to all modules - extra_ignored_labels = torch.full((max_seq_length, 1), -100, device = "cuda:0") - internal_model = model - while hasattr(internal_model, "model"): - internal_model.max_seq_length = max_seq_length - internal_model.extra_ignored_labels = extra_ignored_labels - internal_model = internal_model.model - pass - internal_model.max_seq_length = max_seq_length - internal_model.extra_ignored_labels = extra_ignored_labels - + def post_patch(model, tokenizer): # Torch.compile fails on embedding matrix?? try: old_input_embedding = model.get_input_embeddings ().weight except: return model, tokenizer @@ -2470,6 +2459,18 @@ class FastLlamaModel: ) patch_saving_functions(model) + # Patch cross entropy loss labels + # Fixes https://github.com/unslothai/unsloth/issues/10 + max_seq_length = model.max_seq_length + extra_ignored_labels = torch.full((max_seq_length, 1), -100, device = "cuda:0") + model.model.extra_ignored_labels = extra_ignored_labels + internal_model = model + while hasattr(internal_model, "model"): + internal_model.max_seq_length = max_seq_length + internal_model = internal_model.model + pass + internal_model.max_seq_length = max_seq_length + # Patch tokenizer to pad to the right internal_model = model while hasattr(internal_model, "model"):