"""
2026.6.7
2026.6.9
5.5.0
1.7.0
__UNSLOTH_VERSIONING__
"""
# Unsloth auto generated code
# Copyright 2023-present Daniel Han-Chen, Michael Han-Chen & the Unsloth team. All rights reserved.
#
# This program is free software: you can redistribute it and/or modify
# it under the terms of the GNU Lesser General Public License as published by
# the Free Software Foundation, either version 3 of the License, or
# (at your option) any later version.
#
# This program is distributed in the hope that it will be useful,
# but WITHOUT ANY WARRANTY; without even the implied warranty of
# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
# GNU General Public License for more details.
#
# You should have received a copy of the GNU Lesser General Public License
# along with this program. If not, see .
from torch import Tensor
import torch
import torch.nn as nn
from torch.nn import functional as F
from unsloth_zoo.temporary_patches.common import torch_compile
from typing import Any, List, Optional, Tuple, Union, Dict, Set, Callable
from trl.experimental.kto.kto_trainer import (Any, AutoProcessor, Callable, DataCollator, DataCollatorForUnpairedPreference, DataCollatorForVisionUnpairedPreference, DataLoader, Dataset, EvalLoopOutput, F, Hasher, IterableDataset, IterableDatasetDict, KTOConfig, KTOTrainer, LoraConfig, PartialState, Path, PeftConfig, PeftModel, PreTrainedModel, PreTrainedTokenizerBase, ProcessorMixin, Sampler, SequentialSampler, SyncRefModelCallback, TrainerCallback, Version, _BaseTrainer, _get_kl_completion_ids, apply_chat_template, concatenate_datasets, contextlib, create_model_from_path, dataclass, defaultdict, disable_dropout_in_model, disable_gradient_checkpointing, extract_prompt, flush_left, get_act_offloading_ctx_manager, get_config_model_id, get_dataset_column_names, get_peft_model, has_length, hash_module, is_conversational, is_liger_kernel_available, is_peft_available, is_peft_model, logger, os, pad, peft, prepare_deepspeed, prepare_fsdp, prepare_multimodal_messages, selective_log_softmax, textwrap, torch, tqdm, transformers, unpair_preference_dataset, use_adapter, AutoProcessor, Callable, DataCollator, DataCollatorForUnpairedPreference, DataCollatorForVisionUnpairedPreference, Dataset, EvalLoopOutput, F, IterableDataset, IterableDatasetDict, KTOConfig, KTOTrainer, LoraConfig, PeftConfig, PeftModel, PreTrainedModel, PreTrainedTokenizerBase, ProcessorMixin, SyncRefModelCallback, TrainerCallback, Version, contextlib, create_model_from_path, defaultdict, disable_dropout_in_model, get_act_offloading_ctx_manager, get_config_model_id, get_peft_model, is_liger_kernel_available, is_peft_available, is_peft_model, logger, os, pad, peft, prepare_deepspeed, prepare_fsdp, torch, transformers, unpair_preference_dataset, F, PeftModel, PreTrainedModel, is_peft_available, logger, os, peft, torch)
import os
import math
import logging
from typing import *
from dataclasses import dataclass, field
from packaging.version import Version
import torch
import numpy as np
from contextlib import nullcontext
from torch.nn import functional as F
import inspect
from transformers import DataCollatorForSeq2Seq, DataCollatorForLanguageModeling as TransformersDataCollatorForLanguageModeling
from transformers.training_args import ParallelMode
from unsloth_zoo.device_type import DEVICE_TYPE, device_synchronize
# Wrap trainer with padding to right and enable training mode
import functools
from types import MethodType
try:
from unsloth_zoo.gradient_checkpointing import reset_unsloth_gradient_checkpointing_buffers
except:
def reset_unsloth_gradient_checkpointing_buffers(): pass
# Canonical reset lives in unsloth.models._utils so the SFT auto-packing wrapper and the plain
# Trainer loop can import the same helper; fall back to a no-op only if it can't be imported.
try:
from unsloth.models._utils import _unsloth_reset_stray_compile_cache
except Exception:
def _unsloth_reset_stray_compile_cache(self): pass
def prepare_for_training_mode(f):
@functools.wraps(f)
def wrapper(self, *args, **kwargs):
# Drop any torch.compile graph cache poisoned by a stray pre-train forward.
try:
_unsloth_reset_stray_compile_cache(self)
except Exception:
pass
# Finish the previous W&B run if this is a subsequent train() call.
# We do this at the START of train() (not the end) so that
# evaluate() / log() still work after train() completes.
# HF's WandbCallback.setup() will call wandb.init() for the new run.
# See: https://github.com/unslothai/unsloth/issues/3954
if getattr(self, '_unsloth_training_completed', False):
try:
import wandb
if wandb.run is not None:
wandb.finish()
# Reset HF's WandbCallback so it calls wandb.init() for the new run
for cb in self.callback_handler.callbacks:
if type(cb).__name__ == 'WandbCallback':
cb._initialized = False
break
except:
pass
# Enable training mode
_was_training = None
# Get gradient checkpointing setting from training arguments
use_gc = getattr(self.args, 'gradient_checkpointing', True)
if hasattr(self, 'model') and hasattr(self.model, "training"):
_was_training = self.model.training
if hasattr(self, 'model') and hasattr(self.model, "for_training"):
self.model.for_training(use_gradient_checkpointing=use_gc)
output = f(self, *args, **kwargs)
# Restore previous mode when possible
if hasattr(self, 'model') and hasattr(self.model, "for_inference"):
if _was_training is False:
self.model.for_inference()
elif _was_training is True and hasattr(self.model, "for_training"):
self.model.for_training(use_gradient_checkpointing=use_gc)
# Reset gradient checkpointing buffers to free memory while staying ready for next run
try:
reset_unsloth_gradient_checkpointing_buffers()
except:
pass
# Mark that training completed so the next train() call can
# finish this W&B run before starting a new one
self._unsloth_training_completed = True
return output
return wrapper
pass
torch_compile_options = {
"epilogue_fusion" : True,
"max_autotune" : False,
"shape_padding" : True,
"trace.enabled" : False,
"triton.cudagraphs" : False,
}
@torch.compile(dynamic = True, fullgraph = True, options = torch_compile_options,)
def chunked_hidden_states_selective_log_softmax(
hidden_states: torch.Tensor,
lm_head: torch.Tensor,
index: torch.Tensor,
chunks: int = 4,
logit_scale_multiply: float = 0.0,
logit_scale_divide: float = 0.0,
logit_softcapping: float = 0.0,
temperature: float = 1.0,
) -> torch.Tensor:
# All Unsloth Zoo code licensed under AGPL3
flat_hidden_states = hidden_states.reshape(-1, hidden_states.shape[-1])
flat_index = index.reshape(-1)
chunked_hidden_states = torch.chunk(flat_hidden_states, chunks=chunks, dim=0)
chunked_index = torch.chunk(flat_index, chunks=chunks, dim=0)
all_per_token_logps = []
for chunk_hidden_states, chunk_index in zip(chunked_hidden_states, chunked_index):
chunk_logits = chunk_hidden_states.to(lm_head.dtype) @ lm_head.t()
if logit_scale_multiply != 0.0:
chunk_logits = chunk_logits * logit_scale_multiply
if logit_scale_divide != 0.0:
chunk_logits = chunk_logits / logit_scale_divide
if logit_softcapping != 0.0:
chunk_logits = logit_softcapping * torch.tanh(chunk_logits / logit_softcapping)
chunk_logits = chunk_logits.to(torch.float32)
if temperature != 1.0:
chunk_logits = chunk_logits / temperature
selected_logits = torch.gather(chunk_logits, dim=-1, index=chunk_index.unsqueeze(-1)).squeeze(-1)
logsumexp_values = torch.logsumexp(chunk_logits, dim=-1)
per_token_logps = selected_logits - logsumexp_values
all_per_token_logps.append(per_token_logps)
all_per_token_logps = torch.concat(all_per_token_logps)
all_per_token_logps = all_per_token_logps.reshape((hidden_states.shape[0], hidden_states.shape[1]))
return all_per_token_logps
@torch.compile(dynamic = True, fullgraph = True, options = torch_compile_options,)
def chunked_selective_log_softmax(
logits,
index,
temperature: float = 1.0,
chunks: int = 4,
):
chunked_logits = torch.chunk(logits.reshape(-1, logits.shape[-1]), chunks = chunks, dim = 0)
chunked_index = torch.chunk(index.reshape(-1), chunks = chunks, dim = 0)
all_per_token_logps = []
# Per-chunk selective_log_softmax.
for chunk_logits, chunk_index in zip(chunked_logits, chunked_index):
chunk_logits = chunk_logits.to(torch.float32)
if temperature != 1.0:
chunk_logits = chunk_logits / temperature
selected_logits = torch.gather(chunk_logits, dim = -1, index = chunk_index.unsqueeze(-1)).squeeze(-1)
logsumexp_values = torch.logsumexp(chunk_logits, dim = -1)
per_token_logps = selected_logits - logsumexp_values
all_per_token_logps.append(per_token_logps)
pass
all_per_token_logps = torch.concat(all_per_token_logps)
all_per_token_logps = all_per_token_logps.reshape((logits.shape[0], logits.shape[1]))
return all_per_token_logps
def calculate_pad_tokens_in_prompt(
input_ids: torch.Tensor,
logits_to_keep: int,
pad_token_id: int
) -> torch.Tensor:
"""Count left-padded tokens per sequence, e.g. [pad, pad, pad, cat] -> 3."""
if logits_to_keep >= input_ids.shape[1]:
raise ValueError("logits_to_keep must be smaller than the sequence length.")
prompt_section = input_ids[:, :-logits_to_keep]
padding_mask = (prompt_section == pad_token_id)
pad_token_counts = padding_mask.sum(dim=1)
return pad_token_counts
def create_completion_attention_mask(
completion_input_ids: torch.Tensor,
left_pad_tokens_per_prompt: torch.Tensor,
max_left_pad: int,
pad_token_id: int
) -> torch.Tensor:
"""Build a completion mask that zeros leading prompt and trailing pad tokens.
For [p,p,p,c,c,c,pad,pad,pad] (p=sliced prompt, c=completion, pad=padding)
this returns [0,0,0,1,1,1,0,0,0].
"""
batch_size, completion_len = completion_input_ids.shape
device = completion_input_ids.device
num_tokens_to_mask = max_left_pad - left_pad_tokens_per_prompt
indices = torch.arange(completion_len, device=device).unsqueeze(0)
shift_mask = indices >= num_tokens_to_mask.unsqueeze(1)
non_padding_mask = (completion_input_ids != pad_token_id)
final_mask = shift_mask & non_padding_mask
return final_mask
def left_pack_padding(tensor: torch.Tensor, pad_id: int) -> torch.Tensor:
"""Move all padding tokens in each sequence to the right."""
mask = (tensor != pad_id)
# stable=True since the binary mask is unordered.
sorted_indices = torch.argsort(mask, dim=1, descending=True, stable=True)
packed_tensor = torch.gather(tensor, 1, sorted_indices)
return packed_tensor
def align_logprobs_with_mask(
logprob_tensor: torch.Tensor,
attention_mask: torch.Tensor,
pad_value: float = 0.0
) -> torch.Tensor:
"""Align a log probability tensor with a given attention mask."""
device = logprob_tensor.device
batch_size, logprob_seq_len = logprob_tensor.shape
mask_seq_len = attention_mask.shape[1]
padded_logprobs = torch.full(
attention_mask.shape,
fill_value=pad_value,
dtype=logprob_tensor.dtype,
device=device
)
left_pad_counts = torch.argmax(attention_mask, dim=1)
cols = torch.arange(logprob_seq_len, device=device)
dest_indices = left_pad_counts.unsqueeze(1) + cols
# Destination row indices, shape [batch_size, logprob_seq_len].
row_indices = torch.arange(batch_size, device=device).unsqueeze(1).expand_as(dest_indices)
# Keep only in-bounds destinations, then scatter via advanced indexing.
valid_mask = dest_indices < mask_seq_len
valid_rows = row_indices[valid_mask]
valid_cols = dest_indices[valid_mask]
valid_vals = logprob_tensor[valid_mask]
padded_logprobs[valid_rows, valid_cols] = valid_vals
return padded_logprobs
def align_completion_tool_mask(
tool_mask: torch.Tensor,
completion_mask: torch.Tensor,
) -> torch.Tensor:
"""Align a raw completion-length tool/env mask with Unsloth's repacked loss mask."""
if tool_mask is None:
return completion_mask
if tool_mask.shape[0] != completion_mask.shape[0]:
raise ValueError("tool_mask batch size must match completion_mask batch size.")
tool_mask = tool_mask.to(device=completion_mask.device)
if tool_mask.shape == completion_mask.shape:
aligned_tool_mask = tool_mask
else:
aligned_tool_mask = align_logprobs_with_mask(
tool_mask,
completion_mask,
pad_value=0,
)
return completion_mask * aligned_tool_mask.to(dtype=completion_mask.dtype)
def autotune_batch_and_chunks(
total_input_rows,
seq_len,
hidden_size,
vocab_size,
dtype_bytes=16,
multiplier=None
):
if multiplier is None:
final_m = max(4, seq_len // 4096)
else:
final_m = multiplier
if torch.cuda.is_available():
free_bytes, _ = torch.cuda.mem_get_info()
limit_gb = (free_bytes / (1024**3))*.80
elif hasattr(torch, "xpu") and torch.xpu.is_available():
# XPU: estimate free memory as total - reserved.
total_mem = torch.xpu.get_device_properties(0).total_memory
reserved_mem = torch.xpu.memory_reserved()
free_bytes = total_mem - reserved_mem
limit_gb = (free_bytes / (1024**3)) * 0.80
else:
# Fallback: assume 8GB available.
limit_gb = 8.0
bytes_to_gb = 1024**3
b_vals = torch.arange(total_input_rows, 0, -1, device='cpu', dtype=torch.float32)
hidden_gb = (b_vals * seq_len * hidden_size * dtype_bytes) / bytes_to_gb
base_logits = ((b_vals/total_input_rows) * b_vals * seq_len * vocab_size * dtype_bytes) / bytes_to_gb
logits_gb = base_logits / final_m
total_mem_gb = hidden_gb + logits_gb
valid_mask = total_mem_gb <= limit_gb
valid_indices = torch.nonzero(valid_mask, as_tuple=False)
if valid_indices.shape[0] == 0:
#This means your GPU will OOM
return 4, final_m
best_idx = valid_indices[0].item()
final_b = int(b_vals[best_idx].item())
return final_b, final_m
def sanitize_logprob(logprob):
"""Local port of trl.scripts.vllm_serve.sanitize_logprob.
Filters NaN logprobs from vLLM outputs."""
value = logprob.logprob
if math.isnan(value):
logging.getLogger(__name__).warning(
f"Generated NaN logprob, token logprob '{logprob}' will be ignored"
)
return None
return value
@dataclass
class UnslothKTOConfig(KTOConfig):
"""
KTOConfig(output_dir: str | None = None, per_device_train_batch_size: int = 8, num_train_epochs: float = 3.0, max_steps: int = -1, learning_rate: float = 1e-06, lr_scheduler_type: transformers.trainer_utils.SchedulerType | str = 'linear', lr_scheduler_kwargs: dict | str | None = None, warmup_steps: float = 0, optim: transformers.training_args.OptimizerNames | str = 'adamw_torch_fused', optim_args: str | None = None, weight_decay: float = 0.0, adam_beta1: float = 0.9, adam_beta2: float = 0.999, adam_epsilon: float = 1e-08, optim_target_modules: None | str | list[str] = None, gradient_accumulation_steps: int = 1, average_tokens_across_devices: bool = True, max_grad_norm: float = 1.0, label_smoothing_factor: float = 0.0, bf16: bool | None = None, fp16: bool = False, bf16_full_eval: bool = False, fp16_full_eval: bool = False, tf32: bool | None = None, gradient_checkpointing: bool = True, gradient_checkpointing_kwargs: dict[str, typing.Any] | str | None = None, torch_compile: bool = False, torch_compile_backend: str | None = None, torch_compile_mode: str | None = None, use_liger_kernel: bool = False, liger_kernel_config: dict[str, bool] | None = None, use_cache: bool = False, neftune_noise_alpha: float | None = None, torch_empty_cache_steps: int | None = None, auto_find_batch_size: bool = False, logging_strategy: transformers.trainer_utils.IntervalStrategy | str = 'steps', logging_steps: float = 10, logging_first_step: bool = False, log_on_each_node: bool = True, logging_nan_inf_filter: bool = True, include_num_input_tokens_seen: str | bool = 'no', log_level: str = 'passive', log_level_replica: str = 'warning', disable_tqdm: bool | None = None, report_to: None | str | list[str] = 'none', run_name: str | None = None, project: str = 'huggingface', trackio_space_id: str | None = 'trackio', eval_strategy: transformers.trainer_utils.IntervalStrategy | str = 'no', eval_steps: float | None = None, eval_delay: float = 0, per_device_eval_batch_size: int = 8, prediction_loss_only: bool = False, eval_on_start: bool = False, eval_do_concat_batches: bool = True, eval_use_gather_object: bool = False, eval_accumulation_steps: int | None = None, include_for_metrics: list[str] = , batch_eval_metrics: bool = False, save_only_model: bool = False, save_strategy: transformers.trainer_utils.SaveStrategy | str = 'steps', save_steps: float = 500, save_on_each_node: bool = False, save_total_limit: int | None = None, enable_jit_checkpoint: bool = False, push_to_hub: bool = False, hub_token: str | None = None, hub_private_repo: bool | None = None, hub_model_id: str | None = None, hub_strategy: transformers.trainer_utils.HubStrategy | str = 'every_save', hub_always_push: bool = False, hub_revision: str | None = None, load_best_model_at_end: bool = False, metric_for_best_model: str | None = None, greater_is_better: bool | None = None, ignore_data_skip: bool = False, restore_callback_states_from_checkpoint: bool = False, full_determinism: bool = False, seed: int = 42, data_seed: int | None = None, use_cpu: bool = False, accelerator_config: dict | str | None = None, parallelism_config: accelerate.parallelism_config.ParallelismConfig | None = None, dataloader_drop_last: bool = False, dataloader_num_workers: int = 0, dataloader_pin_memory: bool = True, dataloader_persistent_workers: bool = False, dataloader_prefetch_factor: int | None = None, remove_unused_columns: bool = True, label_names: list[str] | None = None, train_sampling_strategy: str = 'sequential', length_column_name: str = 'length', ddp_find_unused_parameters: bool | None = None, ddp_bucket_cap_mb: int | None = None, ddp_broadcast_buffers: bool | None = None, ddp_backend: str | None = None, ddp_timeout: int = 1800, fsdp: list[transformers.trainer_utils.FSDPOption] | str | None = None, fsdp_config: dict[str, typing.Any] | str | None = None, deepspeed: dict | str | None = None, debug: str | list[transformers.debug_utils.DebugOption] = '', skip_memory_metrics: bool = True, do_train: bool = False, do_eval: bool = False, do_predict: bool = False, resume_from_checkpoint: str | None = None, warmup_ratio: float | None = None, logging_dir: str | None = None, local_rank: int = -1, model_init_kwargs: dict[str, typing.Any] | str | None = None, trust_remote_code: bool = False, disable_dropout: bool = True, dataset_num_proc: int | None = None, max_length: int | None = 1024, pad_to_multiple_of: int | None = None, precompute_ref_log_probs: bool = False, precompute_ref_batch_size: int | None = None, loss_type: str = 'kto', beta: float = 0.1, desirable_weight: float = 1.0, undesirable_weight: float = 1.0, activation_offloading: bool = False, sync_ref_model: bool = False, ref_model_mixup_alpha: float = 0.6, ref_model_sync_steps: int = 512)
"""
vllm_sampling_params: Optional[Any] = field(
default = None,
metadata = {'help': 'vLLM SamplingParams'},
)
unsloth_num_chunks : Optional[int] = field(
default = -1,
metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},
)
unsloth_logit_chunk_multiplier : Optional[int] = field(
default = None,
metadata = {'help': 'Multiplier for chunked logit computations.'},
)
unsloth_grpo_mini_batch : Optional[int] = field(
default = None,
metadata = {'help': 'Mini batch size for GRPO hidden state accumulation. Default is None unless user defines it.'},
)
max_seq_length : Optional[int] = field(
default = None,
metadata = {'help': 'Maximum sequence length to truncate to.'},
)
def __init__(
self,
output_dir = None,
per_device_train_batch_size = 4,
num_train_epochs = 3.0,
max_steps = -1,
learning_rate = 5e-05,
lr_scheduler_type = 'linear',
lr_scheduler_kwargs = None,
warmup_steps = 0.1,
optim = 'adamw_8bit',
optim_args = None,
weight_decay = 0.001,
adam_beta1 = 0.9,
adam_beta2 = 0.999,
adam_epsilon = 1e-08,
optim_target_modules = None,
gradient_accumulation_steps = 2,
average_tokens_across_devices = True,
max_grad_norm = 1.0,
label_smoothing_factor = 0.0,
bf16 = False,
fp16 = False,
bf16_full_eval = False,
fp16_full_eval = False,
tf32 = None,
gradient_checkpointing = True,
gradient_checkpointing_kwargs = None,
torch_compile = False,
torch_compile_backend = None,
torch_compile_mode = None,
use_liger_kernel = False,
liger_kernel_config = None,
use_cache = False,
neftune_noise_alpha = None,
torch_empty_cache_steps = 250,
auto_find_batch_size = False,
logging_strategy = 'steps',
logging_steps = 1,
logging_first_step = False,
log_on_each_node = True,
logging_nan_inf_filter = False,
include_num_input_tokens_seen = False,
log_level = 'passive',
log_level_replica = 'warning',
disable_tqdm = None,
report_to = 'none',
run_name = None,
project = 'huggingface',
trackio_space_id = 'trackio',
eval_strategy = 'no',
eval_steps = None,
eval_delay = 0,
per_device_eval_batch_size = 4,
prediction_loss_only = False,
eval_on_start = False,
eval_do_concat_batches = True,
eval_use_gather_object = False,
eval_accumulation_steps = 2,
batch_eval_metrics = False,
save_only_model = False,
save_strategy = 'steps',
save_steps = 500,
save_on_each_node = False,
save_total_limit = None,
enable_jit_checkpoint = False,
push_to_hub = False,
hub_token = None,
hub_private_repo = None,
hub_model_id = None,
hub_strategy = 'every_save',
hub_always_push = False,
hub_revision = None,
load_best_model_at_end = False,
metric_for_best_model = None,
greater_is_better = None,
ignore_data_skip = False,
restore_callback_states_from_checkpoint = False,
full_determinism = False,
seed = 3407,
data_seed = 3407,
use_cpu = False,
accelerator_config = None,
parallelism_config = None,
dataloader_drop_last = False,
dataloader_num_workers = 0,
dataloader_pin_memory = True,
dataloader_persistent_workers = False,
dataloader_prefetch_factor = None,
remove_unused_columns = True,
label_names = None,
train_sampling_strategy = 'sequential',
length_column_name = 'length',
ddp_find_unused_parameters = None,
ddp_bucket_cap_mb = None,
ddp_broadcast_buffers = None,
ddp_backend = None,
ddp_timeout = 1800,
fsdp = None,
fsdp_config = None,
deepspeed = None,
debug = '',
skip_memory_metrics = True,
do_train = False,
do_eval = False,
do_predict = False,
resume_from_checkpoint = None,
warmup_ratio = None,
logging_dir = None,
local_rank = -1,
model_init_kwargs = None,
trust_remote_code = False,
disable_dropout = True,
dataset_num_proc = None,
max_length = 1024,
pad_to_multiple_of = None,
precompute_ref_log_probs = False,
precompute_ref_batch_size = None,
loss_type = 'kto',
beta = 0.1,
desirable_weight = 1.0,
undesirable_weight = 1.0,
activation_offloading = False,
sync_ref_model = False,
ref_model_mixup_alpha = 0.6,
ref_model_sync_steps = 512,
vllm_sampling_params = None,
unsloth_num_chunks = -1,
unsloth_logit_chunk_multiplier = None,
unsloth_grpo_mini_batch = None,
max_seq_length = None,
**kwargs,
):
if learning_rate < 1e-7: print(f'Unsloth: Your learning rate of `{learning_rate}` is too small and less than 1e-7! Consider increasing it, otherwise gradient updates will be close to 0!')
if learning_rate > 1: print(f'Unsloth: Your learning rate of `{learning_rate}` is way too larger > 1! Consider decreasing it to 1e-1, otherwise gradient updates will explode!')
if num_train_epochs is None:
num_train_epochs = 3.0 # Default to 3 epochs if None, max_steps will override
if output_dir is None and save_strategy == 'steps' and save_steps == 500:
output_dir = 'unsloth_training_checkpoints'
save_strategy = 'no'
import multiprocessing as _mp
if dataset_num_proc is None:
if _mp.get_start_method() != 'fork':
dataset_num_proc = None
else:
import psutil
dataset_num_proc = min(max((psutil.cpu_count() or 1)+4, 2), 64)
memory_gb_left = psutil.virtual_memory().available / (1024**3)
if memory_gb_left <= 2: dataset_num_proc = 1
else: dataset_num_proc = min(dataset_num_proc, int(memory_gb_left))
if os.environ.get('UNSLOTH_ENABLE_FLEX_ATTENTION', '0') == '1':
from unsloth_zoo.flex_attention import HAS_FLEX_ATTENTION
if HAS_FLEX_ATTENTION and pad_to_multiple_of is None:
from unsloth_zoo.flex_attention import FLEX_ATTENTION_BLOCK_SIZE
pad_to_multiple_of = FLEX_ATTENTION_BLOCK_SIZE
super().__init__(
output_dir = output_dir,
per_device_train_batch_size = per_device_train_batch_size,
num_train_epochs = num_train_epochs,
max_steps = max_steps,
learning_rate = learning_rate,
lr_scheduler_type = lr_scheduler_type,
lr_scheduler_kwargs = lr_scheduler_kwargs,
warmup_steps = warmup_steps,
optim = optim,
optim_args = optim_args,
weight_decay = weight_decay,
adam_beta1 = adam_beta1,
adam_beta2 = adam_beta2,
adam_epsilon = adam_epsilon,
optim_target_modules = optim_target_modules,
gradient_accumulation_steps = gradient_accumulation_steps,
average_tokens_across_devices = average_tokens_across_devices,
max_grad_norm = max_grad_norm,
label_smoothing_factor = label_smoothing_factor,
bf16 = bf16,
fp16 = fp16,
bf16_full_eval = bf16_full_eval,
fp16_full_eval = fp16_full_eval,
tf32 = tf32,
gradient_checkpointing = gradient_checkpointing,
gradient_checkpointing_kwargs = gradient_checkpointing_kwargs,
torch_compile = torch_compile,
torch_compile_backend = torch_compile_backend,
torch_compile_mode = torch_compile_mode,
use_liger_kernel = use_liger_kernel,
liger_kernel_config = liger_kernel_config,
use_cache = use_cache,
neftune_noise_alpha = neftune_noise_alpha,
torch_empty_cache_steps = torch_empty_cache_steps,
auto_find_batch_size = auto_find_batch_size,
logging_strategy = logging_strategy,
logging_steps = logging_steps,
logging_first_step = logging_first_step,
log_on_each_node = log_on_each_node,
logging_nan_inf_filter = logging_nan_inf_filter,
include_num_input_tokens_seen = include_num_input_tokens_seen,
log_level = log_level,
log_level_replica = log_level_replica,
disable_tqdm = disable_tqdm,
report_to = report_to,
run_name = run_name,
project = project,
trackio_space_id = trackio_space_id,
eval_strategy = eval_strategy,
eval_steps = eval_steps,
eval_delay = eval_delay,
per_device_eval_batch_size = per_device_eval_batch_size,
prediction_loss_only = prediction_loss_only,
eval_on_start = eval_on_start,
eval_do_concat_batches = eval_do_concat_batches,
eval_use_gather_object = eval_use_gather_object,
eval_accumulation_steps = eval_accumulation_steps,
batch_eval_metrics = batch_eval_metrics,
save_only_model = save_only_model,
save_strategy = save_strategy,
save_steps = save_steps,
save_on_each_node = save_on_each_node,
save_total_limit = save_total_limit,
enable_jit_checkpoint = enable_jit_checkpoint,
push_to_hub = push_to_hub,
hub_token = hub_token,
hub_private_repo = hub_private_repo,
hub_model_id = hub_model_id,
hub_strategy = hub_strategy,
hub_always_push = hub_always_push,
hub_revision = hub_revision,
load_best_model_at_end = load_best_model_at_end,
metric_for_best_model = metric_for_best_model,
greater_is_better = greater_is_better,
ignore_data_skip = ignore_data_skip,
restore_callback_states_from_checkpoint = restore_callback_states_from_checkpoint,
full_determinism = full_determinism,
seed = seed,
data_seed = data_seed,
use_cpu = use_cpu,
accelerator_config = accelerator_config,
parallelism_config = parallelism_config,
dataloader_drop_last = dataloader_drop_last,
dataloader_num_workers = dataloader_num_workers,
dataloader_pin_memory = dataloader_pin_memory,
dataloader_persistent_workers = dataloader_persistent_workers,
dataloader_prefetch_factor = dataloader_prefetch_factor,
remove_unused_columns = remove_unused_columns,
label_names = label_names,
train_sampling_strategy = train_sampling_strategy,
length_column_name = length_column_name,
ddp_find_unused_parameters = ddp_find_unused_parameters,
ddp_bucket_cap_mb = ddp_bucket_cap_mb,
ddp_broadcast_buffers = ddp_broadcast_buffers,
ddp_backend = ddp_backend,
ddp_timeout = ddp_timeout,
fsdp = fsdp,
fsdp_config = fsdp_config,
deepspeed = deepspeed,
debug = debug,
skip_memory_metrics = skip_memory_metrics,
do_train = do_train,
do_eval = do_eval,
do_predict = do_predict,
resume_from_checkpoint = resume_from_checkpoint,
warmup_ratio = warmup_ratio,
logging_dir = logging_dir,
local_rank = local_rank,
model_init_kwargs = model_init_kwargs,
trust_remote_code = trust_remote_code,
disable_dropout = disable_dropout,
dataset_num_proc = dataset_num_proc,
max_length = max_length,
pad_to_multiple_of = pad_to_multiple_of,
precompute_ref_log_probs = precompute_ref_log_probs,
precompute_ref_batch_size = precompute_ref_batch_size,
loss_type = loss_type,
beta = beta,
desirable_weight = desirable_weight,
undesirable_weight = undesirable_weight,
activation_offloading = activation_offloading,
sync_ref_model = sync_ref_model,
ref_model_mixup_alpha = ref_model_mixup_alpha,
ref_model_sync_steps = ref_model_sync_steps,**kwargs)
self.vllm_sampling_params = vllm_sampling_params
self.unsloth_num_chunks = unsloth_num_chunks
if unsloth_grpo_mini_batch is not None:
if self.generation_batch_size >= unsloth_grpo_mini_batch:
self.unsloth_grpo_mini_batch = unsloth_grpo_mini_batch
else:
raise ValueError(
f"Unsloth GRPO mini batch size needs to be less than or equal to the effective generation batch size, "
f"which is self.per_device_train_batch_size * gradient_accumulation_steps."
)
self.unsloth_logit_chunk_multiplier = unsloth_logit_chunk_multiplier
self.max_seq_length = max_seq_length
# Unsloth: Remove use_reentrant=False forced by TRL 0.27.0+
if getattr(self, 'gradient_checkpointing_kwargs', None) is not None:
if 'use_reentrant' in self.gradient_checkpointing_kwargs:
del self.gradient_checkpointing_kwargs['use_reentrant']
pass
class _UnslothKTOTrainer(_BaseTrainer):
"""
Initialize KTOTrainer.
Args:
model (`str` or [`~transformers.PreTrainedModel`] or [`~peft.PeftModel`]):
Model to be trained. Can be either:
- A string, being the *model id* of a pretrained model hosted inside a model repo on huggingface.co, or a
path to a *directory* containing model weights saved using
[`~transformers.PreTrainedModel.save_pretrained`], e.g., `'./my_model_directory/'`. The model is loaded
using `.from_pretrained` (where `` is derived from the model
config) with the keyword arguments in `args.model_init_kwargs`.
- A [`~transformers.PreTrainedModel`] object. Only causal language models are supported.
- A [`~peft.PeftModel`] object. Only causal language models are supported.
ref_model ([`~transformers.PreTrainedModel`], *optional*):
Reference model used to compute the reference log probabilities.
- If provided, this model is used directly as the reference policy.
- If `None`, the trainer will automatically use the initial policy corresponding to `model`, i.e. the model
state before KTO training starts.
args ([`experimental.kto.KTOConfig`], *optional*):
Configuration for this trainer. If `None`, a default configuration is used.
train_dataset ([`~datasets.Dataset`] or [`~datasets.IterableDataset`]):
The dataset to use for training.
eval_dataset ([`~datasets.Dataset`], [`~datasets.IterableDataset`] or `dict[str, Dataset | IterableDataset]`):
The dataset to use for evaluation.
processing_class ([`~transformers.PreTrainedTokenizerBase`] or [`~transformers.ProcessorMixin`], *optional*):
Processing class used to process the data. The padding side must be set to "left". If `None`, the
processing class is loaded from the model's name with [`~transformers.AutoProcessor.from_pretrained`]. A
padding token, `tokenizer.pad_token`, must be set. If the processing class has not set a padding token,
`tokenizer.eos_token` will be used as the default.
data_collator ([`~transformers.DataCollator`], *optional*):
The data collator to use for training. If None is specified, the default data collator
([`~experimental.kto.kto_trainer.DataCollatorForUnpairedPreference`]) will be used which will pad the
sequences to the maximum length of the sequences in the batch.
callbacks (`list[transformers.TrainerCallback]`):
The callbacks to use for training.
optimizers (`tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR]`):
The optimizer and scheduler to use for training.
peft_config ([`~peft.PeftConfig`], *optional*):
PEFT configuration used to wrap the model. If `None`, the model is not wrapped.
compute_metrics (`Callable[[EvalPrediction], dict]`, *optional*):
The function to use to compute the metrics. Must take a `EvalPrediction` and return a dictionary string to
metric values.
"""
_tag_names = ["trl", "kto"]
_name = "KTO"
_paper = {
"title": "KTO: Model Alignment as Prospect Theoretic Optimization",
"id": "2402.01306",
# docstyle-ignore
"citation": textwrap.dedent("""\
@article{ethayarajh2024kto,
title = {{KTO: Model Alignment as Prospect Theoretic Optimization}},
author = {Kawin Ethayarajh and Winnie Xu and Niklas Muennighoff and Dan Jurafsky and Douwe Kiela},
year = 2024,
eprint = {arXiv:2402.01306},
}"""),
}
def __init__(
self,
model: "str | PreTrainedModel | PeftModel",
ref_model: PreTrainedModel | None = None,
args: KTOConfig | None = None,
train_dataset: Dataset | IterableDataset | None = None,
eval_dataset: Dataset | IterableDataset | dict[str, Dataset | IterableDataset] | None = None,
processing_class: PreTrainedTokenizerBase | ProcessorMixin | None = None,
data_collator: DataCollator | None = None,
callbacks: list[TrainerCallback] | None = None,
optimizers: tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),
peft_config: "PeftConfig | None" = None,
compute_metrics: Callable[[EvalLoopOutput], dict] | None = None,
):
# Args
if args is None:
model_name = model if isinstance(model, str) else get_config_model_id(model.config)
model_name = model_name.split("/")[-1]
args = KTOConfig(f"{model_name}-KTO")
if train_dataset is None:
raise ValueError("`train_dataset` is required")
elif isinstance(train_dataset, IterableDataset):
# IterableDataset requires dispatch_batches=False because Accelerate's dispatch mode may try to concatenate
# batches from multiple processes, leading to mismatch errors.
if args.accelerator_config.dispatch_batches is True:
logger.warning(
"You are using an `IterableDataset` for training with `dispatch_batches=True`. `dispatch_batches` "
"is forced to `False` when using an `IterableDataset`. To remove this warning, unset "
"`dispatch_batches` in `KTOConfig` or set it to `False`."
)
args.accelerator_config.dispatch_batches = False
# Model
if isinstance(model, str):
model_init_kwargs = args.model_init_kwargs or {}
# Distributed training requires device_map=None ["auto" fails]
if args.distributed_state.distributed_type in ["MULTI_GPU", "DEEPSPEED"]:
model_init_kwargs["device_map"] = None
model_init_kwargs.setdefault("trust_remote_code", args.trust_remote_code)
model = create_model_from_path(model, **model_init_kwargs)
else:
if args.model_init_kwargs is not None:
logger.warning(
"You passed `model_init_kwargs` to the KTOConfig, but your model is already instantiated. "
"The `model_init_kwargs` will be ignored."
)
# Non-quantized models do not have the `is_loaded_in_{8,4}bit` attributes, whereas quantized models do
_is_quantized_model = getattr(model, "is_loaded_in_4bit", False) or getattr(model, "is_loaded_in_8bit", False)
if ref_model is model:
raise ValueError(
"`model` and `ref_model` cannot be the same object. In most cases you should omit `ref_model` and "
"we'll initialize it to a copy of `model` for you."
)
# Processing class
if processing_class is None:
processing_class = AutoProcessor.from_pretrained(
get_config_model_id(model.config), trust_remote_code=args.trust_remote_code
)
if isinstance(processing_class, ProcessorMixin):
self._tokenizer = processing_class.tokenizer
self._is_vlm = True
elif isinstance(processing_class, PreTrainedTokenizerBase):
self._tokenizer = processing_class
self._is_vlm = False
else:
raise TypeError("The `processing_class` must be either a `PreTrainedTokenizerBase` or a `ProcessorMixin`")
if self._tokenizer.pad_token is None:
self._tokenizer.pad_token = self._tokenizer.eos_token
# PEFT
if False:
if not is_peft_available():
raise ImportError(
"You passed `peft_config` but the `peft` library is not installed. "
"Install it with `pip install trl[peft]`."
)
if not isinstance(peft_config, PeftConfig):
raise TypeError(
f"`peft_config` must be a `peft.PeftConfig` instance (e.g. `peft.LoraConfig`), "
f"got {type(peft_config).__name__}."
)
if is_peft_model(model):
raise ValueError(
"You passed a `PeftModel` instance together with a `peft_config` to the trainer. Please first merge "
"and unload the existing adapter, save the resulting base model, and then pass that base model along "
"with the new `peft_config` to the trainer."
)
# Create PEFT model
# ZeRO-3 + PEFT for non-quantized models:
# - PEFT's default autocast_adapter_dtype=True upcasts LoRA adapter params to fp32 even when the base model is bf16.
# - ZeRO-3's _allgather_params_coalesced allocates output buffers using the dtype of the first persistent parameter,
# so mixed-dtype persistent_parameters [bf16 base + fp32 LoRA] cause a TypeError on the first optimizer step.
# - Passing autocast_adapter_dtype=False keeps adapter params in the base model dtype [bf16], fixing the mismatch.
# - This is safe: the fp32 upcast is a QLoRA-specific concern [low-bit quantized base models], not needed for
# non-quantized bf16 training.
# - See:
# - TRL issue: https://github.com/huggingface/trl/issues/6089
# - Upstream issue: https://github.com/deepspeedai/DeepSpeed/issues/8072
# - autocast_adapter_dtype was introduced in PEFT 0.12.0; before, no upcast existed: no need to pass the kwarg
get_peft_model_kwargs = {}
if (
args.deepspeed_plugin is not None
and args.deepspeed_plugin.zero_stage == 3
and not _is_quantized_model
and Version(peft.__version__) >= Version("0.12.0")
):
get_peft_model_kwargs["autocast_adapter_dtype"] = False
model = get_peft_model(model, peft_config, **get_peft_model_kwargs)
elif is_peft_model(model) and ref_model is None:
# If the model is a PEFT model with a pretrained adapter, we need to create a "ref" adapter that is a copy
# of the "default" adapter, so that we can use it as the reference model during KTO training. PEFT only
# supports one adapter per model when the LoRA config uses `target_parameters` [see peft#3340], so in that
# case we skip the "ref" adapter and compute the reference log probs with adapters disabled, i.e. with the
# base model.
default_config = model.peft_config["default"]
if isinstance(default_config, LoraConfig) and default_config.target_parameters:
logger.warning(
"PEFT can't add a frozen reference adapter alongside one that uses `target_parameters` "
"(peft#3340], so the reference log probs are computed from the base model [adapters disabled]. "
"If you wrapped the model only to apply LoRA, pass a `peft_config` to the trainer instead; if you "
"wrapped it deliberately (pretrained adapter or custom init), note that the base model matches "
"your adapter only when it's freshly zero-initialized. If it is, this warning is safe to ignore."
)
else:
model.add_adapter("ref", default_config)
for name, param in model.named_parameters():
if ".default." in name:
ref_name = name.replace(".default.", ".ref.")
ref_param = model.get_parameter(ref_name)
ref_param.data.copy_(param.data)
# When using gradient checkpointing with PEFT, we need to enable input gradients. transformers.Trainer normally
# handles this, but a bug currently prevents it; see https://github.com/huggingface/transformers/issues/42489
if is_peft_model(model) and args.gradient_checkpointing:
model.enable_input_require_grads()
# When using QLoRA, the PEFT adapter weights are converted to bf16 to follow the recommendations from the
# original paper [see https://huggingface.co/papers/2305.14314, paragraph 3]. Normally, this can be done by
# passing `autocast_adapter_dtype=False` to `get_peft_model`, but this option is not yet supported for
# quantized models. See: https://github.com/huggingface/peft/issues/2889
if _is_quantized_model:
for param in model.parameters():
if param.requires_grad:
param.data = param.data.to(torch.bfloat16)
# Vision dataset detection
dataset_sample = next(iter(train_dataset))
self._is_vision_dataset = "image" in dataset_sample or "images" in dataset_sample
if self._is_vision_dataset and not self._is_vlm:
raise ValueError(
"The dataset appears to be vision-related (contains 'image' or 'images' keys), but the provided "
"model does not seem to be a vision-language model. Please check your model and dataset."
)
if self._is_vision_dataset and args.precompute_ref_log_probs:
raise ValueError(
"`precompute_ref_log_probs=True` is not supported for vision datasets. For vision-language "
"models, all data processing is performed on the fly rather than upfront. "
"Set `precompute_ref_log_probs=False`."
)
if self._is_vision_dataset and ("chosen" in dataset_sample or "rejected" in dataset_sample):
raise ValueError(
"Vision datasets must be in unpaired format with `completion` and `label` columns. "
"Paired format (`chosen`/`rejected`) is not supported for vision datasets because "
"iterating over the full dataset to unpair it would be too expensive for large image "
"collections. Unpair your dataset first: `dataset = unpair_preference_dataset(dataset)`."
)
# Data collator
calculate_kl = args.loss_type not in ["apo_zero_unpaired"]
if data_collator is None and not self._is_vision_dataset:
data_collator = DataCollatorForUnpairedPreference(
pad_token_id=self._tokenizer.pad_token_id,
max_length=args.max_length,
pad_to_multiple_of=args.pad_to_multiple_of,
)
elif data_collator is None and self._is_vision_dataset:
data_collator = DataCollatorForVisionUnpairedPreference(
processor=processing_class,
max_length=args.max_length,
calculate_kl=calculate_kl,
pad_to_multiple_of=args.pad_to_multiple_of,
)
# Training arguments
self.beta = args.beta
self.precompute_ref_logps = args.precompute_ref_log_probs
self.loss_type = args.loss_type
self.desirable_weight = args.desirable_weight
self.undesirable_weight = args.undesirable_weight
self.aux_loss_enabled = getattr(model.config, "output_router_logits", False)
self.aux_loss_coef = getattr(model.config, "router_aux_loss_coef", 0.0)
self.calculate_KL = calculate_kl
if self.calculate_KL and args.train_sampling_strategy != "sequential":
raise ValueError(
f"Loss type `'{args.loss_type}'` estimates the KL divergence term and requires "
f"`train_sampling_strategy='sequential'` because the KL completion for each example is precomputed "
f"against its neighbors in a fixed-order batch; any other strategy breaks that pairing. "
f"Got `train_sampling_strategy='{args.train_sampling_strategy}'`."
)
if self.calculate_KL and args.per_device_train_batch_size <= 1:
raise ValueError(
"Actual (not effective) batch size must be > 1. KTO will not work properly because the KL term will be equivalent to the implied reward."
)
if self.aux_loss_enabled and self.aux_loss_coef == 0.0:
logger.warning(
"You set `output_router_logits` to `True` in the model config, but `router_aux_loss_coef` is set to "
"`0.0`, meaning the auxiliary loss will not be used. Either set `router_aux_loss_coef` to a value "
"greater than `0.0`, or set `output_router_logits` to `False` if you don't want to use the auxiliary "
"loss.",
)
# Dataset
# Skip dataset preparation for VLMs: tokenization and image processing happen on-the-fly in the collator.
if not self._is_vision_dataset:
train_dataset = self._prepare_dataset(train_dataset, processing_class, args, "train")
if eval_dataset is not None:
if isinstance(eval_dataset, dict):
eval_dataset = {
key: self._prepare_dataset(dataset, processing_class, args, key)
for key, dataset in eval_dataset.items()
}
else:
eval_dataset = self._prepare_dataset(eval_dataset, processing_class, args, "eval")
# Transformers explicitly set use_reentrant=True in the past to silence a PyTorch warning, but the default was
# never updated once PyTorch switched to recommending use_reentrant=False. Until that change lands upstream
# [see https://github.com/huggingface/transformers/pull/43203] and is released [most likely in 5.0.0], we
# default to the recommended non-reentrant behavior here, while preserving any user-provided value.
if args.gradient_checkpointing and Version(transformers.__version__) < Version("5.0.0"):
args.gradient_checkpointing_kwargs = args.gradient_checkpointing_kwargs or {}
args.gradient_checkpointing_kwargs.setdefault("use_reentrant", False)
super().__init__(
model=model,
args=args,
data_collator=data_collator,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
processing_class=processing_class,
compute_metrics=compute_metrics,
callbacks=callbacks,
optimizers=optimizers,
)
# Initialize activation offloading context
if self.args.activation_offloading:
self.maybe_activation_offload_context = get_act_offloading_ctx_manager(model=self.model)
else:
self.maybe_activation_offload_context = contextlib.nullcontext()
# Reference model
if ref_model is None:
if is_peft_model(self.model) or args.precompute_ref_log_probs:
# If PEFT is used, the reference model is not needed since the adapter can be disabled to revert to the
# initial model. If precompute_ref_log_probs is True, the reference model does not need to be kept in
# memory during training.
self.ref_model = None
else:
ref_model_init_kwargs = args.model_init_kwargs or {}
# Distributed training requires device_map=None ["auto" fails]
if self.args.distributed_state.distributed_type in ["MULTI_GPU", "DEEPSPEED"]:
ref_model_init_kwargs["device_map"] = None
ref_model_init_kwargs.setdefault("trust_remote_code", args.trust_remote_code)
ref_model_path = get_config_model_id(self.model.config)
self.ref_model = create_model_from_path(ref_model_path, **ref_model_init_kwargs)
else:
self.ref_model = ref_model
# Disable dropout in the model and reference model
if args.disable_dropout:
disable_dropout_in_model(model)
if self.ref_model is not None:
disable_dropout_in_model(self.ref_model)
# Initialize the metrics
self._metrics = {"train": defaultdict(list), "eval": defaultdict(list)}
# Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the
# model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set
# self.model_accepts_loss_kwargs to False to enable scaling.
self.model_accepts_loss_kwargs = False
# Add tags to the model
self.model.add_model_tags(self._tag_names)
if self.ref_model is not None:
if self.is_deepspeed_enabled:
self.ref_model = prepare_deepspeed(self.ref_model, self.accelerator)
elif self.is_fsdp_enabled:
self.ref_model = prepare_fsdp(self.ref_model, self.accelerator)
else:
self.ref_model = self.accelerator.prepare_model(self.ref_model, evaluation_mode=True)
if args.sync_ref_model:
if is_peft_model(self.model):
raise NotImplementedError(
"You passed `sync_ref_model=True` while using a PEFT model, which is currently not supported. "
"With PEFT, KTOTrainer does not keep a separate reference model in memory; instead, it recovers "
"reference behavior by temporarily disabling the adapter. As a result, there is no standalone "
"`ref_model` instance to synchronize. Use `sync_ref_model=False`, or opt for full fine-tuning if "
"you need a synced reference model. If you need `sync_ref_model` to work with PEFT, please open a "
"feature request at https://github.com/huggingface/trl/issues."
)
if args.precompute_ref_log_probs:
raise ValueError(
"You cannot use `sync_ref_model=True` together with `precompute_ref_log_probs=True`. "
"`precompute_ref_log_probs=True` assumes a fixed reference model, but with `sync_ref_model=True` "
"the reference model is periodically updated during training, making any precomputed reference "
"log-probs stale. Set `precompute_ref_log_probs=False` or disable `sync_ref_model`."
)
self.add_callback(SyncRefModelCallback(ref_model=self.ref_model, accelerator=self.accelerator))
self.use_liger_kernel = args.use_liger_kernel
# Import Liger kernel if enabled
if self.use_liger_kernel:
if not is_liger_kernel_available():
raise ImportError(
"You set `use_liger_kernel=True` but the liger kernel is not available. "
"Please install liger-kernel first: `pip install liger-kernel`"
)
if self.loss_type in ["apo_zero_unpaired"]:
raise ValueError(
"You cannot set `loss_type='apo_zero_unpaired'` with liger-kernel."
"Only KTO loss is supported with liger-kernel."
)
if self.precompute_ref_logps:
raise ValueError(
"You cannot use `precompute_ref_log_probs=True` with liger kernel. Please set "
"`precompute_ref_log_probs=False`."
)
if is_peft_model(self.model):
raise ValueError(
"You cannot use `use_liger_kernel=True` with Peft models. Please set `use_liger_kernel=False`."
)
self.liger_loss_fn = LigerFusedLinearKTOLoss(beta=self.beta, use_ref_model=(self.ref_model is not None))
if self.precompute_ref_logps:
if isinstance(self.train_dataset, IterableDataset) or isinstance(
self.eval_dataset, (IterableDataset, IterableDatasetDict)
):
raise ValueError(
"`precompute_ref_log_probs=True` is not supported with IterableDataset. Please use a map-style "
"Dataset or set `precompute_ref_log_probs=False`."
)
self.train_dataset = self._precompute_ref_logps(
self.train_dataset,
"train",
self.args.precompute_ref_batch_size or self.args.per_device_train_batch_size,
)
if self.eval_dataset is not None:
if isinstance(self.eval_dataset, dict):
self.eval_dataset = {
name: self._precompute_ref_logps(
dataset, name, self.args.precompute_ref_batch_size or self.args.per_device_eval_batch_size
)
for name, dataset in self.eval_dataset.items()
}
else:
self.eval_dataset = self._precompute_ref_logps(
self.eval_dataset,
"eval",
self.args.precompute_ref_batch_size or self.args.per_device_eval_batch_size,
)
def _tokenize(
self,
processing_class: PreTrainedTokenizerBase | ProcessorMixin,
input: str | list,
**kwargs,
) -> dict[str, list]:
"""Tokenize a single example for dataset preprocessing.
Dispatches to `apply_chat_template` for conversational input (list of message dicts) and to `__call__` for
non-conversational input (str).
Args:
processing_class ([`~transformers.PreTrainedTokenizerBase`] or [`~transformers.ProcessorMixin`]):
The tokenizer or processor to use.
input (`str` or `list`):
A string for non-conversational input, or a list of message dicts for conversational input.
**kwargs:
Forwarded to `apply_chat_template` (e.g. `add_generation_prompt`, `return_assistant_tokens_mask`).
Returns:
`dict` with at least an `"input_ids"` key mapping to a flat `list[int]`.
"""
if isinstance(input, list): # conversational: list of message dicts
if self._is_vlm:
input = prepare_multimodal_messages(input)
result = processing_class.apply_chat_template(input, tokenize=True, return_dict=True, **kwargs)
else: # non-conversational: plain text string
result = processing_class(text=input)
# VLMs emit a batch dimension even for single examples; unwrap it
if self._is_vlm:
return {k: v[0] for k, v in result.items()}
return result
def _get_kl_dataset(
self,
dataset: Dataset | IterableDataset,
dataset_name: str,
args: KTOConfig,
) -> Dataset | IterableDataset:
"""
Creates the KL dataset by creating mismatched (prompt, completion) pairs for KL divergence estimation.
Args:
dataset (`Dataset` or `IterableDataset`):
Tokenized dataset with `prompt_ids` and `completion_ids` columns.
dataset_name (`str`):
Name used in progress bar descriptions.
args ([`KTOConfig`]):
Training arguments providing `per_device_train_batch_size` and `dataset_num_proc`.
Returns:
`Dataset` or `IterableDataset` with a single `KL_completion_ids` column.
"""
map_kwargs = {}
if isinstance(dataset, Dataset): # IterableDataset does not support num_proc or desc
map_kwargs["num_proc"] = args.dataset_num_proc
map_kwargs["desc"] = f"Extracting KL {dataset_name} dataset"
kl_dataset = dataset.map(
_get_kl_completion_ids, batched=True, batch_size=args.per_device_train_batch_size, **map_kwargs
)
def rename_kl_fn(example):
return {"KL_completion_ids": example["completion_ids"]}
if isinstance(dataset, Dataset): # `IterableDataset.map` does not support `desc`
map_kwargs["desc"] = f"Assembling KL {dataset_name} dataset"
column_names = get_dataset_column_names(dataset)
kl_dataset = kl_dataset.map(
rename_kl_fn,
remove_columns=[c for c in get_dataset_column_names(kl_dataset) if c in column_names],
**map_kwargs,
)
return kl_dataset
def _prepare_dataset(
self,
dataset: Dataset | IterableDataset,
processing_class: PreTrainedTokenizerBase | ProcessorMixin,
args: KTOConfig | None,
dataset_name: str,
) -> Dataset | IterableDataset:
# Build the kwargs for the `map` function
map_kwargs = {}
if isinstance(dataset, Dataset): # IterableDataset does not support num_proc
map_kwargs["num_proc"] = args.dataset_num_proc
# Compute that only on the main process for faster data processing.
# see: https://github.com/huggingface/trl/pull/1255
with PartialState().main_process_first():
# Extract the prompt if needed
first_example = next(iter(dataset))
if "prompt" not in first_example:
if isinstance(dataset, Dataset): # `IterableDataset.map` does not support `desc`
map_kwargs["desc"] = f"Extracting prompt from {dataset_name} dataset"
dataset = dataset.map(extract_prompt, **map_kwargs)
# Unpair the dataset if needed
first_example = next(iter(dataset))
if "chosen" in first_example and "rejected" in first_example:
if isinstance(dataset, Dataset): # `IterableDataset.map` does not support `desc`
map_kwargs["desc"] = f"Unpairing {dataset_name} dataset"
dataset = unpair_preference_dataset(dataset, **map_kwargs)
# Add EOS token if needed: non-conversational only
first_example = next(iter(dataset))
if not is_conversational(first_example):
if isinstance(dataset, Dataset): # `IterableDataset.map` does not support `desc`
map_kwargs["desc"] = f"Adding EOS to {dataset_name} dataset"
def add_eos(example, eos_token):
if not example["completion"].endswith(eos_token):
example["completion"] = example["completion"] + eos_token
return example
dataset = dataset.map(add_eos, fn_kwargs={"eos_token": self._tokenizer.eos_token}, **map_kwargs)
# Tokenize dataset
if isinstance(dataset, Dataset): # `IterableDataset.map` does not support `desc`
map_kwargs["desc"] = f"Tokenizing {dataset_name} dataset"
def tokenize_fn(example, processing_class):
if is_conversational(example):
chat_template_kwargs = example.get("chat_template_kwargs", {})
prompt_ids = self._tokenize(
processing_class,
example["prompt"],
add_generation_prompt=True,
**chat_template_kwargs,
)["input_ids"]
prompt_completion_ids = self._tokenize(
processing_class,
example["prompt"] + example["completion"],
**chat_template_kwargs,
)["input_ids"]
else:
prompt_ids = self._tokenize(processing_class, example["prompt"])["input_ids"]
prompt_completion_ids = self._tokenize(
processing_class, example["prompt"] + example["completion"]
)["input_ids"]
if not prompt_completion_ids[: len(prompt_ids)] == prompt_ids:
logger.warning(
"Mismatch between tokenized prompt and the start of tokenized prompt+completion. "
"This may be due to unexpected tokenizer behavior, whitespace issues, or special "
"token handling. Verify that the tokenizer is processing text consistently."
)
return {
"prompt_ids": prompt_ids,
"completion_ids": prompt_completion_ids[len(prompt_ids) :],
}
dataset = dataset.map(tokenize_fn, fn_kwargs={"processing_class": processing_class}, **map_kwargs)
# Get KL datasets if needed
if self.calculate_KL:
# create pairs for estimating the KL term by flipping the matched pairs in each batch of size total_batch_size
# i.e., (x_1, y_1), ..., (x_n, y_n) --> (x_1, y_n), ..., (x_n, y_1) = (x'_1, y'_1), ..., (x'_n, y'_n)
kl_dataset = self._get_kl_dataset(dataset, dataset_name, args)
dataset = concatenate_datasets([dataset, kl_dataset], axis=1)
# Calculate dataset desirability balance
if dataset_name == "train" and isinstance(dataset, Dataset): # IterableDataset does not support len
num_desirable = max(sum(dataset["label"]), 1)
num_undesirable = max(len(dataset["label"]) - num_desirable, 1) # "label" is binary
if num_desirable != num_undesirable:
# The lower and upper bounds come from Eq. (8) of https://huggingface.co/papers/2402.01306
des_weight_lower_bound = round((num_undesirable * self.undesirable_weight / num_desirable) * 1, 2)
des_weight_upper_bound = round(
(num_undesirable * self.undesirable_weight / num_desirable) * 1.33, 2
)
und_weight_lower_bound = round((num_desirable * self.desirable_weight / num_undesirable) / 1.33, 2)
und_weight_upper_bound = round((num_desirable * self.desirable_weight / num_undesirable) / 1, 2)
des_weight_in_range = des_weight_lower_bound <= self.desirable_weight <= des_weight_upper_bound
und_weight_in_range = und_weight_lower_bound <= self.undesirable_weight <= und_weight_upper_bound
if not (des_weight_in_range or und_weight_in_range):
logger.warning(
"You have different amounts of desirable/positive and undesirable/negative examples but the "
"weights on the desirable and undesirable losses don't seem to be in an ideal range. Based "
f"on your data, we recommend EITHER "
f"desirable_weight in [{des_weight_lower_bound}, {des_weight_upper_bound}] or "
f"undesirable_weight in [{und_weight_lower_bound}, {und_weight_upper_bound}] (but NOT BOTH). "
"See the documentation on how to optimally set these weights.",
)
return dataset
def _set_signature_columns_if_needed(self):
# If `self.args.remove_unused_columns` is True, non-signature columns are removed.
# By default, this method sets `self._signature_columns` to the model's expected inputs (usually, "input_ids"
# and "attention_mask").
if self._signature_columns is None:
if self._is_vision_dataset:
self._signature_columns = [
"prompt",
"completion",
"image",
"images",
"label",
"chat_template_kwargs",
]
else:
self._signature_columns = [
"prompt_ids",
"completion_ids",
"KL_completion_ids",
"label",
"ref_logps",
"ref_KL_logps",
]
def _get_train_sampler(self, train_dataset: Dataset | None = None) -> Sampler | None:
if self.calculate_KL and Version(transformers.__version__) < Version("5.2.0"):
if train_dataset is None:
train_dataset = self.train_dataset
if train_dataset is None or not has_length(train_dataset):
return None
return SequentialSampler(train_dataset)
return super()._get_train_sampler(
train_dataset
) # Override training step to add activation offloading context.
def _precompute_ref_logps(self, dataset: Dataset, name: str, batch_size: int) -> Dataset:
model_hash = hash_module(self.ref_model or self.model)
fingerprint = Hasher.hash((dataset._fingerprint, model_hash, self.calculate_KL))
cache_file = dataset._get_cache_file_path(fingerprint)
if os.path.exists(cache_file):
return concatenate_datasets([dataset, Dataset.from_file(cache_file)], axis=1)
dataloader = DataLoader(
dataset,
batch_size=batch_size,
collate_fn=self.data_collator,
num_workers=self.args.dataloader_num_workers,
pin_memory=self.args.dataloader_pin_memory,
shuffle=False,
)
data_loader = self.accelerator.prepare(dataloader)
ref_logps = []
ref_KL_logps = []
for padded_batch in tqdm(iterable=data_loader, desc=f"Computing reference log probs for {name} dataset"):
ref_logp, ref_KL_logp = self.compute_ref_log_probs(padded_batch)
if self.calculate_KL:
ref_logp, ref_KL_logp = self.accelerator.gather_for_metrics((ref_logp, ref_KL_logp))
ref_KL_logps.append(ref_KL_logp.cpu())
else:
ref_logp = self.accelerator.gather_for_metrics(ref_logp)
ref_logps.append(ref_logp.cpu())
ref_logps = torch.cat(ref_logps)
if self.calculate_KL:
ref_KL_logps = torch.cat(ref_KL_logps)
if self.accelerator.is_main_process:
def add_ref_logps(batch, indices):
result = {"ref_logps": ref_logps[indices]}
if self.calculate_KL:
result.update({"ref_KL_logps": ref_KL_logps[indices]})
return result
dataset.map(
add_ref_logps,
with_indices=True,
batched=True,
remove_columns=dataset.column_names,
new_fingerprint=fingerprint,
desc=f"Caching reference log probs for {name} dataset",
)
self.accelerator.wait_for_everyone()
return concatenate_datasets([dataset, Dataset.from_file(cache_file)], axis=1)
def compute_ref_log_probs(self, inputs):
"""Computes reference log probabilities for a single padded batch."""
with torch.no_grad(), disable_gradient_checkpointing(self.model, self.args.gradient_checkpointing_kwargs):
if self.ref_model is None:
if is_peft_model(self.model):
model = self.accelerator.unwrap_model(self.model)
with use_adapter(model, adapter_name="ref" if "ref" in model.peft_config else None):
completion_logits = self.model(
inputs["completion_input_ids"],
attention_mask=inputs["completion_attention_mask"],
).logits
if self.calculate_KL:
KL_logits = self.model(
inputs["KL_completion_input_ids"],
attention_mask=inputs["KL_completion_attention_mask"],
).logits
else:
completion_logits = self.model(
inputs["completion_input_ids"],
attention_mask=inputs["completion_attention_mask"],
).logits
if self.calculate_KL:
KL_logits = self.model(
inputs["KL_completion_input_ids"],
attention_mask=inputs["KL_completion_attention_mask"],
).logits
else:
completion_logits = self.ref_model(
inputs["completion_input_ids"], attention_mask=inputs["completion_attention_mask"]
).logits
if self.calculate_KL:
KL_logits = self.ref_model(
inputs["KL_completion_input_ids"],
attention_mask=inputs["KL_completion_attention_mask"],
).logits
shift_logits = completion_logits[:, :-1, :]
per_token_logps = selective_log_softmax(shift_logits, inputs["completion_input_ids"][:, 1:])
per_token_logps[inputs["completion_mask"][:, 1:] == 0] = 0.0
completion_logps = per_token_logps.sum(-1)
if self.calculate_KL:
shift_KL_logits = KL_logits[:, :-1, :]
KL_per_token_logps = selective_log_softmax(shift_KL_logits, inputs["KL_completion_input_ids"][:, 1:])
KL_per_token_logps[inputs["KL_completion_mask"][:, 1:] == 0] = 0.0
KL_logps = KL_per_token_logps.sum(-1)
else:
KL_logps = None
return completion_logps, KL_logps
def _compute_kl_logps(self, model, batch):
"""Compute KL log probabilities for a given batch."""
KL_logps = None
if self.calculate_KL:
_non_model_keys = {
"completion_input_ids",
"completion_attention_mask",
"completion_mask",
"KL_completion_mask",
"KL_completion_token_type_ids",
"KL_completion_mm_token_type_ids",
"label",
"ref_logps",
"ref_KL_logps",
}
KL_model_kwargs = {k: v for k, v in batch.items() if k not in _non_model_keys}
KL_model_kwargs["input_ids"] = KL_model_kwargs.pop("KL_completion_input_ids")
KL_model_kwargs["attention_mask"] = KL_model_kwargs.pop("KL_completion_attention_mask")
# KL sequences have different widths from the main completion after flush_left; override token-type
# tensors with the KL-specific ones the collator built for exactly this purpose.
if "KL_completion_token_type_ids" in batch:
KL_model_kwargs["token_type_ids"] = batch["KL_completion_token_type_ids"]
if "KL_completion_mm_token_type_ids" in batch:
KL_model_kwargs["mm_token_type_ids"] = batch["KL_completion_mm_token_type_ids"]
with torch.no_grad():
KL_logits = model(**KL_model_kwargs).logits
shift_KL_logits = KL_logits[:, :-1, :]
KL_per_token_logps = selective_log_softmax(shift_KL_logits, batch["KL_completion_input_ids"][:, 1:])
KL_per_token_logps[batch["KL_completion_mask"][:, 1:] == 0] = 0.0
KL_logps = KL_per_token_logps.sum(-1)
return KL_logps
def _compute_loss_liger(self, model, inputs, return_outputs):
if return_outputs:
raise RuntimeError(
"return_outputs=True is not supported with the Liger KTO loss. The Liger loss computes the loss "
"without materializing logits, so outputs cannot be returned."
)
mode = "train" if self.model.training else "eval"
batch = {k: (v.to(self.accelerator.device) if isinstance(v, torch.Tensor) else v) for k, v in inputs.items()}
labels = torch.tensor(batch["label"])
num_chosen = labels.sum().to(self.accelerator.device)
num_rejected = (len(labels) - num_chosen).to(self.accelerator.device)
policy_KL_logps = self._compute_kl_logps(model, batch)
ref_KL_logps = self._compute_kl_logps(self.ref_model, batch)
if self.calculate_KL:
kl = (policy_KL_logps - ref_KL_logps).mean().detach()
kl = self.accelerator.gather_for_metrics(kl).mean().clamp(min=0)
else:
kl = torch.zeros(1).to(self.accelerator.device)
_non_model_keys = {
"completion_mask",
"KL_completion_input_ids",
"KL_completion_attention_mask",
"KL_completion_mask",
"KL_completion_token_type_ids",
"KL_completion_mm_token_type_ids",
"label",
"ref_logps",
"ref_KL_logps",
}
model_kwargs = {k: v for k, v in batch.items() if k not in _non_model_keys}
model_kwargs["input_ids"] = model_kwargs.pop("completion_input_ids")
model_kwargs["attention_mask"] = model_kwargs.pop("completion_attention_mask")
model_kwargs["use_cache"] = False
if self.aux_loss_enabled:
model_kwargs["output_router_logits"] = True
# `base_model` gives the inner module (skipping `lm_head`) — text decoder for LMs, multimodal wrapper for
# VLMs (so vision-token injection runs before the text decoder). `get_decoder()` won't do: on VLMs it
# returns just the text stack and feeds image-placeholder IDs through it.
# Pre-5.0 transformers VLMs set `base_model_prefix = ""` so `base_model is self` (re-runs `lm_head`).
# Fall back to `.model` there.
if self._is_vlm and Version(transformers.__version__) < Version("5.0.0"):
backbone, ref_backbone = model.model, self.ref_model.model
else:
backbone, ref_backbone = model.base_model, self.ref_model.base_model
outputs = backbone(**model_kwargs)
# reference model
with torch.no_grad(), disable_gradient_checkpointing(self.model, self.args.gradient_checkpointing_kwargs):
ref_outputs = ref_backbone(**{k: v for k, v in model_kwargs.items() if k != "output_router_logits"})
lm_head = model.get_output_embeddings()
ref_lm_head = self.ref_model.get_output_embeddings()
shift_completion_mask = batch["completion_mask"][:, 1:]
target = batch["completion_input_ids"][:, 1:].clone()
target[shift_completion_mask == 0] = -100
(
loss,
(
chosen_logps_sum,
rejected_logps_sum,
chosen_logits_sum,
rejected_logits_sum,
chosen_rewards_sum,
rejected_rewards_sum,
),
) = self.liger_loss_fn(
_input=outputs.last_hidden_state[:, :-1],
lin_weight=lm_head.weight,
target=target,
bias=lm_head.bias if hasattr(lm_head, "bias") else None,
preference_labels=torch.tensor(batch["label"], dtype=torch.bool).to(self.accelerator.device),
ref_input=ref_outputs.last_hidden_state[:, :-1],
ref_weight=ref_lm_head.weight,
ref_bias=ref_lm_head.bias if hasattr(lm_head, "bias") else None,
kl=kl,
)
if self.aux_loss_enabled:
loss += self.aux_loss_coef * outputs.aux_loss
self._metrics[mode]["kl"].append(kl.item())
all_num_chosen = self.accelerator.gather_for_metrics(num_chosen).sum().item()
all_num_rejected = self.accelerator.gather_for_metrics(num_rejected).sum().item()
if all_num_chosen > 0:
self._metrics[mode]["rewards/chosen"].append(
self.accelerator.gather_for_metrics(chosen_rewards_sum.nansum()).nansum().item() / all_num_chosen
)
self._metrics[mode]["logps/chosen"].append(
self.accelerator.gather_for_metrics(chosen_logps_sum.nansum()).nansum().item() / all_num_chosen
)
self._metrics[mode]["logits/chosen"].append(
self.accelerator.gather_for_metrics(chosen_logits_sum.nansum()).nansum().item() / all_num_chosen
)
if all_num_rejected > 0:
self._metrics[mode]["rewards/rejected"].append(
self.accelerator.gather_for_metrics(rejected_rewards_sum.nansum()).nansum().item() / all_num_rejected
)
self._metrics[mode]["logps/rejected"].append(
self.accelerator.gather_for_metrics(rejected_logps_sum.nansum()).nansum().item() / all_num_rejected
)
self._metrics[mode]["logits/rejected"].append(
self.accelerator.gather_for_metrics(rejected_logits_sum.nansum()).nansum().item() / all_num_rejected
)
if all_num_chosen > 0 and all_num_rejected > 0:
self._metrics[mode]["rewards/margins"].append(
self._metrics[mode]["rewards/chosen"][-1] - self._metrics[mode]["rewards/rejected"][-1]
)
return loss
def _compute_loss(self, model, inputs, return_outputs):
"""Compute the KTO loss and other metrics for the given batch of inputs for train or test."""
mode = "train" if self.model.training else "eval"
batch = {k: (v.to(self.accelerator.device) if isinstance(v, torch.Tensor) else v) for k, v in inputs.items()}
labels = torch.tensor(batch["label"])
num_chosen = labels.sum().to(self.accelerator.device)
num_rejected = (len(labels) - num_chosen).to(self.accelerator.device)
policy_KL_logps = self._compute_kl_logps(model, batch)
_non_model_keys = {
"completion_mask",
"KL_completion_input_ids",
"KL_completion_attention_mask",
"KL_completion_mask",
"KL_completion_token_type_ids",
"KL_completion_mm_token_type_ids",
"label",
"ref_logps",
"ref_KL_logps",
}
model_kwargs = {k: v for k, v in batch.items() if k not in _non_model_keys}
model_kwargs["input_ids"] = model_kwargs.pop("completion_input_ids")
model_kwargs["attention_mask"] = model_kwargs.pop("completion_attention_mask")
if self.aux_loss_enabled:
model_kwargs["output_router_logits"] = True
outputs = model(**model_kwargs)
if self.aux_loss_enabled:
aux_loss = outputs.aux_loss
shift_logits = outputs.logits[:, :-1, :]
per_token_logps = selective_log_softmax(shift_logits, batch["completion_input_ids"][:, 1:])
per_token_logps[batch["completion_mask"][:, 1:] == 0] = 0.0
completion_logps = per_token_logps.sum(-1)
if completion_logps.shape[0] != len(batch["label"]):
raise ValueError(
"There is a mismatch between the number of examples in this batch and the number of "
"examples for which an output sequence was predicted."
)
device = outputs.logits.device
bool_labels = torch.as_tensor(batch["label"], dtype=torch.bool, device=device)
chosen_idx = torch.nonzero(bool_labels, as_tuple=False).view(-1)
rejected_idx = torch.nonzero(~bool_labels, as_tuple=False).view(-1)
policy_chosen_logps = completion_logps.index_select(0, chosen_idx)
policy_rejected_logps = completion_logps.index_select(0, rejected_idx)
policy_chosen_logits = outputs.logits.index_select(0, chosen_idx)
policy_rejected_logits = outputs.logits.index_select(0, rejected_idx)
if self.precompute_ref_logps:
ref_chosen_logps = batch["ref_logps"].index_select(0, chosen_idx)
ref_rejected_logps = batch["ref_logps"].index_select(0, rejected_idx)
if self.calculate_KL:
ref_KL_logps = batch["ref_KL_logps"]
else:
ref_KL_logps = None
else:
ref_model_kwargs = {k: v for k, v in model_kwargs.items() if k != "output_router_logits"}
with torch.no_grad(), disable_gradient_checkpointing(self.model, self.args.gradient_checkpointing_kwargs):
if is_peft_model(self.model) and self.ref_model is None:
ref_model_unwrapped = self.accelerator.unwrap_model(self.model)
with use_adapter(
ref_model_unwrapped, adapter_name="ref" if "ref" in ref_model_unwrapped.peft_config else None
):
ref_KL_logps = self._compute_kl_logps(self.model, batch)
ref_outputs = self.model(**ref_model_kwargs)
else:
ref_KL_logps = self._compute_kl_logps(self.ref_model, batch)
ref_outputs = self.ref_model(**ref_model_kwargs)
ref_shift_logits = ref_outputs.logits[:, :-1, :]
ref_per_token_logps = selective_log_softmax(ref_shift_logits, batch["completion_input_ids"][:, 1:])
ref_per_token_logps[batch["completion_mask"][:, 1:] == 0] = 0.0
ref_completion_logps = ref_per_token_logps.sum(-1)
ref_chosen_logps = ref_completion_logps.index_select(0, chosen_idx)
ref_rejected_logps = ref_completion_logps.index_select(0, rejected_idx)
if self.calculate_KL:
kl = (policy_KL_logps - ref_KL_logps).mean().detach()
kl = self.accelerator.gather_for_metrics(kl).mean().clamp(min=0)
else:
kl = torch.zeros(1).to(policy_chosen_logps.device)
# Chosen losses
if policy_chosen_logps.shape[0] != 0 or ref_chosen_logps.shape[0] != 0:
chosen_logratios = policy_chosen_logps - ref_chosen_logps
if self.loss_type == "kto":
# Eqn (7) of the KTO paper (https://huggingface.co/papers/2402.01306)
chosen_losses = 1 - F.sigmoid(self.beta * (chosen_logratios - kl))
elif self.loss_type == "apo_zero_unpaired":
# Unpaired variant of Eqn (7) of the APO paper (https://huggingface.co/papers/2408.06266)
# Use this loss when you believe the chosen outputs are better than your model's default output
chosen_losses = 1 - F.sigmoid(self.beta * chosen_logratios)
chosen_rewards = self.beta * chosen_logratios.detach()
else:
# lists can't be empty -- if they are, then accelerate.gather will hang
chosen_losses = torch.Tensor([]).to(self.accelerator.device)
chosen_rewards = torch.Tensor([]).to(self.accelerator.device)
# Rejected losses
if policy_rejected_logps.shape[0] != 0 or ref_rejected_logps.shape[0] != 0:
rejected_logratios = policy_rejected_logps - ref_rejected_logps
if self.loss_type == "kto":
rejected_losses = 1 - F.sigmoid(self.beta * (kl - rejected_logratios))
elif self.loss_type == "apo_zero_unpaired":
rejected_losses = F.sigmoid(self.beta * rejected_logratios)
rejected_rewards = self.beta * rejected_logratios.detach()
else:
# lists can't be empty -- if they are, then accelerate.gather will hang
rejected_losses = torch.Tensor([]).to(self.accelerator.device)
rejected_rewards = torch.Tensor([]).to(self.accelerator.device)
losses = torch.cat(
(self.desirable_weight * chosen_losses, self.undesirable_weight * rejected_losses),
0,
)
self._metrics[mode]["kl"].append(kl.item())
all_num_chosen = self.accelerator.gather_for_metrics(num_chosen).sum().item()
all_num_rejected = self.accelerator.gather_for_metrics(num_rejected).sum().item()
if all_num_chosen > 0:
self._metrics[mode]["rewards/chosen"].append(
self.accelerator.gather_for_metrics(chosen_rewards.nansum()).nansum().item() / all_num_chosen
)
self._metrics[mode]["logps/chosen"].append(
self.accelerator.gather_for_metrics(policy_chosen_logps.nansum()).nansum().item() / all_num_chosen
)
self._metrics[mode]["logits/chosen"].append(
self.accelerator.gather_for_metrics(policy_chosen_logits.nansum()).nansum().item() / all_num_chosen
)
if all_num_rejected > 0:
self._metrics[mode]["rewards/rejected"].append(
self.accelerator.gather_for_metrics(rejected_rewards.nansum()).nansum().item() / all_num_rejected
)
self._metrics[mode]["logps/rejected"].append(
self.accelerator.gather_for_metrics(policy_rejected_logps.nansum()).nansum().item() / all_num_rejected
)
self._metrics[mode]["logits/rejected"].append(
self.accelerator.gather_for_metrics(policy_rejected_logits.nansum()).nansum().item() / all_num_rejected
)
if all_num_chosen > 0 and all_num_rejected > 0:
self._metrics[mode]["rewards/margins"].append(
self._metrics[mode]["rewards/chosen"][-1] - self._metrics[mode]["rewards/rejected"][-1]
)
loss = losses.nanmean()
if self.aux_loss_enabled:
loss += self.aux_loss_coef * aux_loss
return (loss, outputs) if return_outputs else loss
def evaluate(
self,
eval_dataset: Dataset | dict[str, Dataset] | None = None,
ignore_keys: list[str] | None = None,
metric_key_prefix: str = "eval",
) -> dict[str, float]:
# When a dataset is passed directly to `evaluate` (e.g. a held-out test set), preprocess it the same way
# `__init__` does, so that `evaluate` accepts the same dataset types as the trainer. `_prepare_dataset` is
# idempotent: it skips datasets that are already tokenized. A `str` selects a dataset that was already prepared
# at init time, so it's left untouched.
if not self._is_vision_dataset and eval_dataset is not None and not isinstance(eval_dataset, str):
if isinstance(eval_dataset, dict):
eval_dataset = {
key: self._prepare_dataset(dataset, self.processing_class, self.args, key)
for key, dataset in eval_dataset.items()
}
else:
eval_dataset = self._prepare_dataset(eval_dataset, self.processing_class, self.args, "eval")
# With `precompute_ref_log_probs`, `_compute_loss` reads the reference log-probs from the batch, so they
# must be precomputed here as well, mirroring `__init__`.
if self.precompute_ref_logps:
batch_size = self.args.precompute_ref_batch_size or self.args.per_device_eval_batch_size
if isinstance(eval_dataset, dict):
eval_dataset = {
name: self._precompute_ref_logps(dataset, name, batch_size)
for name, dataset in eval_dataset.items()
}
else:
eval_dataset = self._precompute_ref_logps(eval_dataset, "eval", batch_size)
return super().evaluate(
eval_dataset=eval_dataset, ignore_keys=ignore_keys, metric_key_prefix=metric_key_prefix
)
def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None):
try:
if self.use_liger_kernel:
return self._compute_loss_liger(model, inputs, return_outputs)
return self._compute_loss(model, inputs, return_outputs)
except ValueError as e:
if "Image features and image tokens do not match" in str(e) and self.args.max_length is not None:
raise ValueError(
f"The current `max_length` ({self.args.max_length}) is too short and causes image placeholder "
f"tokens in `input_ids` to be truncated, while the corresponding image features remain intact. "
f"Please increase `max_length` or set it to `None` to disable truncation."
) from e
raise
# Override training step to add activation offloading context.
def training_step(self, *args, **kwargs):
with self.maybe_activation_offload_context:
return super().training_step(*args, **kwargs)
def log(self, logs: dict[str, float], start_time: float | None = None) -> None:
mode = "train" if self.model.training else "eval"
metrics = {key: sum(val) / len(val) for key, val in self._metrics[mode].items()} # average the metrics
# This method can be called both in training and evaluation. When called in evaluation, the keys in `logs`
# start with "eval_". We need to add the prefix "eval_" to the keys in `metrics` to match the format.
if mode == "eval":
metrics = {f"eval_{key}": val for key, val in metrics.items()}
logs.update(metrics)
super().log(logs, start_time)
self._metrics[mode].clear()
# During eval, Trainer calls prediction_step. If no labels are present in the inputs, it only runs forward and
# returns logits. We override prediction_step to force compute_loss, because this trainer doesn't involve labels.
def prediction_step(self, model, inputs, prediction_loss_only, ignore_keys: list[str] | None = None):
inputs = self._prepare_inputs(inputs)
with torch.no_grad(), self.compute_loss_context_manager():
if prediction_loss_only:
loss = self.compute_loss(model, inputs, return_outputs=False) # logits aren't materialized with liger
logits, labels = None, None
else:
loss, outputs = self.compute_loss(model, inputs, return_outputs=True)
logits, labels = outputs.logits, inputs["completion_input_ids"]
return loss, logits, labels
# Ensure the model card is saved along with the checkpoint
def _save_checkpoint(self, model, trial):
if self.args.hub_model_id is None:
model_name = Path(self.args.output_dir).name
else:
model_name = self.args.hub_model_id.split("/")[-1]
self.create_model_card(model_name=model_name)
super()._save_checkpoint(model, trial)
class UnslothKTOTrainer(_UnslothKTOTrainer):
"""
KTOTrainer(*args, **kwargs)
"""
def __init__(
self,
model,
ref_model = None,
args = None,
train_dataset = None,
eval_dataset = None,
processing_class = None,
data_collator = None,
callbacks = None,
peft_config = None,
compute_metrics = None,
**kwargs
):
if args is None: args = UnslothKTOConfig()
use_bf16 = getattr(args, 'bf16', False)
if type(use_bf16) is not bool: use_bf16 = False
use_fp16 = getattr(args, 'fp16', False)
if type(use_fp16) is not bool: use_fp16 = False
force_float32 = False
full_finetuning = os.environ.get('UNSLOTH_ENABLE_FULL_FINETUNING', '0') == '1'
if not full_finetuning and (os.environ.get('UNSLOTH_FORCE_FLOAT32', '0') == '1'):
print('Unsloth: Switching to float32 training since model cannot work with float16')
force_float32 = True
mixed_precision_dtype = os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32')
dtype = getattr(model.config, 'dtype', None) or getattr(model.config, 'torch_dtype', None)
if dtype is None: dtype = model.get_input_embeddings().weight.dtype
from unsloth_zoo.utils import _get_dtype
dtype = _get_dtype(dtype)
float16 = dtype == torch.float16
if not force_float32 and (float16 and use_bf16): raise TypeError('Unsloth: Model is in float16 precision but you want to use bfloat16 precision. Set fp16 to `True` and bf16 to `False`')
if not force_float32 and (not float16 and use_fp16): raise TypeError('Unsloth: Model is in bfloat16 precision but you want to use float16 precision. Set fp16 to `False` and bf16 to `True`')
if force_float32:
# Forced float32 training
args.fp16 = False
args.bf16 = False
os.environ['ACCELERATE_MIXED_PRECISION'] = 'no'
if hasattr(args, 'mixed_precision'): args.mixed_precision = 'no'
# args.mixed_precision is a new argument which needs to be set now
elif (not use_bf16 and not use_fp16) and mixed_precision_dtype == 'float32':
# Mixed precision training
args.fp16 = float16
args.bf16 = not float16
os.environ['ACCELERATE_MIXED_PRECISION'] = 'fp16' if float16 else 'bf16'
if hasattr(args, 'mixed_precision'): args.mixed_precision = 'fp16' if float16 else 'bf16'
# args.mixed_precision is a new argument which needs to be set now
elif mixed_precision_dtype == 'bfloat16':
# Both False since bfloat16 full finetuning doesn't do any autocasting.
args.fp16 = False
args.bf16 = False
os.environ['ACCELERATE_MIXED_PRECISION'] = 'no'
if hasattr(args, 'mixed_precision'): args.mixed_precision = 'no'
# args.mixed_precision is a new argument which needs to be set now
if getattr(args, 'eval_dataset', None) is not None and getattr(args, 'eval_strategy', 'no') == 'no':
args.eval_strategy = 'steps'
if getattr(args, 'eval_steps', None) is None: args.eval_steps = 0.1
ga_steps = getattr(args, 'gradient_accumulation_steps', None)
if ga_steps is not None and ga_steps > 1:
from transformers import __version__ as transformers_version
if Version(transformers_version) <= Version('4.45.2'):
print('**** Unsloth: Please use our fixed gradient_accumulation_steps by updating transformers, TRL and Unsloth!\n'
'`pip install --upgrade --no-cache-dir --force-reinstall --no-deps unsloth transformers trl unsloth_zoo`')
if getattr(args, 'eval_strategy', 'no') != 'no':
eval_bsz = getattr(args, 'per_device_eval_batch_size', 8)
if eval_bsz == 8 and args.per_device_train_batch_size < eval_bsz: args.per_device_eval_batch_size = args.per_device_train_batch_size
if getattr(args, 'eval_accumulation_steps', None) is None and ga_steps is not None: args.eval_accumulation_steps = ga_steps
fp16_full_eval = getattr(args, 'fp16_full_eval', False)
if type(fp16_full_eval) is not bool: fp16_full_eval = False
bf16_full_eval = getattr(args, 'bf16_full_eval', False)
if type(bf16_full_eval) is not bool: bf16_full_eval = False
if args.fp16 and bf16_full_eval: args.bf16_full_eval = False; args.fp16_full_eval = True
if args.bf16 and fp16_full_eval: args.bf16_full_eval = True; args.fp16_full_eval = False
if force_float32:
args.bf16_full_eval = False
args.fp16_full_eval = False
elif os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32') == 'bfloat16':
args.bf16_full_eval = True
args.fp16_full_eval = False
elif not bf16_full_eval and not fp16_full_eval:
args.bf16_full_eval = args.bf16
args.fp16_full_eval = args.fp16
_output_logits = False
if locals().get('compute_metrics', None) is not None: _output_logits = True
if locals().get('preprocess_logits_for_metrics', None) is not None: _output_logits = True
if _output_logits:
os.environ['UNSLOTH_RETURN_LOGITS'] = '1'
if model is not None:
_warnings_issued = getattr(model, 'warnings_issued', None)
if _warnings_issued is None:
model.warnings_issued = {}
elif not isinstance(_warnings_issued, dict):
try:
model.warnings_issued = dict(_warnings_issued)
except Exception:
model.warnings_issued = {}
if 'max_seq_length' not in locals() and not hasattr(args, 'max_seq_length'):
pass
else:
model_max_seq_length = getattr(model, 'max_seq_length', None)
args_max_seq_length = getattr(args, 'max_seq_length', None)
if args_max_seq_length is None and model_max_seq_length is not None:
max_seq_length = model.max_seq_length
if hasattr(args, 'max_seq_length'): args.max_seq_length = max_seq_length
elif args_max_seq_length is not None and model_max_seq_length is not None:
if args_max_seq_length > model_max_seq_length:
print('Unsloth: You set `max_seq_length` as ' + str(args_max_seq_length) + ' but '
'the maximum the model supports is ' + str(model_max_seq_length) + '. We shall reduce it.')
args.max_seq_length = model_max_seq_length
if model is not None and hasattr(model, 'for_training'):
model.for_training(use_gradient_checkpointing=getattr(args, 'gradient_checkpointing', True))
if 'tokenizer' in locals() and hasattr(tokenizer, 'padding_side'): tokenizer.padding_side = 'right'
if 'processing_class' in locals():
if hasattr(processing_class, 'padding_side'): processing_class.padding_side = 'right'
if hasattr(processing_class, 'tokenizer') and hasattr(processing_class.tokenizer, 'padding_side'): processing_class.tokenizer.padding_side = 'right'
__tokenizer = processing_class if 'processing_class' in locals() else tokenizer
from unsloth_zoo.vision_utils import UnslothVisionDataCollator
if not isinstance(data_collator, UnslothVisionDataCollator):
if isinstance(data_collator, DataCollatorForSeq2Seq) and 'labels' not in train_dataset.column_names:
data_collator = TransformersDataCollatorForLanguageModeling(
__tokenizer,
mlm = False,
mlm_probability = 0.0,
pad_to_multiple_of = getattr(args, 'pad_to_multiple_of', None),
)
elif isinstance(data_collator, TransformersDataCollatorForLanguageModeling) and 'labels' in train_dataset.column_names:
data_collator = DataCollatorForSeq2Seq(
__tokenizer,
pad_to_multiple_of = getattr(args, 'pad_to_multiple_of', None),
)
else:
if hasattr(args, 'remove_unused_columns'): args.remove_unused_columns = False
if hasattr(args, 'dataset_text_field'): args.dataset_text_field = ''
if hasattr(args, 'dataset_kwargs'): args.dataset_kwargs = {'skip_prepare_dataset': True}
if not isinstance(data_collator, UnslothVisionDataCollator):
if not hasattr(__tokenizer, 'pad') and hasattr(__tokenizer, 'tokenizer'):
if isinstance(data_collator, DataCollatorForSeq2Seq):
data_collator = DataCollatorForSeq2Seq(
__tokenizer.tokenizer,
pad_to_multiple_of = getattr(args, 'pad_to_multiple_of', None),
)
elif isinstance(data_collator, TransformersDataCollatorForLanguageModeling):
data_collator = TransformersDataCollatorForLanguageModeling(
__tokenizer.tokenizer,
mlm = False,
mlm_probability = 0.0,
pad_to_multiple_of = getattr(args, 'pad_to_multiple_of', None),
)
other_metrics = []
from unsloth_zoo.logging_utils import PatchRLStatistics
PatchRLStatistics('kto_trainer', other_metrics)
# [TODO] Fix up DataParallel multiplying batch sizes
# [TODO] DDP works, but DP seems to not work? [TODO]
if getattr(args, "parallel_mode", None) == ParallelMode.NOT_DISTRIBUTED and args.n_gpu > 1:
if getattr(args, "_n_gpu", 1) != 1:
args._n_gpu = 1
if "model" in locals() and hasattr(model, "for_training"):
model.for_training(use_gradient_checkpointing=getattr(args, 'gradient_checkpointing', True))
super().__init__(
model = model,
ref_model = ref_model,
args = args,
train_dataset = train_dataset,
eval_dataset = eval_dataset,
processing_class = processing_class,
data_collator = data_collator,
callbacks = callbacks,
peft_config = peft_config,
compute_metrics = compute_metrics,**kwargs)
if "model" in locals() and hasattr(model, "for_inference"):
model.for_inference()
if hasattr(self, 'neftune_hook_handle'):
self.neftune_hook_handle.remove()
if hasattr(self, 'neftune_hook_handle'): del self.neftune_hook_handle
if getattr(args, 'neftune_noise_alpha', None) is not None:
model.get_input_embeddings().neftune_noise_alpha = self.neftune_noise_alpha
pass
if hasattr(self, 'accelerator'):
scaler = self.accelerator.scaler
current_model = model
while hasattr(current_model, 'model'):
current_model.accelerator_scaler = scaler
current_model = current_model.model
current_model.accelerator_scaler = scaler
pass
if hasattr(self, 'train'):
self.train = MethodType(prepare_for_training_mode(self.__class__.train), self)
pass
if hasattr(self, 'llm') and self.llm is not None and hasattr(self.llm, 'get_tokenizer'):
_vllm_tok = self.llm.get_tokenizer()
_pc = getattr(self, 'processing_class', None) or getattr(self, 'tokenizer', None)
if _vllm_tok is not None and _pc is not None and getattr(_pc, 'chat_template', None) is not None and getattr(_vllm_tok, 'chat_template', None) is None:
_vllm_tok.chat_template = _pc.chat_template
pass
pass
if hasattr(logger, "addFilter"):
import logging
class HideLoggingMessage(logging.Filter):
def __init__(self, text): self.text = text
def filter(self, x): return not (self.text in x.getMessage())
pass
logger.addFilter(HideLoggingMessage("`use_cache=True`"))