Update _utils.py

This commit is contained in:
Daniel Han 2024-12-27 00:59:16 -08:00
commit 7fdcb7f24d

View file

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