| """ |
| 2026.6.7 |
| 2026.6.9 |
| 5.5.0 |
| 1.7.0 |
| __UNSLOTH_VERSIONING__ |
| """ |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| 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 |
|
|
| |
| 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 |
| |
| |
| 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): |
| |
| try: |
| _unsloth_reset_stray_compile_cache(self) |
| except Exception: |
| pass |
| |
| |
| |
| |
| |
| if getattr(self, '_unsloth_training_completed', False): |
| try: |
| import wandb |
| if wandb.run is not None: |
| wandb.finish() |
| |
| for cb in self.callback_handler.callbacks: |
| if type(cb).__name__ == 'WandbCallback': |
| cb._initialized = False |
| break |
| except: |
| pass |
| |
| _was_training = None |
| |
| 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) |
| |
| 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) |
| |
| try: |
| reset_unsloth_gradient_checkpointing_buffers() |
| except: |
| pass |
| |
| |
| 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: |
| |
| 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 = [] |
| |
| 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) |
| |
| 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 |
|
|
| |
| row_indices = torch.arange(batch_size, device=device).unsqueeze(1).expand_as(dest_indices) |
|
|
| |
| 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(): |
| |
| 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: |
| |
| 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: |
| |
| 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] = <factory>, 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 |
| 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 |
| |
| 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 `<ModelArchitecture>.from_pretrained` (where `<ModelArchitecture>` 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", |
| |
| "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, |
| ): |
| |
| 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): |
| |
| |
| 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 |
|
|
| |
| if isinstance(model, str): |
| model_init_kwargs = args.model_init_kwargs or {} |
| |
| 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." |
| ) |
| |
| _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." |
| ) |
|
|
| |
| 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 |
|
|
| |
| 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." |
| ) |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| 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: |
| |
| |
| |
| |
| |
| 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) |
|
|
| |
| |
| if is_peft_model(model) and args.gradient_checkpointing: |
| model.enable_input_require_grads() |
|
|
| |
| |
| |
| |
| if _is_quantized_model: |
| for param in model.parameters(): |
| if param.requires_grad: |
| param.data = param.data.to(torch.bfloat16) |
|
|
| |
| 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)`." |
| ) |
|
|
| |
| 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, |
| ) |
|
|
| |
| 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.", |
| ) |
|
|
| |
| |
| 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") |
|
|
| |
| |
| |
| |
| 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, |
| ) |
|
|
| |
| 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() |
|
|
| |
| if ref_model is None: |
| if is_peft_model(self.model) or args.precompute_ref_log_probs: |
| |
| |
| |
| self.ref_model = None |
| else: |
| ref_model_init_kwargs = args.model_init_kwargs or {} |
| |
| 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 |
|
|
| |
| if args.disable_dropout: |
| disable_dropout_in_model(model) |
| if self.ref_model is not None: |
| disable_dropout_in_model(self.ref_model) |
|
|
| |
| self._metrics = {"train": defaultdict(list), "eval": defaultdict(list)} |
|
|
| |
| |
| |
| self.model_accepts_loss_kwargs = False |
|
|
| |
| 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 |
| |
| 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): |
| if self._is_vlm: |
| input = prepare_multimodal_messages(input) |
| result = processing_class.apply_chat_template(input, tokenize=True, return_dict=True, **kwargs) |
| else: |
| result = processing_class(text=input) |
| |
| 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): |
| 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): |
| 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: |
| |
| map_kwargs = {} |
| if isinstance(dataset, Dataset): |
| map_kwargs["num_proc"] = args.dataset_num_proc |
|
|
| |
| |
| with PartialState().main_process_first(): |
| |
| first_example = next(iter(dataset)) |
| if "prompt" not in first_example: |
| if isinstance(dataset, Dataset): |
| map_kwargs["desc"] = f"Extracting prompt from {dataset_name} dataset" |
| dataset = dataset.map(extract_prompt, **map_kwargs) |
|
|
| |
| first_example = next(iter(dataset)) |
| if "chosen" in first_example and "rejected" in first_example: |
| if isinstance(dataset, Dataset): |
| map_kwargs["desc"] = f"Unpairing {dataset_name} dataset" |
| dataset = unpair_preference_dataset(dataset, **map_kwargs) |
|
|
| |
| first_example = next(iter(dataset)) |
| if not is_conversational(first_example): |
| if isinstance(dataset, Dataset): |
| 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) |
|
|
| |
| if isinstance(dataset, Dataset): |
| 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) |
|
|
| |
| if self.calculate_KL: |
| |
| |
| kl_dataset = self._get_kl_dataset(dataset, dataset_name, args) |
| dataset = concatenate_datasets([dataset, kl_dataset], axis=1) |
|
|
| |
| if dataset_name == "train" and isinstance(dataset, Dataset): |
| num_desirable = max(sum(dataset["label"]), 1) |
| num_undesirable = max(len(dataset["label"]) - num_desirable, 1) |
|
|
| if num_desirable != num_undesirable: |
| |
| 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._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 |
| ) |
|
|
| 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") |
| |
| |
| 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 |
|
|
| |
| |
| |
| |
| |
| 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) |
|
|
| |
| 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) |
| |
| 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": |
| |
| chosen_losses = 1 - F.sigmoid(self.beta * (chosen_logratios - kl)) |
| elif self.loss_type == "apo_zero_unpaired": |
| |
| |
| chosen_losses = 1 - F.sigmoid(self.beta * chosen_logratios) |
|
|
| chosen_rewards = self.beta * chosen_logratios.detach() |
|
|
| else: |
| |
| chosen_losses = torch.Tensor([]).to(self.accelerator.device) |
| chosen_rewards = torch.Tensor([]).to(self.accelerator.device) |
| |
| 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: |
| |
| 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]: |
| |
| |
| |
| |
| 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") |
| |
| |
| 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 |
|
|
| |
| 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()} |
| |
| |
| 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() |
|
|
| |
| |
| 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, 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 |
|
|
| |
| 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: |
| |
| args.fp16 = False |
| args.bf16 = False |
| os.environ['ACCELERATE_MIXED_PRECISION'] = 'no' |
| if hasattr(args, 'mixed_precision'): args.mixed_precision = 'no' |
| |
| elif (not use_bf16 and not use_fp16) and mixed_precision_dtype == 'float32': |
| |
| 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' |
| |
| elif mixed_precision_dtype == 'bfloat16': |
| |
| args.fp16 = False |
| args.bf16 = False |
| os.environ['ACCELERATE_MIXED_PRECISION'] = 'no' |
| if hasattr(args, 'mixed_precision'): args.mixed_precision = 'no' |
| |
| |
| 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) |
| |
| |
| |
| 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`")) |
|
|
|
|