diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 762ebd1fdf..af1d35bd9f 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -1040,19 +1040,24 @@ pass def _unsloth_pre_compute_loss(self, model, inputs, *args, **kwargs): + num_items_in_batch = None + if "num_items_in_batch" in kwargs: - if kwargs["num_items_in_batch"] is None: + 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"] = kwargs["num_items_in_batch"] + inputs["num_items_in_batch"] = num_items_in_batch pass - else: + 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 on our blog post: https://unsloth.ai/blog/gradient" + "Read more on gradient accumulation issues here: https://unsloth.ai/blog/gradient" ) pass return self._old_compute_loss(model, inputs, *args, **kwargs)