Fix DPO, ORPO

This commit is contained in:
Daniel Han 2024-10-24 00:17:26 -07:00
commit 4ff247ab18
3 changed files with 32 additions and 4 deletions

View file

@ -375,7 +375,6 @@ def fast_cross_entropy_loss(
)
if n_items is None:
n_items = torch.count_nonzero(labels != -100)
print(n_items)
return loss.sum() / n_items
pass

View file

@ -1172,10 +1172,10 @@ pass
def patch_gradient_accumulation_fix(Trainer):
# Fixes gradient accumulation
import inspect
if hasattr(Trainer, "get_batch_samples"):
from inspect import getsource
if \
not getsource(Trainer.get_batch_samples).strip()\
not inspect.getsource(Trainer.get_batch_samples).strip()\
.endswith("return batch_samples, num_items_in_batch"):
raise NotImplementedError("Unsloth: Please make a Github issue immediately!!")
@ -1198,4 +1198,33 @@ def patch_gradient_accumulation_fix(Trainer):
'`pip install --upgrade --no-cache-dir unsloth git+https://github.com/huggingface/transformers.git git+https://github.com/huggingface/trl.git`'
)
pass
# Also fix up loss scaling ie negate loss *= self.args.gradient_accumulation_steps
if "num_items_in_batch" not in inspect.signature(Trainer.training_step).parameters: return
function = inspect.getsource(Trainer.training_step)
where = function.find("def")
function = function.split("\n")
function = "\n".join(x[where:] for x in function)
# Import all variables that need importing
import transformers.trainer
items_in_trainer = dir(transformers.trainer)
good_items = []
for item in items_in_trainer:
# TODO: Support Deepspeed
if item.startswith(("deepspeed", "xm", "met", "smp")): continue
if item in function: good_items.append(item)
pass
exec("from transformers.trainer import (" + ", ".join(x for x in good_items) + ")", globals())
# Accelerate does / self.args.gradient_accumulation_steps internally, so if we already
# summed it up and did the division before hand, we have to negate it.
function = function.replace(
"loss *= self.args.gradient_accumulation_steps",
"if num_items_in_batch is not None: loss *= self.args.gradient_accumulation_steps",
)
function = function.replace("def training_step", "def _unsloth_training_step", 1)
exec(function, globals())
Trainer.training_step = _unsloth_training_step
pass

View file

@ -145,7 +145,7 @@ pass
def _merge_lora(layer, name):
bias = None
bias = getattr(layer, "bias", None)
if isinstance(layer, (Bnb_Linear4bit, Peft_Linear4bit, Peft_Linear)):
# Is LoRA so we need to merge!
W, quant_state, A, B, s, bias = get_lora_parameters_bias(layer)