Zwounds/Boolean_Search_Query_Model
0
1"""22025.3.1532025.3.1744.50.0.dev050.15.26__UNSLOTH_VERSIONING__7"""8from torch import Tensor9import torch10import torch.nn as nn11from torch.nn import functional as F12from trl.trainer.dpo_trainer import (Any, AutoModelForCausalLM, BaseImageProcessor, Callable, DPOConfig, DPOTrainer, DataCollator, DataCollatorForPreference, DataLoader, Dataset, EvalLoopOutput, F, FDivergenceConstants, FDivergenceType, FeatureExtractionMixin, IterableDataset, Literal, MODEL_FOR_VISION_2_SEQ_MAPPING_NAMES, Optional, PartialState, PeftModel, PreTrainedModel, PreTrainedModelWrapper, PreTrainedTokenizerBase, ProcessorMixin, RunningMoments, SyncRefModelCallback, Trainer, TrainerCallback, Union, amp, cap_exp, contextmanager, create_reference_model, dataclass, deepcopy, defaultdict, deprecate_kwarg, disable_dropout_in_model, empty_cache, flush_left, generate_model_card, get_comet_experiment_url, inspect, is_comet_available, is_peft_available, is_torch_xpu_available, is_wandb_available, log_table_to_comet_experiment, maybe_apply_chat_template, maybe_extract_prompt, nn, nullcontext, os, pad, pad_to_length, pd, peft_module_casting_to_bf16, prepare_model_for_kbit_training, random, textwrap, torch, tqdm, transformers, version, warnings)13 14 15import os16from typing import *17from dataclasses import dataclass, field18from packaging.version import Version19import torch20import numpy as np21from contextlib import nullcontext22from torch.nn import functional as F23from transformers import DataCollatorForSeq2Seq, DataCollatorForLanguageModeling24 25torch_compile_options = {26 "epilogue_fusion" : True,27 "max_autotune" : False,28 "shape_padding" : True,29 "trace.enabled" : False,30 "triton.cudagraphs" : False,31}32 33@torch.compile(dynamic = True, fullgraph = True, options = torch_compile_options,)34def selective_log_softmax(logits, index):35 logits = logits.to(torch.float32)36 selected_logits = torch.gather(logits, dim = -1, index = index.unsqueeze(-1)).squeeze(-1)37 # loop to reduce peak mem consumption38 # logsumexp_values = torch.stack([torch.logsumexp(lg, dim=-1) for lg in logits])39 logsumexp_values = torch.logsumexp(logits, dim = -1)40 per_token_logps = selected_logits - logsumexp_values # log_softmax(x_i) = x_i - logsumexp(x)41 return per_token_logps42@dataclass43class UnslothDPOConfig(DPOConfig):44 """45 46 Configuration class for the [`DPOTrainer`].47 48 Using [`~transformers.HfArgumentParser`] we can turn this class into49 [argparse](https://docs.python.org/3/library/argparse#module-argparse) arguments that can be specified on the50 command line.51 52 Parameters:53 > Parameters that control the model and reference model54 55 model_init_kwargs (`dict[str, Any]` or `None`, *optional*, defaults to `None`):56 Keyword arguments for `AutoModelForCausalLM.from_pretrained`, used when the `model` argument of the57 [`DPOTrainer`] is provided as a string.58 ref_model_init_kwargs (`dict[str, Any]` or `None`, *optional*, defaults to `None`):59 Keyword arguments for `AutoModelForCausalLM.from_pretrained`, used when the `ref_model` argument of the60 [`DPOTrainer`] is provided as a string.61 model_adapter_name (`str` or `None`, *optional*, defaults to `None`):62 Name of the train target PEFT adapter, when using LoRA with multiple adapters.63 ref_adapter_name (`str` or `None`, *optional*, defaults to `None`):64 Name of the reference PEFT adapter, when using LoRA with multiple adapters.65 force_use_ref_model (`bool`, *optional*, defaults to `False`):66 If you provide a PEFT model as the active model and wish to use a different model for the `ref_model`, set67 this flag to `True`.68 disable_dropout (`bool`, *optional*, defaults to `True`):69 Whether to disable dropout in the model and reference model.70 use_logits_to_keep (`bool`, *optional*, defaults to `False`):71 If `True`, only a specified number of logits are computed in the forward pass. This can be useful for72 saving memory and speeding up training by not computing the logits for all tokens, especially in73 scenarios when working with very long prompts where labels are ignored (-100).74 75 > Parameters that control the data preprocessing76 77 dataset_num_proc (`int` or `None`, *optional*, defaults to `None`):78 Number of processes to use for processing the dataset.79 padding_value (`int` or `None`, *optional*, defaults to `None`):80 Padding value to use. If `None`, the padding value of the tokenizer is used.81 label_pad_token_id (`int`, *optional*, defaults to `-100`):82 Padding value to use for labels.83 max_prompt_length (`int` or `None`, *optional*, defaults to `512`):84 Maximum length of the prompt.85 max_completion_length (`int` or `None`, *optional*, defaults to `None`):86 Maximum length of the completion.87 max_length (`int` or `None`, *optional*, defaults to `1024`):88 Maximum length of the full sequence (prompt + completion).89 truncation_mode (`str`, *optional*, defaults to `"keep_end"`):90 Truncation mode to use when the sequence exceeds `max_length`. Possible values are `"keep_end"` and91 `"keep_start"`.92 padding_free (`bool`, *optional*, defaults to `False`):93 Whether forward passes are performed without padding by flattening all sequences in the batch94 into a single continuous sequence. This approach requires associating a `position_ids` vector to track95 positional information. Currently, this is only supported with the `flash_attention_2` mechanism, as it96 can handle the flattened batch structure.97 precompute_ref_log_probs (`bool`, *optional*, defaults to `False`):98 Whether to precompute the log probabilities from the reference model. Setting this to `True` allows99 training without needing the reference model during training, which can help reduce GPU memory usage. If100 set to `False` (default), the reference model will be used during training to compute log probabilities101 on-the-fly.102 precompute_ref_batch_size (`int` or `None`, *optional*, defaults to `None`):103 Batch size to use when precomputing reference model log probabilities. This can be set higher than the104 training batch size to speed up preprocessing. If `None`, defaults to `per_device_train_batch_size` for105 training and `per_device_eval_batch_size` for evaluation.106 tools (`Optional[list[Union[dict, Callable]]]`, *optional*, defaults to `None`):107 List of tools (callable functions) that will be accessible to the model.108 If the template does not support function calling, this argument will have no effect.109 110 > Parameters that control the training111 112 learning_rate (`float`, *optional*, defaults to `1e-6`):113 Initial learning rate for [`AdamW`] optimizer. The default value replaces that of114 [`~transformers.TrainingArguments`].115 loss_type (`str`, *optional*, defaults to `"sigmoid"`):116 Type of loss to use. Possible values are:117 118 - `"sigmoid"`: sigmoid loss from the original [DPO](https://huggingface.co/papers/2305.18290) paper.119 - `"hinge"`: hinge loss on the normalized likelihood from the [SLiC](https://huggingface.co/papers/2305.10425) paper.120 - `"ipo"`: IPO loss from the [IPO](https://huggingface.co/papers/2310.12036) paper.121 - `"exo_pair"`: pairwise EXO loss from the [EXO](https://huggingface.co/papers/2402.00856) paper.122 - `"nca_pair"`: pairwise NCA loss from the [NCA](https://huggingface.co/papers/2402.05369) paper.123 - `"robust"`: unbiased estimate of the DPO loss that is robust to preference noise from the [Robust DPO](https://huggingface.co/papers/2403.00409) paper.124 - `"bco_pair"`: pairwise BCO loss from the [BCO](https://huggingface.co/papers/2404.04656) paper.125 - `"sppo_hard"`: SPPO loss with hard label from the [SPPO](https://huggingface.co/papers/2405.00675) paper.126 - `"aot"`: AOT loss for paired datasets from the [AOT](https://huggingface.co/papers/2406.05882) paper.127 - `"aot_pair"`: AOT loss for unpaired datasets from the [AOT](https://huggingface.co/papers/2406.05882) paper.128 - `"discopop"`: DiscoPOP (a.k.a Log-Ratio Modulated Loss, LRML) loss from the [DiscoPOP](https://huggingface.co/papers/2406.08414) paper.129 - `"apo_zero"`: APO-zero loss from the [APO](https://huggingface.co/papers/2408.06266) paper.130 - `"apo_down"`: APO-down loss from the [APO](https://huggingface.co/papers/2408.06266) paper.131 132 beta (`float`, *optional*, defaults to `0.1`):133 Parameter controlling the deviation from the reference model. Higher β means less deviation from the134 reference model. For the IPO loss (`loss_type="ipo"`), β is the regularization parameter denoted by τ in135 the [paper](https://huggingface.co/papers/2310.12036).136 f_divergence_type (`str`, *optional*, defaults to `FDivergenceType.REVERSE_KL`):137 Type of f-divergence regularization function to compute divergence between policy and reference model.138 f_alpha_divergence_coef (`float`, *optional*, defaults to `1.0`):139 α coefficient in the α-divergence u^-α regularization function for DPO loss.140 reference_free (`bool`, *optional*, defaults to `False`):141 Whether to ignore the provided reference model and implicitly use a reference model that assigns equal142 probability to all responses.143 label_smoothing (`float`, *optional*, defaults to `0.0`):144 Robust DPO label smoothing parameter from the [cDPO](https://ericmitchell.ai/cdpo.pdf) report and145 [Robust DPO](https://huggingface.co/papers/2403.00409) paper that should be between `0.0` and `0.5`.146 use_weighting (`bool`, *optional*, defaults to `False`):147 Whether to weight the loss as done in the [WPO](https://huggingface.co/papers/2406.11827) paper.148 rpo_alpha (`float`, *optional*, defaults to `None`):149 α parameter from the [RPO](https://huggingface.co/papers/2404.19733) paper (v3), which controls the150 weighting of the NLL term in the loss. If `None`, no weighting is applied and the loss is the same as the151 DPO loss. The paper recommends `rpo_alpha=1.0`.152 discopop_tau (`float`, *optional*, defaults to `0.05`):153 τ/temperature parameter from the [DiscoPOP](https://huggingface.co/papers/2406.08414) paper, which controls154 the shape of log ratio modulated loss. The paper recommends the default value `discopop_tau=0.05`.155 sync_ref_model (`bool`, *optional*, defaults to `False`):156 Whether to synchronize the reference model with the active model every `ref_model_sync_steps` steps, using157 the `ref_model_mixup_alpha` parameter. This synchronization originites from the158 [TR-DPO](https://huggingface.co/papers/2404.09656) paper.159 ref_model_mixup_alpha (`float`, *optional*, defaults to `0.9`):160 α parameter from the [TR-DPO](https://huggingface.co/papers/2404.09656) paper, which controls the mix161 between the current policy and the previous reference policy during updates. The reference policy is162 updated according to the equation: `π_ref = α * π_θ + (1 - α) * π_ref_prev`. To use this parameter, you163 must set `sync_ref_model=True`.164 ref_model_sync_steps (`int`, *optional*, defaults to `64`):165 τ parameter from the [TR-DPO](https://huggingface.co/papers/2404.09656) paper, which determines how166 frequently the current policy is synchronized with the reference policy. To use this parameter, you must167 set `sync_ref_model=True`.168 169 > Parameters that control the logging170 171 generate_during_eval (`bool`, *optional*, defaults to `False`):172 Whether to generate and log completions from both the model and the reference model to W&B or Comet during173 evaluation.174 175 """176 vllm_sampling_params: Optional[Any] = field(177 default = None,178 metadata = {'help': 'vLLM SamplingParams'},179 )180 unsloth_num_chunks : Optional[int] = field(181 default = -1,182 metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},183 )184 def __init__(185 self,186 output_dir = None,187 overwrite_output_dir = None,188 do_train = False,189 do_eval = False,190 do_predict = False,191 eval_strategy = 'no',192 prediction_loss_only = False,193 per_device_train_batch_size = 4,194 per_device_eval_batch_size = 4,195 per_gpu_train_batch_size = None,196 per_gpu_eval_batch_size = None,197 gradient_accumulation_steps = 2,198 eval_accumulation_steps = 2,199 eval_delay = 0,200 torch_empty_cache_steps = 250,201 learning_rate = 5e-05,202 weight_decay = 0.01,203 adam_beta1 = 0.9,204 adam_beta2 = 0.999,205 adam_epsilon = 1e-08,206 max_grad_norm = 1.0,207 num_train_epochs = 3.0,208 max_steps = -1,209 lr_scheduler_type = 'linear',210 warmup_ratio = 0.1,211 warmup_steps = 0,212 log_level = 'passive',213 log_level_replica = 'warning',214 log_on_each_node = True,215 logging_dir = None,216 logging_strategy = 'steps',217 logging_first_step = False,218 logging_steps = 1,219 logging_nan_inf_filter = False,220 save_strategy = 'steps',221 save_steps = 500,222 save_total_limit = None,223 save_safetensors = True,224 save_on_each_node = False,225 save_only_model = False,226 restore_callback_states_from_checkpoint = False,227 no_cuda = False,228 use_cpu = False,229 use_mps_device = False,230 seed = 3407,231 data_seed = 3407,232 jit_mode_eval = False,233 use_ipex = False,234 bf16 = False,235 fp16 = False,236 fp16_opt_level = 'O1',237 half_precision_backend = 'auto',238 bf16_full_eval = False,239 fp16_full_eval = False,240 tf32 = None,241 local_rank = -1,242 ddp_backend = None,243 tpu_num_cores = None,244 tpu_metrics_debug = False,245 debug = '',246 dataloader_drop_last = False,247 eval_steps = None,248 dataloader_num_workers = 0,249 dataloader_prefetch_factor = None,250 past_index = -1,251 run_name = None,252 disable_tqdm = None,253 remove_unused_columns = True,254 label_names = None,255 load_best_model_at_end = False,256 metric_for_best_model = None,257 greater_is_better = None,258 ignore_data_skip = False,259 fsdp = '',260 fsdp_min_num_params = 0,261 fsdp_config = None,262 tp_size = 0,263 fsdp_transformer_layer_cls_to_wrap = None,264 accelerator_config = None,265 deepspeed = None,266 label_smoothing_factor = 0.0,267 optim = 'adamw_8bit',268 optim_args = None,269 adafactor = False,270 group_by_length = False,271 length_column_name = 'length',272 report_to = None,273 ddp_find_unused_parameters = None,274 ddp_bucket_cap_mb = None,275 ddp_broadcast_buffers = None,276 dataloader_pin_memory = True,277 dataloader_persistent_workers = False,278 skip_memory_metrics = True,279 use_legacy_prediction_loop = False,280 push_to_hub = False,281 resume_from_checkpoint = None,282 hub_model_id = None,283 hub_strategy = 'every_save',284 hub_token = None,285 hub_private_repo = None,286 hub_always_push = False,287 gradient_checkpointing = False,288 gradient_checkpointing_kwargs = None,289 include_inputs_for_metrics = False,290 eval_do_concat_batches = True,291 fp16_backend = 'auto',292 evaluation_strategy = None,293 push_to_hub_model_id = None,294 push_to_hub_organization = None,295 push_to_hub_token = None,296 mp_parameters = '',297 auto_find_batch_size = False,298 full_determinism = False,299 torchdynamo = None,300 ray_scope = 'last',301 ddp_timeout = 1800,302 torch_compile = False,303 torch_compile_backend = None,304 torch_compile_mode = None,305 dispatch_batches = None,306 split_batches = None,307 include_tokens_per_second = False,308 include_num_input_tokens_seen = False,309 neftune_noise_alpha = None,310 optim_target_modules = None,311 batch_eval_metrics = False,312 eval_on_start = False,313 use_liger_kernel = False,314 eval_use_gather_object = False,315 average_tokens_across_devices = False,316 model_init_kwargs = None,317 ref_model_init_kwargs = None,318 model_adapter_name = None,319 ref_adapter_name = None,320 force_use_ref_model = False,321 disable_dropout = True,322 use_logits_to_keep = False,323 dataset_num_proc = None,324 padding_value = None,325 label_pad_token_id = -100,326 max_prompt_length = 512,327 max_completion_length = None,328 max_length = 1024,329 truncation_mode = 'keep_end',330 padding_free = False,331 precompute_ref_log_probs = False,332 precompute_ref_batch_size = None,333 tools = None,334 loss_type = 'sigmoid',335 beta = 0.1,336 f_alpha_divergence_coef = 1.0,337 reference_free = False,338 label_smoothing = 0.0,339 use_weighting = False,340 rpo_alpha = None,341 discopop_tau = 0.05,342 sync_ref_model = False,343 ref_model_mixup_alpha = 0.9,344 ref_model_sync_steps = 64,345 generate_during_eval = False,346 use_num_logits_to_keep = False,347 vllm_sampling_params = None,348 unsloth_num_chunks = -1,349 **kwargs,350 ):351 if learning_rate < 1e-7: raise FloatingPointError(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!')352 if learning_rate > 1: raise OverflowError(f'Unsloth: Your learning rate of `{learning_rate}` is way too larger > 1! Consider decreasing it to 1e-1, otherwise gradient updates will explode!')353 if output_dir is None and save_strategy == 'steps' and save_steps == 500:354 output_dir = 'unsloth_training_checkpoints'355 save_strategy = 'no'356 if dataset_num_proc is None:357 from multiprocessing import cpu_count358 dataset_num_proc = cpu_count()359 360 super().__init__(361 output_dir = output_dir,362 overwrite_output_dir = overwrite_output_dir,363 do_train = do_train,364 do_eval = do_eval,365 do_predict = do_predict,366 eval_strategy = eval_strategy,367 prediction_loss_only = prediction_loss_only,368 per_device_train_batch_size = per_device_train_batch_size,369 per_device_eval_batch_size = per_device_eval_batch_size,370 per_gpu_train_batch_size = per_gpu_train_batch_size,371 per_gpu_eval_batch_size = per_gpu_eval_batch_size,372 gradient_accumulation_steps = gradient_accumulation_steps,373 eval_accumulation_steps = eval_accumulation_steps,374 eval_delay = eval_delay,375 torch_empty_cache_steps = torch_empty_cache_steps,376 learning_rate = learning_rate,377 weight_decay = weight_decay,378 adam_beta1 = adam_beta1,379 adam_beta2 = adam_beta2,380 adam_epsilon = adam_epsilon,381 max_grad_norm = max_grad_norm,382 num_train_epochs = num_train_epochs,383 max_steps = max_steps,384 lr_scheduler_type = lr_scheduler_type,385 warmup_ratio = warmup_ratio,386 warmup_steps = warmup_steps,387 log_level = log_level,388 log_level_replica = log_level_replica,389 log_on_each_node = log_on_each_node,390 logging_dir = logging_dir,391 logging_strategy = logging_strategy,392 logging_first_step = logging_first_step,393 logging_steps = logging_steps,394 logging_nan_inf_filter = logging_nan_inf_filter,395 save_strategy = save_strategy,396 save_steps = save_steps,397 save_total_limit = save_total_limit,398 save_safetensors = save_safetensors,399 save_on_each_node = save_on_each_node,400 save_only_model = save_only_model,401 restore_callback_states_from_checkpoint = restore_callback_states_from_checkpoint,402 no_cuda = no_cuda,403 use_cpu = use_cpu,404 use_mps_device = use_mps_device,405 seed = seed,406 data_seed = data_seed,407 jit_mode_eval = jit_mode_eval,408 use_ipex = use_ipex,409 bf16 = bf16,410 fp16 = fp16,411 fp16_opt_level = fp16_opt_level,412 half_precision_backend = half_precision_backend,413 bf16_full_eval = bf16_full_eval,414 fp16_full_eval = fp16_full_eval,415 tf32 = tf32,416 local_rank = local_rank,417 ddp_backend = ddp_backend,418 tpu_num_cores = tpu_num_cores,419 tpu_metrics_debug = tpu_metrics_debug,420 debug = debug,421 dataloader_drop_last = dataloader_drop_last,422 eval_steps = eval_steps,423 dataloader_num_workers = dataloader_num_workers,424 dataloader_prefetch_factor = dataloader_prefetch_factor,425 past_index = past_index,426 run_name = run_name,427 disable_tqdm = disable_tqdm,428 remove_unused_columns = remove_unused_columns,429 label_names = label_names,430 load_best_model_at_end = load_best_model_at_end,431 metric_for_best_model = metric_for_best_model,432 greater_is_better = greater_is_better,433 ignore_data_skip = ignore_data_skip,434 fsdp = fsdp,435 fsdp_min_num_params = fsdp_min_num_params,436 fsdp_config = fsdp_config,437 tp_size = tp_size,438 fsdp_transformer_layer_cls_to_wrap = fsdp_transformer_layer_cls_to_wrap,439 accelerator_config = accelerator_config,440 deepspeed = deepspeed,441 label_smoothing_factor = label_smoothing_factor,442 optim = optim,443 optim_args = optim_args,444 adafactor = adafactor,445 group_by_length = group_by_length,446 length_column_name = length_column_name,447 report_to = report_to,448 ddp_find_unused_parameters = ddp_find_unused_parameters,449 ddp_bucket_cap_mb = ddp_bucket_cap_mb,450 ddp_broadcast_buffers = ddp_broadcast_buffers,451 dataloader_pin_memory = dataloader_pin_memory,452 dataloader_persistent_workers = dataloader_persistent_workers,453 skip_memory_metrics = skip_memory_metrics,454 use_legacy_prediction_loop = use_legacy_prediction_loop,455 push_to_hub = push_to_hub,456 resume_from_checkpoint = resume_from_checkpoint,457 hub_model_id = hub_model_id,458 hub_strategy = hub_strategy,459 hub_token = hub_token,460 hub_private_repo = hub_private_repo,461 hub_always_push = hub_always_push,462 gradient_checkpointing = gradient_checkpointing,463 gradient_checkpointing_kwargs = gradient_checkpointing_kwargs,464 include_inputs_for_metrics = include_inputs_for_metrics,465 eval_do_concat_batches = eval_do_concat_batches,466 fp16_backend = fp16_backend,467 evaluation_strategy = evaluation_strategy,468 push_to_hub_model_id = push_to_hub_model_id,469 push_to_hub_organization = push_to_hub_organization,470 push_to_hub_token = push_to_hub_token,471 mp_parameters = mp_parameters,472 auto_find_batch_size = auto_find_batch_size,473 full_determinism = full_determinism,474 torchdynamo = torchdynamo,475 ray_scope = ray_scope,476 ddp_timeout = ddp_timeout,477 torch_compile = torch_compile,478 torch_compile_backend = torch_compile_backend,479 torch_compile_mode = torch_compile_mode,480 dispatch_batches = dispatch_batches,481 split_batches = split_batches,482 include_tokens_per_second = include_tokens_per_second,483 include_num_input_tokens_seen = include_num_input_tokens_seen,484 neftune_noise_alpha = neftune_noise_alpha,485 optim_target_modules = optim_target_modules,486 batch_eval_metrics = batch_eval_metrics,487 eval_on_start = eval_on_start,488 use_liger_kernel = use_liger_kernel,489 eval_use_gather_object = eval_use_gather_object,490 average_tokens_across_devices = average_tokens_across_devices,491 model_init_kwargs = model_init_kwargs,492 ref_model_init_kwargs = ref_model_init_kwargs,493 model_adapter_name = model_adapter_name,494 ref_adapter_name = ref_adapter_name,495 force_use_ref_model = force_use_ref_model,496 disable_dropout = disable_dropout,497 use_logits_to_keep = use_logits_to_keep,498 dataset_num_proc = dataset_num_proc,499 padding_value = padding_value,500 label_pad_token_id = label_pad_token_id,501 max_prompt_length = max_prompt_length,502 max_completion_length = max_completion_length,503 max_length = max_length,504 truncation_mode = truncation_mode,505 padding_free = padding_free,506 precompute_ref_log_probs = precompute_ref_log_probs,507 precompute_ref_batch_size = precompute_ref_batch_size,508 tools = tools,509 loss_type = loss_type,510 beta = beta,511 f_alpha_divergence_coef = f_alpha_divergence_coef,512 reference_free = reference_free,513 label_smoothing = label_smoothing,514 use_weighting = use_weighting,515 rpo_alpha = rpo_alpha,516 discopop_tau = discopop_tau,517 sync_ref_model = sync_ref_model,518 ref_model_mixup_alpha = ref_model_mixup_alpha,519 ref_model_sync_steps = ref_model_sync_steps,520 generate_during_eval = generate_during_eval,521 use_num_logits_to_keep = use_num_logits_to_keep,**kwargs)522 self.vllm_sampling_params = vllm_sampling_params523 self.unsloth_num_chunks = unsloth_num_chunks524pass525 526class _UnslothDPOTrainer(Trainer):527 r""""""528 529 _tag_names = ["trl", "dpo"]530 531 @deprecate_kwarg(532 "tokenizer", "0.16.0", "processing_class", warn_if_greater_or_equal_version=True, raise_if_both_names=True533 )534 def __init__(535 self,536 model: Optional[Union[PreTrainedModel, nn.Module, str]] = None,537 ref_model: Optional[Union[PreTrainedModel, nn.Module, str]] = None,538 args: Optional[DPOConfig] = None,539 data_collator: Optional[DataCollator] = None,540 train_dataset: Optional[Dataset] = None,541 eval_dataset: Optional[Union[Dataset, dict[str, Dataset]]] = None,542 processing_class: Optional[543 Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin]544 ] = None,545 model_init: Optional[Callable[[], PreTrainedModel]] = None,546 compute_metrics: Optional[Callable[[EvalLoopOutput], dict]] = None,547 callbacks: Optional[list[TrainerCallback]] = None,548 optimizers: tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),549 preprocess_logits_for_metrics: Optional[Callable[[torch.Tensor, torch.Tensor], torch.Tensor]] = None,550 peft_config: Optional[dict] = None,551 ):552 if model is None:553 raise ValueError("No model provided. Please provide a model to train.")554 555 if not isinstance(model, str) and ref_model is model:556 raise ValueError(557 "`model` and `ref_model` cannot be the same object. If you want `ref_model` to be the "558 "same as `model`, you must mass a copy of it, or `None` if you use peft."559 )560 561 if args.model_init_kwargs is None:562 model_init_kwargs = {}563 elif not isinstance(model, str):564 raise ValueError(565 "You passed model_init_kwargs to the DPOTrainer/DPOConfig, but your model is already instantiated."566 )567 else:568 model_init_kwargs = args.model_init_kwargs569 torch_dtype = model_init_kwargs.get("torch_dtype")570 if torch_dtype is not None:571 # Convert to `torch.dtype` if an str is passed572 if isinstance(torch_dtype, str) and torch_dtype != "auto":573 torch_dtype = getattr(torch, torch_dtype)574 if torch_dtype != "auto" and not isinstance(torch_dtype, torch.dtype):575 raise ValueError(576 f"Invalid `torch_dtype` passed to the DPOConfig. Expected a string with either `torch.dtype` or 'auto', but got {torch_dtype}."577 )578 model_init_kwargs["torch_dtype"] = torch_dtype579 580 if args.ref_model_init_kwargs is None:581 ref_model_init_kwargs = {}582 elif not isinstance(ref_model, str):583 raise ValueError(584 "You passed ref_model_init_kwargs to the DPOTrainer/DPOConfig, but your ref_model is already instantiated."585 )586 else:587 ref_model_init_kwargs = args.ref_model_init_kwargs588 torch_dtype = ref_model_init_kwargs.get("torch_dtype")589 if torch_dtype is not None:590 # Convert to `torch.dtype` if an str is passed591 if isinstance(torch_dtype, str) and torch_dtype != "auto":592 torch_dtype = getattr(torch, torch_dtype)593 if torch_dtype != "auto" and not isinstance(torch_dtype, torch.dtype):594 raise ValueError(595 f"Invalid `torch_dtype` passed to the DPOConfig. Expected a string with either `torch.dtype` or 'auto', but got {torch_dtype}."596 )597 ref_model_init_kwargs["torch_dtype"] = torch_dtype598 599 if isinstance(model, str):600 model = AutoModelForCausalLM.from_pretrained(model, **model_init_kwargs)601 602 if isinstance(ref_model, str):603 ref_model = AutoModelForCausalLM.from_pretrained(ref_model, **ref_model_init_kwargs)604 605 # Initialize this variable to False. This helps tracking the case when `peft_module_casting_to_bf16`606 # has been called in order to properly call autocast if needed.607 self._peft_has_been_casted_to_bf16 = False608 609 if not is_peft_available() and peft_config is not None:610 raise ValueError(611 "PEFT is not installed and you passed a `peft_config` in the trainer's kwargs, please install it to use the PEFT models"612 )613 elif is_peft_available() and peft_config is not None:614 # if model is a peft model and we have a peft_config, we merge and unload it first615 if isinstance(model, PeftModel):616 model = model.merge_and_unload()617 618 if ref_model is not None and not args.force_use_ref_model:619 raise ValueError(620 "You passed both a ref_model and a peft_config. For training PEFT adapters with DPO there is no need to pass a reference"621 " model. Please pass `ref_model=None` in case you want to train PEFT adapters, or pass a ref_model with `force_use_ref_model=True` in DPOTrainer's init."622 " if you want to use a different ref_model."623 )624 625 if getattr(model, "is_loaded_in_8bit", False) or getattr(model, "is_loaded_in_4bit", False):626 _support_gc_kwargs = hasattr(627 args, "gradient_checkpointing_kwargs"628 ) and "gradient_checkpointing_kwargs" in list(629 inspect.signature(prepare_model_for_kbit_training).parameters630 )631 632 prepare_model_kwargs = {"use_gradient_checkpointing": args.gradient_checkpointing}633 634 if _support_gc_kwargs:635 prepare_model_kwargs["gradient_checkpointing_kwargs"] = args.gradient_checkpointing_kwargs636 637 model = prepare_model_for_kbit_training(model, **prepare_model_kwargs)638 elif getattr(args, "gradient_checkpointing", False):639 # For backward compatibility with older versions of transformers640 if hasattr(model, "enable_input_require_grads"):641 model.enable_input_require_grads()642 else:643 644 def make_inputs_require_grad(module, input, output):645 output.requires_grad_(True)646 647 model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)648 649 # get peft model with the given config650 model = model651 if args.bf16 and getattr(model, "is_loaded_in_4bit", False):652 peft_module_casting_to_bf16(model)653 # If args.bf16 we need to explicitly call `generate` with torch amp autocast context manager654 self._peft_has_been_casted_to_bf16 = True655 656 # For models that use gradient_checkpointing, we need to attach a hook that enables input657 # to explicitly have `requires_grad=True`, otherwise training will either silently658 # fail or completely fail.659 elif getattr(args, "gradient_checkpointing", False):660 # For backward compatibility with older versions of transformers661 if hasattr(model, "enable_input_require_grads"):662 model.enable_input_require_grads()663 else:664 665 def make_inputs_require_grad(module, input, output):666 output.requires_grad_(True)667 668 model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)669 670 if args.generate_during_eval and not (is_wandb_available() or is_comet_available()):671 raise ValueError(672 "`generate_during_eval=True` requires Weights and Biases or Comet to be installed."673 " Please install `wandb` or `comet-ml` to resolve."674 )675 676 self.is_encoder_decoder = model.config.is_encoder_decoder677 self.is_vision_model = model.config.model_type in MODEL_FOR_VISION_2_SEQ_MAPPING_NAMES.keys()678 self.is_peft_model = is_peft_available() and isinstance(model, PeftModel)679 self.model_adapter_name = args.model_adapter_name680 self.ref_adapter_name = args.ref_adapter_name681 self.reference_free = args.reference_free682 683 if ref_model:684 self.ref_model = ref_model685 elif self.is_peft_model or args.precompute_ref_log_probs:686 # The `model` with adapters turned off will be used as the reference model687 self.ref_model = None688 else:689 self.ref_model = create_reference_model(model)690 691 if processing_class is None:692 raise ValueError("processing_class must be specified to tokenize a DPO dataset.")693 694 if args.padding_value is not None:695 self.padding_value = args.padding_value696 else:697 if hasattr(processing_class, "pad_token_id") and processing_class.pad_token_id is not None:698 self.padding_value = processing_class.pad_token_id699 elif hasattr(processing_class, "tokenizer") and processing_class.tokenizer.pad_token_id is not None:700 self.padding_value = processing_class.tokenizer.pad_token_id701 else:702 raise ValueError(703 "`padding_value` is not specified in `DPOConfig`, and `pad_token_id` is missing in the "704 "`processing_class`. Please either set the `padding_value` argument in `DPOConfig`, or set "705 "`tokenizer.pad_token` (e.g., `tokenizer.pad_token = tokenizer.eos_token`) before instantiating "706 "the trainer."707 )708 709 if data_collator is None:710 data_collator = DataCollatorForPreference(pad_token_id=self.padding_value)711 712 # Disable dropout in the model and reference model713 if args.disable_dropout:714 disable_dropout_in_model(model)715 if self.ref_model is not None:716 disable_dropout_in_model(self.ref_model)717 718 self.generate_during_eval = args.generate_during_eval719 self.label_pad_token_id = args.label_pad_token_id720 self.max_prompt_length = args.max_prompt_length721 self.max_completion_length = args.max_completion_length722 self.max_length = args.max_length723 self.truncation_mode = args.truncation_mode724 self.precompute_ref_log_probs = args.precompute_ref_log_probs725 self.use_logits_to_keep = args.use_logits_to_keep726 727 if args.padding_free:728 if model.config._attn_implementation != "flash_attention_2":729 warnings.warn(730 "Padding-free training is enabled, but the attention implementation is not set to "731 "'flash_attention_2'. Padding-free training flattens batches into a single sequence, and "732 "'flash_attention_2' is the only known attention mechanism that reliably supports this. Using "733 "other implementations may lead to unexpected behavior. To ensure compatibility, set "734 "`attn_implementation='flash_attention_2'` in the model configuration, or verify that your "735 "attention mechanism can handle flattened sequences."736 )737 self.padding_free = args.padding_free738 739 # Since ref_logs are precomputed on the first call to get_train/eval_dataloader740 # keep track of first called to avoid computation of future calls741 self._precomputed_train_ref_log_probs = False742 self._precomputed_eval_ref_log_probs = False743 744 if (745 args.loss_type in ["hinge", "ipo", "bco_pair", "sppo_hard", "nca_pair", "apo_zero", "apo_down"]746 and args.label_smoothing > 0747 ):748 warnings.warn(749 f"You are using the {args.loss_type} loss type that does not support label smoothing. The "750 "`label_smoothing` parameter will be ignored. Set `label_smoothing` to `0.0` to remove this warning.",751 UserWarning,752 )753 if args.loss_type == "kto_pair":754 raise ValueError("Support for kto_pair has been removed in DPOTrainer. Please use KTOTrainer.")755 756 self.beta = args.beta757 self.label_smoothing = args.label_smoothing758 self.loss_type = args.loss_type759 self.aux_loss_enabled = getattr(model.config, "output_router_logits", False)760 self.use_weighting = args.use_weighting761 self.aux_loss_coef = getattr(model.config, "router_aux_loss_coef", 0.0)762 if self.aux_loss_enabled and self.aux_loss_coef == 0.0:763 warnings.warn(764 "You set `output_router_logits` to `True` in the model config, but `router_aux_loss_coef` is set to "765 "`0.0`, meaning the auxiliary loss will not be used. Either set `router_aux_loss_coef` to a value "766 "greater than `0.0`, or set `output_router_logits` to `False` if you don't want to use the auxiliary "767 "loss.",768 UserWarning,769 )770 771 self._stored_metrics = defaultdict(lambda: defaultdict(list))772 self.f_divergence_type = args.f_divergence_type773 self.f_divergence_params = {FDivergenceConstants.ALPHA_DIVERGENCE_COEF_KEY: args.f_alpha_divergence_coef}774 self.dataset_num_proc = args.dataset_num_proc775 776 # The trainer estimates the number of FLOPs (floating-point operations) using the number of elements in the777 # input tensor associated with the key "input_ids". However, in DPO, the sampled data does not include the778 # "input_ids" key. Instead, the available keys are "prompt_input_ids", "chosen_input_ids", and779 # "rejected_input_ids". As a result, the trainer issues the warning: "Could not estimate the number of tokens780 # of the input, floating-point operations will not be computed." To suppress this warning, we set the781 # "estimate_tokens" key in the model's "warnings_issued" dictionary to True. This acts as a flag to indicate782 # that the warning has already been issued.783 model.warnings_issued["estimate_tokens"] = True784 785 # Dataset preparation786 train_dataset = self._prepare_dataset(train_dataset, processing_class, args, "train")787 if eval_dataset is not None:788 if isinstance(eval_dataset, dict):789 eval_dataset = {790 key: self._prepare_dataset(dataset, processing_class, args, key)791 for key, dataset in eval_dataset.items()792 }793 else:794 eval_dataset = self._prepare_dataset(eval_dataset, processing_class, args, "eval")795 796 super().__init__(797 model=model,798 args=args,799 data_collator=data_collator,800 train_dataset=train_dataset,801 eval_dataset=eval_dataset,802 processing_class=processing_class,803 model_init=model_init,804 compute_metrics=compute_metrics,805 callbacks=callbacks,806 optimizers=optimizers,807 preprocess_logits_for_metrics=preprocess_logits_for_metrics,808 )809 810 # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the811 # model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set812 # self.model_accepts_loss_kwargs to False to enable scaling.813 self.model_accepts_loss_kwargs = False814 815 # Add tags for models that have been loaded with the correct transformers version816 if hasattr(self.model, "add_model_tags"):817 self.model.add_model_tags(self._tag_names)818 819 if not hasattr(self, "accelerator"):820 raise AttributeError(821 "Your `Trainer` does not have an `accelerator` object. Consider upgrading `transformers`."822 )823 824 # Deepspeed Zero-3 does not support precompute_ref_log_probs825 if self.is_deepspeed_enabled:826 if self.accelerator.state.deepspeed_plugin.zero_stage == 3 and self.precompute_ref_log_probs:827 raise ValueError(828 "You cannot use `precompute_ref_log_probs=True` with Deepspeed ZeRO-3. Please set `precompute_ref_log_probs=False`."829 )830 831 if self.ref_model is None:832 if not (self.is_peft_model or self.precompute_ref_log_probs):833 raise ValueError(834 "No reference model and model is not a Peft model. Try setting `precompute_ref_log_probs=True`"835 )836 if args.sync_ref_model:837 raise ValueError(838 "You currently cannot use `ref_model=None` with TR-DPO method. Please provide `ref_model`."839 )840 else:841 if self.is_deepspeed_enabled:842 self.ref_model = self._prepare_deepspeed(self.ref_model)843 else:844 self.ref_model = self.accelerator.prepare_model(self.ref_model, evaluation_mode=True)845 846 if args.sync_ref_model:847 if self.precompute_ref_log_probs:848 raise ValueError(849 "You cannot use `precompute_ref_log_probs=True` with TR-DPO method. Please set `precompute_ref_log_probs=False`."850 )851 852 self.add_callback(SyncRefModelCallback(ref_model=self.ref_model, accelerator=self.accelerator))853 854 if self.loss_type == "bco_pair":855 self.running = RunningMoments(self.accelerator)856 857 def _prepare_dataset(858 self,859 dataset: Union[Dataset, IterableDataset],860 processing_class: Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin],861 args: DPOConfig,862 dataset_name: str,863 ) -> Union[Dataset, IterableDataset]:864 # Build the kwargs for the `map` function865 map_kwargs = {"writer_batch_size": 10}866 if isinstance(dataset, Dataset): # IterableDataset does not support num_proc867 map_kwargs["num_proc"] = args.dataset_num_proc868 869 with PartialState().local_main_process_first():870 # Extract prompt if needed871 if isinstance(dataset, Dataset): # `IterableDataset.map` does not support `desc`872 map_kwargs["desc"] = f"Extracting prompt in {dataset_name} dataset"873 dataset = dataset.map(maybe_extract_prompt, **map_kwargs)874 875 # Apply the chat template if needed876 if isinstance(dataset, Dataset): # `IterableDataset.map` does not support `desc`877 map_kwargs["desc"] = f"Applying chat template to {dataset_name} dataset"878 dataset = dataset.map(879 maybe_apply_chat_template, fn_kwargs={"tokenizer": processing_class, "tools": args.tools}, **map_kwargs880 )881 882 # Tokenize the dataset883 if isinstance(dataset, Dataset): # `IterableDataset.map` does not support `desc`884 map_kwargs["desc"] = f"Tokenizing {dataset_name} dataset"885 886 dataset = dataset.map(887 self.tokenize_row if not self.is_vision_model else self.process_row,888 remove_columns=["prompt", "chosen", "rejected"],889 fn_kwargs={890 "processing_class": processing_class,891 "max_prompt_length": args.max_prompt_length,892 "max_completion_length": args.max_completion_length,893 # for enc-dec, we add the special tokens ([bos_token] + prompt + [eos_token]; completion + [eos_token])894 "add_special_tokens": False,895 },896 **map_kwargs,897 )898 899 return dataset900 901 @staticmethod902 def tokenize_row(features, processing_class, max_prompt_length, max_completion_length, add_special_tokens):903 """904 Tokenize a row of the dataset.905 906 Args:907 features (`dict[str, str]`):908 Row of the dataset, should contain the keys `"prompt"`, `"chosen"`, and `"rejected"`.909 processing_class (`PreTrainedTokenizerBase`):910 Processing class used to process the data.911 max_prompt_length (`int` or `None`):912 Maximum length of the prompt sequence. If `None`, the prompt sequence is not truncated.913 max_completion_length (`int` or `None`):914 Maximum length of the completion sequences. If `None`, the completion sequences are not truncated.915 add_special_tokens (`bool`):916 Whether to add special tokens to the sequences. Typically used for encoder-decoder models. If `True`,917 the prompt sequence will have a bos token prepended and an eos token appended. In any case, the918 completion sequences will have an eos token appended.919 920 Returns:921 `dict[str, list[int]]`:922 Tokenized sequences with the keys `"prompt_input_ids"`, `"chosen_input_ids"`, and923 `"rejected_input_ids".924 925 Example:926 ```python927 >>> from transformers import GPT2Tokenizer928 >>> tokenizer = GPT2Tokenizer.from_pretrained("gpt2")929 >>> features = {"prompt": "The sky is", "chosen": " blue", "rejected": " green"}930 >>> DPOTrainer.tokenize_row(931 ... features, tokenizer, max_prompt_length=3, max_completion_length=3, add_special_tokens=False932 ... )933 {'prompt_input_ids': [464, 6766, 318], 'chosen_input_ids': [4171, 50256], 'rejected_input_ids': [4077, 50256]}934 ```935 """936 tokenizer = processing_class # the processing class is a tokenizer937 prompt_input_ids = tokenizer(features["prompt"], add_special_tokens=False)["input_ids"]938 chosen_input_ids = tokenizer(features["chosen"], add_special_tokens=False)["input_ids"]939 rejected_input_ids = tokenizer(features["rejected"], add_special_tokens=False)["input_ids"]940 941 # Add special tokens (typically for encoder-decoder models)942 if add_special_tokens:943 if tokenizer.bos_token_id is not None:944 prompt_input_ids = [tokenizer.bos_token_id] + prompt_input_ids945 if tokenizer.eos_token_id is not None:946 prompt_input_ids = prompt_input_ids + [tokenizer.eos_token_id]947 chosen_input_ids = chosen_input_ids + [tokenizer.eos_token_id]948 rejected_input_ids = rejected_input_ids + [tokenizer.eos_token_id]949 950 # Truncate prompt and completion sequences951 if max_prompt_length is not None:952 prompt_input_ids = prompt_input_ids[-max_prompt_length:]953 if max_completion_length is not None:954 chosen_input_ids = chosen_input_ids[:max_completion_length]955 rejected_input_ids = rejected_input_ids[:max_completion_length]956 957 return {958 "prompt_input_ids": prompt_input_ids,959 "chosen_input_ids": chosen_input_ids,960 "rejected_input_ids": rejected_input_ids,961 }962 963 @staticmethod964 def process_row(features, processing_class, max_prompt_length, max_completion_length, add_special_tokens):965 """966 Same as `tokenize_row` but for vision models. Please refer to `tokenize_row` for more information.967 """968 processor, tokenizer = processing_class, processing_class.tokenizer # the processing class is a processor969 processed_features = processor(images=features["images"], text=features["prompt"], add_special_tokens=False)970 971 prompt_input_ids = processed_features["input_ids"][0]972 pixel_values = processed_features["pixel_values"][0]973 chosen_input_ids = tokenizer(features["chosen"], add_special_tokens=False)["input_ids"]974 rejected_input_ids = tokenizer(features["rejected"], add_special_tokens=False)["input_ids"]975 976 # Add special tokens (typically for encoder-decoder models)977 if add_special_tokens:978 if tokenizer.bos_token_id is not None:979 prompt_input_ids = [tokenizer.bos_token_id] + prompt_input_ids980 if tokenizer.eos_token_id is not None:981 prompt_input_ids = prompt_input_ids + [tokenizer.eos_token_id]982 chosen_input_ids = chosen_input_ids + [tokenizer.eos_token_id]983 rejected_input_ids = rejected_input_ids + [tokenizer.eos_token_id]984 985 # Truncate prompt and completion sequences986 if max_prompt_length is not None:987 prompt_input_ids = prompt_input_ids[-max_prompt_length:]988 if max_completion_length is not None:989 chosen_input_ids = chosen_input_ids[:max_completion_length]990 rejected_input_ids = rejected_input_ids[:max_completion_length]991 992 output = {993 "prompt_input_ids": prompt_input_ids,994 "pixel_values": pixel_values,995 "chosen_input_ids": chosen_input_ids,996 "rejected_input_ids": rejected_input_ids,997 }998 999 if "pixel_attention_mask" in processed_features:1000 output["pixel_attention_mask"] = processed_features["pixel_attention_mask"][0]1001 if "image_sizes" in processed_features:1002 output["image_sizes"] = processed_features["image_sizes"][0]1003 1004 return output1005 1006 def _prepare_deepspeed(self, model: PreTrainedModelWrapper):1007 # Adapted from accelerate: https://github.com/huggingface/accelerate/blob/739b135f8367becb67ffaada12fe76e3aa60fefd/src/accelerate/accelerator.py#L14731008 deepspeed_plugin = self.accelerator.state.deepspeed_plugin1009 config_kwargs = deepcopy(deepspeed_plugin.deepspeed_config)1010 1011 if model is not None:1012 if hasattr(model, "config"):1013 hidden_size = (1014 max(model.config.hidden_sizes)1015 if getattr(model.config, "hidden_sizes", None)1016 else getattr(model.config, "hidden_size", None)1017 )1018 if hidden_size is not None and config_kwargs["zero_optimization"]["stage"] == 3:1019 # Note that `stage3_prefetch_bucket_size` can produce DeepSpeed messages like: `Invalidate trace cache @ step 0: expected module 1, but got module 0`1020 # This is expected and is not an error, see: https://github.com/microsoft/DeepSpeed/discussions/40811021 config_kwargs.update(1022 {1023 "zero_optimization.reduce_bucket_size": hidden_size * hidden_size,1024 "zero_optimization.stage3_param_persistence_threshold": 10 * hidden_size,1025 "zero_optimization.stage3_prefetch_bucket_size": 0.9 * hidden_size * hidden_size,1026 }1027 )1028 1029 # If ZeRO-3 is used, we shard both the active and reference model.1030 # Otherwise, we assume the reference model fits in memory and is initialized on each device with ZeRO disabled (stage 0)1031 if config_kwargs["zero_optimization"]["stage"] != 3:1032 config_kwargs["zero_optimization"]["stage"] = 01033 model, *_ = deepspeed.initialize(model=model, config=config_kwargs)1034 model.eval()1035 return model1036 1037 def _set_signature_columns_if_needed(self):1038 # If `self.args.remove_unused_columns` is True, non-signature columns are removed.1039 # By default, this method sets `self._signature_columns` to the model's expected inputs.1040 # In DPOTrainer, we preprocess data, so using the model's signature columns doesn't work.1041 # Instead, we set them to the columns expected by `DataCollatorForPreference`, hence the override.1042 if self._signature_columns is None:1043 self._signature_columns = [1044 "prompt_input_ids",1045 "chosen_input_ids",1046 "rejected_input_ids",1047 "image_sizes",1048 "ref_chosen_logps",1049 "ref_rejected_logps",1050 ]1051 1052 def get_train_dataloader(self) -> DataLoader:1053 """1054 Returns the training [`~torch.utils.data.DataLoader`].1055 1056 Subclass of transformers.src.transformers.trainer.get_train_dataloader to precompute `ref_log_probs`.1057 """1058 1059 if self.precompute_ref_log_probs and not self._precomputed_train_ref_log_probs:1060 batch_size = self.args.precompute_ref_batch_size or self.args.per_device_train_batch_size1061 dataloader_params = {1062 "batch_size": batch_size,1063 "collate_fn": self.data_collator,1064 "num_workers": self.args.dataloader_num_workers,1065 "pin_memory": self.args.dataloader_pin_memory,1066 "shuffle": False,1067 }1068 1069 # prepare dataloader1070 data_loader = self.accelerator.prepare(DataLoader(self.train_dataset, **dataloader_params))1071 1072 ref_chosen_logps = []1073 ref_rejected_logps = []1074 for padded_batch in tqdm(iterable=data_loader, desc="Train dataset reference log probs"):1075 ref_chosen_logp, ref_rejected_logp = self.compute_ref_log_probs(padded_batch)1076 ref_chosen_logp, ref_rejected_logp = self.accelerator.gather_for_metrics(1077 (ref_chosen_logp, ref_rejected_logp)1078 )1079 ref_chosen_logps.append(ref_chosen_logp.cpu())1080 ref_rejected_logps.append(ref_rejected_logp.cpu())1081 1082 # Unnecessary cache clearing to avoid OOM1083 empty_cache()1084 self.accelerator.free_memory()1085 1086 all_ref_chosen_logps = torch.cat(ref_chosen_logps).float().numpy()1087 all_ref_rejected_logps = torch.cat(ref_rejected_logps).float().numpy()1088 1089 self.train_dataset = self.train_dataset.add_column(name="ref_chosen_logps", column=all_ref_chosen_logps)1090 self.train_dataset = self.train_dataset.add_column(1091 name="ref_rejected_logps", column=all_ref_rejected_logps1092 )1093 1094 self._precomputed_train_ref_log_probs = True1095 1096 return super().get_train_dataloader()1097 1098 def get_eval_dataloader(self, eval_dataset: Optional[Dataset] = None) -> DataLoader:1099 """1100 Returns the evaluation [`~torch.utils.data.DataLoader`].1101 1102 Subclass of transformers.src.transformers.trainer.get_eval_dataloader to precompute `ref_log_probs`.1103 1104 Args:1105 eval_dataset (`torch.utils.data.Dataset`, *optional*):1106 If provided, will override `self.eval_dataset`. If it is a [`~datasets.Dataset`], columns not accepted1107 by the `model.forward()` method are automatically removed. It must implement `__len__`.1108 """1109 if eval_dataset is None and self.eval_dataset is None:1110 raise ValueError("Trainer: evaluation requires an eval_dataset.")1111 eval_dataset = eval_dataset if eval_dataset is not None else self.eval_dataset1112 1113 if self.precompute_ref_log_probs and not self._precomputed_eval_ref_log_probs:1114 batch_size = self.args.precompute_ref_batch_size or self.args.per_device_eval_batch_size1115 dataloader_params = {1116 "batch_size": batch_size,1117 "collate_fn": self.data_collator,1118 "num_workers": self.args.dataloader_num_workers,1119 "pin_memory": self.args.dataloader_pin_memory,1120 "shuffle": False,1121 }1122 1123 # prepare dataloader1124 data_loader = self.accelerator.prepare(DataLoader(eval_dataset, **dataloader_params))1125 1126 ref_chosen_logps = []1127 ref_rejected_logps = []1128 for padded_batch in tqdm(iterable=data_loader, desc="Eval dataset reference log probs"):1129 ref_chosen_logp, ref_rejected_logp = self.compute_ref_log_probs(padded_batch)1130 ref_chosen_logp, ref_rejected_logp = self.accelerator.gather_for_metrics(1131 (ref_chosen_logp, ref_rejected_logp)1132 )1133 ref_chosen_logps.append(ref_chosen_logp.cpu())1134 ref_rejected_logps.append(ref_rejected_logp.cpu())1135 1136 all_ref_chosen_logps = torch.cat(ref_chosen_logps).float().numpy()1137 all_ref_rejected_logps = torch.cat(ref_rejected_logps).float().numpy()1138 1139 eval_dataset = eval_dataset.add_column(name="ref_chosen_logps", column=all_ref_chosen_logps)1140 eval_dataset = eval_dataset.add_column(name="ref_rejected_logps", column=all_ref_rejected_logps)1141 1142 # Save calculated ref_chosen_logps and ref_rejected_logps to the eval_dataset for subsequent runs1143 if self.eval_dataset is not None:1144 self.eval_dataset = eval_dataset1145 self._precomputed_eval_ref_log_probs = True1146 1147 return super().get_eval_dataloader(eval_dataset=eval_dataset)1148 1149 @contextmanager1150 def null_ref_context(self):1151 """Context manager for handling null reference model (that is, peft adapter manipulation)."""1152 with (1153 self.accelerator.unwrap_model(self.model).disable_adapter()1154 if self.is_peft_model and not self.ref_adapter_name1155 else nullcontext()1156 ):1157 if self.ref_adapter_name:1158 self.model.set_adapter(self.ref_adapter_name)1159 yield1160 if self.ref_adapter_name:1161 self.model.set_adapter(self.model_adapter_name or "default")1162 1163 def compute_ref_log_probs(self, batch: dict[str, torch.LongTensor]) -> dict:1164 """Computes log probabilities of the reference model for a single padded batch of a DPO specific dataset."""1165 device_type = "xpu" if is_torch_xpu_available() else "cuda"1166 compte_ref_context_manager = amp.autocast(device_type) if self._peft_has_been_casted_to_bf16 else nullcontext()1167 with torch.no_grad(), compte_ref_context_manager:1168 if self.ref_model is None:1169 with self.null_ref_context():1170 ref_model_output = self.concatenated_forward(self.model, batch)1171 else:1172 ref_model_output = self.concatenated_forward(self.ref_model, batch)1173 return ref_model_output["chosen_logps"], ref_model_output["rejected_logps"]1174 1175 @staticmethod1176 def concatenated_inputs(1177 batch: dict[str, Union[list, torch.LongTensor]], padding_value: int1178 ) -> dict[str, torch.LongTensor]:1179 """1180 Concatenate the `chosen` and `rejected` inputs from the batch into a single tensor for both the prompt1181 and completion sequences.1182 1183 Args:1184 batch (`dict[str, Union[list, torch.LongTensor]]`):1185 A batch of input data. The batch must contain the following keys:1186 1187 - `"prompt_input_ids"`: Tensor of shape `(batch_size, prompt_length)` representing the prompt input IDs.1188 - `"chosen_input_ids"`: Tensor of shape `(batch_size, chosen_length)` representing the chosen completion input IDs.1189 - `"rejected_input_ids"`: Tensor of shape `(batch_size, rejected_length)` representing the rejected completion input IDs.1190 - `"prompt_pixel_values"` (optional): Tensor for pixel values, if available.1191 - `"prompt_pixel_attention_mask"` (optional): Tensor for pixel attention masks, if available.1192 1193 padding_value (`int`):1194 The padding value to use for the concatenated completion sequences (`chosen_input_ids` and1195 `rejected_input_ids`).1196 1197 Returns:1198 `dict[str, torch.LongTensor]`: A dictionary containing:1199 1200 - `"prompt_input_ids"`: Concatenated prompt input IDs of shape `(2 * batch_size, prompt_length)`.