diff --git a/pyproject.toml b/pyproject.toml index 9abe7a5d88..ce3301547b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -39,7 +39,7 @@ triton = [ "triton @ https://github.com/woct0rdho/triton-windows/releases/download/v3.1.0-windows.post5/triton-3.1.0-cp312-cp312-win_amd64.whl ; python_version=='3.12' and platform_system == 'Windows'", ] huggingface = [ - "unsloth_zoo>=2024.12.6", + "unsloth_zoo>=2024.12.7", "packaging", "tyro", "transformers>=4.46.1,!=4.47.0", @@ -285,7 +285,7 @@ colab-ampere-torch220 = [ "flash-attn>=2.6.3", ] colab-new = [ - "unsloth_zoo>=2024.12.6", + "unsloth_zoo>=2024.12.7", "packaging", "tyro", "transformers>=4.46.1,!=4.47.0", diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 4f1b40884a..86346d7e2e 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -1003,28 +1003,62 @@ pass def _unsloth_get_batch_samples(self, epoch_iterator, num_batches): batch_samples = [] num_items_in_batch = None + + # Check if model allows **kwargs + model = self.model + f = model.base_model.model.forward if hasattr(model, "base_model") else model.forward + has_kwargs = tuple(inspect.signature(f).parameters.values())[-1].kind == inspect._VAR_KEYWORD + + # Iterate to find all batches for _ in range(num_batches): try: batch_samples += [next(epoch_iterator)] except StopIteration: break - if len(batch_samples) > 0 and "labels" in batch_samples[0]: + pass + + # Get num_items_in_batch + if has_kwargs and len(batch_samples) > 0 and "labels" in batch_samples[0]: try: num_items_in_batch = sum( - [torch.count_nonzero(x["labels"][..., 1:] != -100) for x in batch_samples] + [(x["labels"][..., 1:] != -100).sum() for x in batch_samples] ) - except TypeError: - pass + + if self.args.average_tokens_across_devices: + num_items_in_batch = self.accelerator.gather(num_items_in_batch).sum().item() + + if torch.is_tensor(num_items_in_batch): + num_items_in_batch = num_items_in_batch.item() + + except Exception as exception: + logger.warning_once(exception) + pass + return batch_samples, num_items_in_batch pass def _unsloth_pre_compute_loss(self, model, inputs, *args, **kwargs): + num_items_in_batch = None + if "num_items_in_batch" in kwargs: - if "num_items_in_batch" not in inputs: - inputs["num_items_in_batch"] = kwargs["num_items_in_batch"] + num_items_in_batch = kwargs["num_items_in_batch"] + if num_items_in_batch is None: + # Remove it since the model does not support it! + kwargs.pop("num_items_in_batch") + elif "num_items_in_batch" not in inputs: + inputs["num_items_in_batch"] = num_items_in_batch pass pass + + if num_items_in_batch is None: + name = (model.base_model.model if hasattr(model, "base_model") else model).__class__.__name__ + logger.warning_once( + f"Unsloth: Not an error, but {name} does not accept `num_items_in_batch`.\n"\ + "Using gradient accumulation will be very slightly less accurate.\n"\ + "Read more on gradient accumulation issues here: https://unsloth.ai/blog/gradient" + ) + pass return self._old_compute_loss(model, inputs, *args, **kwargs) pass @@ -1104,7 +1138,7 @@ def patch_gradient_accumulation_fix(Trainer): "else:\n"\ "\2if num_items_in_batch is None:\n"\ - "\3loss /= self.args.gradient_accumulation_steps\n"\ + "\3loss = loss / self.args.gradient_accumulation_steps\n"\ "\1self.accelerator.backward(loss, **kwargs)", function, @@ -1148,6 +1182,7 @@ def unsloth_compile_transformers( manual_replacements = True, fast_lora_forwards = True, fast_residual_stream = True, + accurate_accumulation = True, epilogue_fusion = True, max_autotune = False, shape_padding = True, @@ -1194,6 +1229,7 @@ def unsloth_compile_transformers( manual_replacements = manual_replacements, fast_lora_forwards = fast_lora_forwards, fast_residual_stream = fast_residual_stream, + accurate_accumulation = accurate_accumulation, epilogue_fusion = epilogue_fusion, max_autotune = max_autotune, shape_padding = shape_padding, diff --git a/unsloth/models/loader.py b/unsloth/models/loader.py index d1c8b1e07b..113c4fbc70 100644 --- a/unsloth/models/loader.py +++ b/unsloth/models/loader.py @@ -472,6 +472,7 @@ class FastVisionModel(FastBaseVisionModel): manual_replacements = True, fast_lora_forwards = False, fast_residual_stream = False, + accurate_accumulation = True, epilogue_fusion = True, max_autotune = False, shape_padding = True, diff --git a/unsloth/models/vision.py b/unsloth/models/vision.py index 709cd1cb5c..2dc4b88dfa 100644 --- a/unsloth/models/vision.py +++ b/unsloth/models/vision.py @@ -186,6 +186,10 @@ class FastBaseVisionModel: patch_saving_functions(model, vision = True) patch_saving_functions(tokenizer, vision = True) + # Fix gradient accumulation + from transformers.trainer import Trainer + patch_gradient_accumulation_fix(Trainer) + # Save tokenizer for inference purposes tokenizer.padding_side = "left" # Force inference tokenizer.tokenizer.padding_side = "left" # Force inference diff --git a/unsloth/save.py b/unsloth/save.py index 8db3b6dc35..d3ba1928c4 100644 --- a/unsloth/save.py +++ b/unsloth/save.py @@ -2131,7 +2131,8 @@ def unsloth_generic_save( if token is None and push_to_hub: token = get_token() merge_and_overwrite_lora( get_model_name, - model, + model = model, + tokenizer = tokenizer, save_directory = save_directory, push_to_hub = push_to_hub, private = private,