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.orpo_trainer import (Any, AutoModelForCausalLM, BaseImageProcessor, Callable, DPODataCollatorWithPadding, DataCollator, DataLoader, Dataset, EvalLoopOutput, F, FeatureExtractionMixin, Literal, ORPOConfig, ORPOTrainer, Optional, PartialState, PeftModel, PreTrainedModel, PreTrainedModelWrapper, PreTrainedTokenizerBase, ProcessorMixin, Trainer, TrainerCallback, Union, add_bos_token_if_needed, add_eos_token_if_needed, amp, deepcopy, defaultdict, disable_dropout_in_model, generate_model_card, get_comet_experiment_url, inspect, is_comet_available, is_peft_available, is_torch_fx_proxy, is_torch_xla_available, is_wandb_available, log_table_to_comet_experiment, maybe_apply_chat_template, maybe_extract_prompt, nn, np, nullcontext, os, pad_to_length, pd, peft_module_casting_to_bf16, prepare_model_for_kbit_training, random, textwrap, torch, 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 UnslothORPOConfig(ORPOConfig):44 """45 46 Configuration class for the [`ORPOTrainer`].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 learning_rate (`float`, *optional*, defaults to `1e-6`):54 Initial learning rate for [`AdamW`] optimizer. The default value replaces that of55 [`~transformers.TrainingArguments`].56 max_length (`int` or `None`, *optional*, defaults to `1024`):57 Maximum length of the sequences (prompt + completion) in the batch. This argument is required if you want58 to use the default data collator.59 max_prompt_length (`int` or `None`, *optional*, defaults to `512`):60 Maximum length of the prompt. This argument is required if you want to use the default data collator.61 max_completion_length (`int` or `None`, *optional*, defaults to `None`):62 Maximum length of the completion. This argument is required if you want to use the default data collator63 and your model is an encoder-decoder.64 beta (`float`, *optional*, defaults to `0.1`):65 Parameter controlling the relative ratio loss weight in the ORPO loss. In the [paper](https://huggingface.co/papers/2403.07691),66 it is denoted by λ. In the [code](https://github.com/xfactlab/orpo), it is denoted by `alpha`.67 disable_dropout (`bool`, *optional*, defaults to `True`):68 Whether to disable dropout in the model.69 label_pad_token_id (`int`, *optional*, defaults to `-100`):70 Label pad token id. This argument is required if you want to use the default data collator.71 padding_value (`int` or `None`, *optional*, defaults to `None`):72 Padding value to use. If `None`, the padding value of the tokenizer is used.73 truncation_mode (`str`, *optional*, defaults to `"keep_end"`):74 Truncation mode to use when the prompt is too long. Possible values are `"keep_end"` or `"keep_start"`.75 This argument is required if you want to use the default data collator.76 generate_during_eval (`bool`, *optional*, defaults to `False`):77 If `True`, generates and logs completions from the model to W&B or Comet during evaluation.78 is_encoder_decoder (`bool` or `None`, *optional*, defaults to `None`):79 When using the `model_init` argument (callable) to instantiate the model instead of the `model` argument,80 you need to specify if the model returned by the callable is an encoder-decoder model.81 model_init_kwargs (`dict[str, Any]` or `None`, *optional*, defaults to `None`):82 Keyword arguments to pass to `AutoModelForCausalLM.from_pretrained` when instantiating the model from a83 string.84 dataset_num_proc (`int` or `None`, *optional*, defaults to `None`):85 Number of processes to use for processing the dataset.86 87 """88 vllm_sampling_params: Optional[Any] = field(89 default = None,90 metadata = {'help': 'vLLM SamplingParams'},91 )92 unsloth_num_chunks : Optional[int] = field(93 default = -1,94 metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},95 )96 def __init__(97 self,98 output_dir = None,99 overwrite_output_dir = None,100 do_train = False,101 do_eval = False,102 do_predict = False,103 eval_strategy = 'no',104 prediction_loss_only = False,105 per_device_train_batch_size = 4,106 per_device_eval_batch_size = 4,107 per_gpu_train_batch_size = None,108 per_gpu_eval_batch_size = None,109 gradient_accumulation_steps = 2,110 eval_accumulation_steps = 2,111 eval_delay = 0,112 torch_empty_cache_steps = 250,113 learning_rate = 5e-05,114 weight_decay = 0.01,115 adam_beta1 = 0.9,116 adam_beta2 = 0.999,117 adam_epsilon = 1e-08,118 max_grad_norm = 1.0,119 num_train_epochs = 3.0,120 max_steps = -1,121 lr_scheduler_type = 'linear',122 warmup_ratio = 0.1,123 warmup_steps = 0,124 log_level = 'passive',125 log_level_replica = 'warning',126 log_on_each_node = True,127 logging_dir = None,128 logging_strategy = 'steps',129 logging_first_step = False,130 logging_steps = 1,131 logging_nan_inf_filter = False,132 save_strategy = 'steps',133 save_steps = 500,134 save_total_limit = None,135 save_safetensors = True,136 save_on_each_node = False,137 save_only_model = False,138 restore_callback_states_from_checkpoint = False,139 no_cuda = False,140 use_cpu = False,141 use_mps_device = False,142 seed = 3407,143 data_seed = 3407,144 jit_mode_eval = False,145 use_ipex = False,146 bf16 = False,147 fp16 = False,148 fp16_opt_level = 'O1',149 half_precision_backend = 'auto',150 bf16_full_eval = False,151 fp16_full_eval = False,152 tf32 = None,153 local_rank = -1,154 ddp_backend = None,155 tpu_num_cores = None,156 tpu_metrics_debug = False,157 debug = '',158 dataloader_drop_last = False,159 eval_steps = None,160 dataloader_num_workers = 0,161 dataloader_prefetch_factor = None,162 past_index = -1,163 run_name = None,164 disable_tqdm = None,165 remove_unused_columns = True,166 label_names = None,167 load_best_model_at_end = False,168 metric_for_best_model = None,169 greater_is_better = None,170 ignore_data_skip = False,171 fsdp = '',172 fsdp_min_num_params = 0,173 fsdp_config = None,174 tp_size = 0,175 fsdp_transformer_layer_cls_to_wrap = None,176 accelerator_config = None,177 deepspeed = None,178 label_smoothing_factor = 0.0,179 optim = 'adamw_8bit',180 optim_args = None,181 adafactor = False,182 group_by_length = False,183 length_column_name = 'length',184 report_to = None,185 ddp_find_unused_parameters = None,186 ddp_bucket_cap_mb = None,187 ddp_broadcast_buffers = None,188 dataloader_pin_memory = True,189 dataloader_persistent_workers = False,190 skip_memory_metrics = True,191 use_legacy_prediction_loop = False,192 push_to_hub = False,193 resume_from_checkpoint = None,194 hub_model_id = None,195 hub_strategy = 'every_save',196 hub_token = None,197 hub_private_repo = None,198 hub_always_push = False,199 gradient_checkpointing = False,200 gradient_checkpointing_kwargs = None,201 include_inputs_for_metrics = False,202 eval_do_concat_batches = True,203 fp16_backend = 'auto',204 evaluation_strategy = None,205 push_to_hub_model_id = None,206 push_to_hub_organization = None,207 push_to_hub_token = None,208 mp_parameters = '',209 auto_find_batch_size = False,210 full_determinism = False,211 torchdynamo = None,212 ray_scope = 'last',213 ddp_timeout = 1800,214 torch_compile = False,215 torch_compile_backend = None,216 torch_compile_mode = None,217 dispatch_batches = None,218 split_batches = None,219 include_tokens_per_second = False,220 include_num_input_tokens_seen = False,221 neftune_noise_alpha = None,222 optim_target_modules = None,223 batch_eval_metrics = False,224 eval_on_start = False,225 use_liger_kernel = False,226 eval_use_gather_object = False,227 average_tokens_across_devices = False,228 max_length = 1024,229 max_prompt_length = 512,230 max_completion_length = None,231 beta = 0.1,232 disable_dropout = True,233 label_pad_token_id = -100,234 padding_value = None,235 truncation_mode = 'keep_end',236 generate_during_eval = False,237 is_encoder_decoder = None,238 model_init_kwargs = None,239 dataset_num_proc = None,240 vllm_sampling_params = None,241 unsloth_num_chunks = -1,242 **kwargs,243 ):244 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!')245 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!')246 if output_dir is None and save_strategy == 'steps' and save_steps == 500:247 output_dir = 'unsloth_training_checkpoints'248 save_strategy = 'no'249 if dataset_num_proc is None:250 from multiprocessing import cpu_count251 dataset_num_proc = cpu_count()252 253 super().__init__(254 output_dir = output_dir,255 overwrite_output_dir = overwrite_output_dir,256 do_train = do_train,257 do_eval = do_eval,258 do_predict = do_predict,259 eval_strategy = eval_strategy,260 prediction_loss_only = prediction_loss_only,261 per_device_train_batch_size = per_device_train_batch_size,262 per_device_eval_batch_size = per_device_eval_batch_size,263 per_gpu_train_batch_size = per_gpu_train_batch_size,264 per_gpu_eval_batch_size = per_gpu_eval_batch_size,265 gradient_accumulation_steps = gradient_accumulation_steps,266 eval_accumulation_steps = eval_accumulation_steps,267 eval_delay = eval_delay,268 torch_empty_cache_steps = torch_empty_cache_steps,269 learning_rate = learning_rate,270 weight_decay = weight_decay,271 adam_beta1 = adam_beta1,272 adam_beta2 = adam_beta2,273 adam_epsilon = adam_epsilon,274 max_grad_norm = max_grad_norm,275 num_train_epochs = num_train_epochs,276 max_steps = max_steps,277 lr_scheduler_type = lr_scheduler_type,278 warmup_ratio = warmup_ratio,279 warmup_steps = warmup_steps,280 log_level = log_level,281 log_level_replica = log_level_replica,282 log_on_each_node = log_on_each_node,283 logging_dir = logging_dir,284 logging_strategy = logging_strategy,285 logging_first_step = logging_first_step,286 logging_steps = logging_steps,287 logging_nan_inf_filter = logging_nan_inf_filter,288 save_strategy = save_strategy,289 save_steps = save_steps,290 save_total_limit = save_total_limit,291 save_safetensors = save_safetensors,292 save_on_each_node = save_on_each_node,293 save_only_model = save_only_model,294 restore_callback_states_from_checkpoint = restore_callback_states_from_checkpoint,295 no_cuda = no_cuda,296 use_cpu = use_cpu,297 use_mps_device = use_mps_device,298 seed = seed,299 data_seed = data_seed,300 jit_mode_eval = jit_mode_eval,301 use_ipex = use_ipex,302 bf16 = bf16,303 fp16 = fp16,304 fp16_opt_level = fp16_opt_level,305 half_precision_backend = half_precision_backend,306 bf16_full_eval = bf16_full_eval,307 fp16_full_eval = fp16_full_eval,308 tf32 = tf32,309 local_rank = local_rank,310 ddp_backend = ddp_backend,311 tpu_num_cores = tpu_num_cores,312 tpu_metrics_debug = tpu_metrics_debug,313 debug = debug,314 dataloader_drop_last = dataloader_drop_last,315 eval_steps = eval_steps,316 dataloader_num_workers = dataloader_num_workers,317 dataloader_prefetch_factor = dataloader_prefetch_factor,318 past_index = past_index,319 run_name = run_name,320 disable_tqdm = disable_tqdm,321 remove_unused_columns = remove_unused_columns,322 label_names = label_names,323 load_best_model_at_end = load_best_model_at_end,324 metric_for_best_model = metric_for_best_model,325 greater_is_better = greater_is_better,326 ignore_data_skip = ignore_data_skip,327 fsdp = fsdp,328 fsdp_min_num_params = fsdp_min_num_params,329 fsdp_config = fsdp_config,330 tp_size = tp_size,331 fsdp_transformer_layer_cls_to_wrap = fsdp_transformer_layer_cls_to_wrap,332 accelerator_config = accelerator_config,333 deepspeed = deepspeed,334 label_smoothing_factor = label_smoothing_factor,335 optim = optim,336 optim_args = optim_args,337 adafactor = adafactor,338 group_by_length = group_by_length,339 length_column_name = length_column_name,340 report_to = report_to,341 ddp_find_unused_parameters = ddp_find_unused_parameters,342 ddp_bucket_cap_mb = ddp_bucket_cap_mb,343 ddp_broadcast_buffers = ddp_broadcast_buffers,344 dataloader_pin_memory = dataloader_pin_memory,345 dataloader_persistent_workers = dataloader_persistent_workers,346 skip_memory_metrics = skip_memory_metrics,347 use_legacy_prediction_loop = use_legacy_prediction_loop,348 push_to_hub = push_to_hub,349 resume_from_checkpoint = resume_from_checkpoint,350 hub_model_id = hub_model_id,351 hub_strategy = hub_strategy,352 hub_token = hub_token,353 hub_private_repo = hub_private_repo,354 hub_always_push = hub_always_push,355 gradient_checkpointing = gradient_checkpointing,356 gradient_checkpointing_kwargs = gradient_checkpointing_kwargs,357 include_inputs_for_metrics = include_inputs_for_metrics,358 eval_do_concat_batches = eval_do_concat_batches,359 fp16_backend = fp16_backend,360 evaluation_strategy = evaluation_strategy,361 push_to_hub_model_id = push_to_hub_model_id,362 push_to_hub_organization = push_to_hub_organization,363 push_to_hub_token = push_to_hub_token,364 mp_parameters = mp_parameters,365 auto_find_batch_size = auto_find_batch_size,366 full_determinism = full_determinism,367 torchdynamo = torchdynamo,368 ray_scope = ray_scope,369 ddp_timeout = ddp_timeout,370 torch_compile = torch_compile,371 torch_compile_backend = torch_compile_backend,372 torch_compile_mode = torch_compile_mode,373 dispatch_batches = dispatch_batches,374 split_batches = split_batches,375 include_tokens_per_second = include_tokens_per_second,376 include_num_input_tokens_seen = include_num_input_tokens_seen,377 neftune_noise_alpha = neftune_noise_alpha,378 optim_target_modules = optim_target_modules,379 batch_eval_metrics = batch_eval_metrics,380 eval_on_start = eval_on_start,381 use_liger_kernel = use_liger_kernel,382 eval_use_gather_object = eval_use_gather_object,383 average_tokens_across_devices = average_tokens_across_devices,384 max_length = max_length,385 max_prompt_length = max_prompt_length,386 max_completion_length = max_completion_length,387 beta = beta,388 disable_dropout = disable_dropout,389 label_pad_token_id = label_pad_token_id,390 padding_value = padding_value,391 truncation_mode = truncation_mode,392 generate_during_eval = generate_during_eval,393 is_encoder_decoder = is_encoder_decoder,394 model_init_kwargs = model_init_kwargs,395 dataset_num_proc = dataset_num_proc,**kwargs)396 self.vllm_sampling_params = vllm_sampling_params397 self.unsloth_num_chunks = unsloth_num_chunks398pass399 400class _UnslothORPOTrainer(Trainer):401 r""""""402 403 _tag_names = ["trl", "orpo"]404 405 def __init__(406 self,407 model: Optional[Union[PreTrainedModel, nn.Module, str]] = None,408 args: Optional[ORPOConfig] = None,409 data_collator: Optional[DataCollator] = None,410 train_dataset: Optional[Dataset] = None,411 eval_dataset: Optional[Union[Dataset, dict[str, Dataset]]] = None,412 processing_class: Optional[413 Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin]414 ] = None,415 model_init: Optional[Callable[[], PreTrainedModel]] = None,416 callbacks: Optional[list[TrainerCallback]] = None,417 optimizers: tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),418 preprocess_logits_for_metrics: Optional[Callable[[torch.Tensor, torch.Tensor], torch.Tensor]] = None,419 peft_config: Optional[dict] = None,420 compute_metrics: Optional[Callable[[EvalLoopOutput], dict]] = None,421 ):422 if args.model_init_kwargs is None:423 model_init_kwargs = {}424 elif not isinstance(model, str):425 raise ValueError("You passed model_kwargs to the ORPOTrainer. But your model is already instantiated.")426 else:427 model_init_kwargs = args.model_init_kwargs428 torch_dtype = model_init_kwargs.get("torch_dtype")429 if torch_dtype is not None:430 # Convert to `torch.dtype` if an str is passed431 if isinstance(torch_dtype, str) and torch_dtype != "auto":432 torch_dtype = getattr(torch, torch_dtype)433 if torch_dtype != "auto" and not isinstance(torch_dtype, torch.dtype):434 raise ValueError(435 f"Invalid `torch_dtype` passed to the ORPOConfig. Expected a string with either `torch.dtype` or 'auto', but got {torch_dtype}."436 )437 model_init_kwargs["torch_dtype"] = torch_dtype438 439 if isinstance(model, str):440 model = AutoModelForCausalLM.from_pretrained(model, **model_init_kwargs)441 442 # Initialize this variable to False. This helps tracking the case when `peft_module_casting_to_bf16`443 # has been called in order to properly call autocast if needed.444 self._peft_has_been_casted_to_bf16 = False445 446 if not is_peft_available() and peft_config is not None:447 raise ValueError(448 "PEFT is not installed and you passed a `peft_config` in the trainer's kwargs, please install it to use the PEFT models"449 )450 elif is_peft_available() and peft_config is not None:451 # if model is a peft model and we have a peft_config, we merge and unload it first452 if isinstance(model, PeftModel):453 model = model.merge_and_unload()454 455 if getattr(model, "is_loaded_in_8bit", False) or getattr(model, "is_loaded_in_4bit", False):456 _support_gc_kwargs = hasattr(457 args, "gradient_checkpointing_kwargs"458 ) and "gradient_checkpointing_kwargs" in list(459 inspect.signature(prepare_model_for_kbit_training).parameters460 )461 462 prepare_model_kwargs = {"use_gradient_checkpointing": args.gradient_checkpointing}463 464 if _support_gc_kwargs:465 prepare_model_kwargs["gradient_checkpointing_kwargs"] = args.gradient_checkpointing_kwargs466 467 model = prepare_model_for_kbit_training(model, **prepare_model_kwargs)468 elif getattr(args, "gradient_checkpointing", False):469 # For backward compatibility with older versions of transformers470 if hasattr(model, "enable_input_require_grads"):471 model.enable_input_require_grads()472 else:473 474 def make_inputs_require_grad(module, input, output):475 output.requires_grad_(True)476 477 model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)478 479 # get peft model with the given config480 model = model481 if args.bf16 and getattr(model, "is_loaded_in_4bit", False):482 peft_module_casting_to_bf16(model)483 # If args.bf16 we need to explicitly call `generate` with torch amp autocast context manager484 self._peft_has_been_casted_to_bf16 = True485 486 # For models that use gradient_checkpointing, we need to attach a hook that enables input487 # to explicitly have `requires_grad=True`, otherwise training will either silently488 # fail or completely fail.489 elif getattr(args, "gradient_checkpointing", False):490 # For backward compatibility with older versions of transformers491 if hasattr(model, "enable_input_require_grads"):492 model.enable_input_require_grads()493 else:494 495 def make_inputs_require_grad(module, input, output):496 output.requires_grad_(True)497 498 model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)499 500 if args.generate_during_eval and not (is_wandb_available() or is_comet_available()):501 raise ValueError(502 "`generate_during_eval=True` requires Weights and Biases or Comet to be installed."503 " Please install `wandb` or `comet-ml` to resolve."504 )505 506 if model is not None:507 self.is_encoder_decoder = model.config.is_encoder_decoder508 elif args.is_encoder_decoder is None:509 raise ValueError("When no model is provided, you need to pass the parameter is_encoder_decoder.")510 else:511 self.is_encoder_decoder = args.is_encoder_decoder512 513 if self.is_encoder_decoder:514 self.decoder_start_token_id = model.config.decoder_start_token_id515 self.pad_token_id = model.config.pad_token_id516 517 if processing_class is None:518 raise ValueError("processing_class must be specified to tokenize a ORPO dataset.")519 if args.max_length is None:520 warnings.warn(521 "`max_length` is not set in the ORPOConfig's init"522 " it will default to `512` by default, but you should do it yourself in the future.",523 UserWarning,524 )525 max_length = 512526 else:527 max_length = args.max_length528 if args.max_prompt_length is None:529 warnings.warn(530 "`max_prompt_length` is not set in the ORPOConfig's init"531 " it will default to `128` by default, but you should do it yourself in the future.",532 UserWarning,533 )534 max_prompt_length = 128535 else:536 max_prompt_length = args.max_prompt_length537 538 if args.max_completion_length is None and self.is_encoder_decoder:539 warnings.warn(540 "When using an encoder decoder architecture, you should set `max_completion_length` in the ORPOConfig's init"541 " it will default to `128` by default, but you should do it yourself in the future.",542 UserWarning,543 )544 self.max_completion_length = 128545 else:546 self.max_completion_length = args.max_completion_length547 548 if data_collator is None:549 data_collator = DPODataCollatorWithPadding(550 pad_token_id=processing_class.pad_token_id,551 label_pad_token_id=args.label_pad_token_id,552 is_encoder_decoder=self.is_encoder_decoder,553 )554 555 if args.remove_unused_columns:556 args.remove_unused_columns = False557 # warn users558 warnings.warn(559 "When using DPODataCollatorWithPadding, you should set `remove_unused_columns=False` in your TrainingArguments"560 " we have set it for you, but you should do it yourself in the future.",561 UserWarning,562 )563 564 self.use_dpo_data_collator = True565 else:566 self.use_dpo_data_collator = False567 568 # Disable dropout in the model and reference model569 if args.disable_dropout:570 disable_dropout_in_model(model)571 572 self.max_length = max_length573 self.generate_during_eval = args.generate_during_eval574 self.label_pad_token_id = args.label_pad_token_id575 self.padding_value = args.padding_value if args.padding_value is not None else processing_class.pad_token_id576 self.max_prompt_length = max_prompt_length577 self.truncation_mode = args.truncation_mode578 self.processing_class = processing_class579 580 self.beta = args.beta581 self.aux_loss_enabled = getattr(model.config, "output_router_logits", False)582 self.aux_loss_coef = getattr(model.config, "router_aux_loss_coef", 0.0)583 if self.aux_loss_enabled and self.aux_loss_coef == 0.0:584 warnings.warn(585 "You set `output_router_logits` to `True` in the model config, but `router_aux_loss_coef` is set to "586 "`0.0`, meaning the auxiliary loss will not be used. Either set `router_aux_loss_coef` to a value "587 "greater than `0.0`, or set `output_router_logits` to `False` if you don't want to use the auxiliary "588 "loss.",589 UserWarning,590 )591 592 self._stored_metrics = defaultdict(lambda: defaultdict(list))593 594 # The trainer estimates the number of FLOPs (floating-point operations) using the number of elements in the595 # input tensor associated with the key "input_ids". However, in ORPO, the sampled data does not include the596 # "input_ids" key. Instead, the available keys are "prompt_input_ids", "chosen_input_ids", and597 # "rejected_input_ids". As a result, the trainer issues the warning: "Could not estimate the number of tokens598 # of the input, floating-point operations will not be computed." To suppress this warning, we set the599 # "estimate_tokens" key in the model's "warnings_issued" dictionary to True. This acts as a flag to indicate600 # that the warning has already been issued.601 model.warnings_issued["estimate_tokens"] = True602 603 # Compute that only on the main process for faster data processing.604 # see: https://github.com/huggingface/trl/pull/1255605 with PartialState().local_main_process_first():606 # Extract the prompt if needed, and apply the chat template if needed607 train_dataset = train_dataset.map(maybe_extract_prompt, num_proc=args.dataset_num_proc)608 train_dataset = train_dataset.map(609 maybe_apply_chat_template, fn_kwargs={"tokenizer": processing_class}, num_proc=args.dataset_num_proc610 )611 train_dataset = train_dataset.map(self.tokenize_row, num_proc=args.dataset_num_proc)612 if eval_dataset is not None:613 eval_dataset = eval_dataset.map(maybe_extract_prompt, num_proc=args.dataset_num_proc)614 eval_dataset = eval_dataset.map(615 maybe_apply_chat_template,616 fn_kwargs={"tokenizer": processing_class},617 num_proc=args.dataset_num_proc,618 )619 eval_dataset = eval_dataset.map(self.tokenize_row, num_proc=args.dataset_num_proc)620 621 super().__init__(622 model=model,623 args=args,624 data_collator=data_collator,625 train_dataset=train_dataset,626 eval_dataset=eval_dataset,627 processing_class=processing_class,628 model_init=model_init,629 compute_metrics=compute_metrics,630 callbacks=callbacks,631 optimizers=optimizers,632 preprocess_logits_for_metrics=preprocess_logits_for_metrics,633 )634 635 # Add tags for models that have been loaded with the correct transformers version636 if hasattr(self.model, "add_model_tags"):637 self.model.add_model_tags(self._tag_names)638 639 if not hasattr(self, "accelerator"):640 raise AttributeError(641 "Your `Trainer` does not have an `accelerator` object. Consider upgrading `transformers`."642 )643 644 def _prepare_deepspeed(self, model: PreTrainedModelWrapper):645 # Adapted from accelerate: https://github.com/huggingface/accelerate/blob/739b135f8367becb67ffaada12fe76e3aa60fefd/src/accelerate/accelerator.py#L1473646 deepspeed_plugin = self.accelerator.state.deepspeed_plugin647 config_kwargs = deepcopy(deepspeed_plugin.deepspeed_config)648 649 if model is not None:650 if hasattr(model, "config"):651 hidden_size = (652 max(model.config.hidden_sizes)653 if getattr(model.config, "hidden_sizes", None)654 else getattr(model.config, "hidden_size", None)655 )656 if hidden_size is not None and config_kwargs["zero_optimization"]["stage"] == 3:657 # Note that `stage3_prefetch_bucket_size` can produce DeepSpeed messages like: `Invalidate trace cache @ step 0: expected module 1, but got module 0`658 # This is expected and is not an error, see: https://github.com/microsoft/DeepSpeed/discussions/4081659 config_kwargs.update(660 {661 "zero_optimization.reduce_bucket_size": hidden_size * hidden_size,662 "zero_optimization.stage3_param_persistence_threshold": 10 * hidden_size,663 "zero_optimization.stage3_prefetch_bucket_size": 0.9 * hidden_size * hidden_size,664 }665 )666 667 # If ZeRO-3 is used, we shard both the active and reference model.668 # Otherwise, we assume the reference model fits in memory and is initialized on each device with ZeRO disabled (stage 0)669 if config_kwargs["zero_optimization"]["stage"] != 3:670 config_kwargs["zero_optimization"]["stage"] = 0671 model, *_ = deepspeed.initialize(model=model, config=config_kwargs)672 model.eval()673 return model674 675 def build_tokenized_answer(self, prompt, answer):676 """677 Llama tokenizer does satisfy `enc(a + b) = enc(a) + enc(b)`.678 It does ensure `enc(a + b) = enc(a) + enc(a + b)[len(enc(a)):]`.679 Reference:680 https://github.com/EleutherAI/lm-evaluation-harness/pull/531#issuecomment-1595586257681 """682 683 full_tokenized = self.processing_class(prompt + answer, add_special_tokens=False)684 prompt_input_ids = self.processing_class(prompt, add_special_tokens=False)["input_ids"]685 686 answer_input_ids = full_tokenized["input_ids"][len(prompt_input_ids) :]687 answer_attention_mask = full_tokenized["attention_mask"][len(prompt_input_ids) :]688 689 # Concat tokens to form `enc(a) + enc(a + b)[len(enc(a)):]`690 full_concat_input_ids = np.concatenate([prompt_input_ids, answer_input_ids])691 692 # Prepare input tokens for token by token comparison693 full_input_ids = np.array(full_tokenized["input_ids"])694 695 if len(full_input_ids) != len(full_concat_input_ids):696 raise ValueError("Prompt input ids and answer input ids should have the same length.")697 698 # On some tokenizers, like Llama-2 tokenizer, there are occasions where tokens699 # can be merged together when tokenizing prompt+answer. This could result700 # on the last token from the prompt being different when tokenized on its own701 # vs when done as prompt+answer.702 response_token_ids_start_idx = len(prompt_input_ids)703 704 # If tokenized prompt is different than both prompt+answer, then it means the705 # last token has changed due to merging.706 if prompt_input_ids != full_tokenized["input_ids"][:response_token_ids_start_idx]:707 response_token_ids_start_idx -= 1708 709 prompt_input_ids = full_tokenized["input_ids"][:response_token_ids_start_idx]710 prompt_attention_mask = full_tokenized["attention_mask"][:response_token_ids_start_idx]711 712 if len(prompt_input_ids) != len(prompt_attention_mask):713 raise ValueError("Prompt input ids and attention mask should have the same length.")714 715 answer_input_ids = full_tokenized["input_ids"][response_token_ids_start_idx:]716 answer_attention_mask = full_tokenized["attention_mask"][response_token_ids_start_idx:]717 718 return dict(719 prompt_input_ids=prompt_input_ids,720 prompt_attention_mask=prompt_attention_mask,721 input_ids=answer_input_ids,722 attention_mask=answer_attention_mask,723 )724 725 def tokenize_row(self, feature, model: Optional[Union[PreTrainedModel, nn.Module]] = None) -> dict:726 """Tokenize a single row from a ORPO specific dataset.727 728 At this stage, we don't convert to PyTorch tensors yet; we just handle the truncation729 in case the prompt + chosen or prompt + rejected responses is/are too long. First730 we truncate the prompt; if we're still too long, we truncate the chosen/rejected.731 732 We also create the labels for the chosen/rejected responses, which are of length equal to733 the sum of the length of the prompt and the chosen/rejected response, with734 label_pad_token_id for the prompt tokens.735 """736 batch = {}737 prompt = feature["prompt"]738 chosen = feature["chosen"]739 rejected = feature["rejected"]740 741 if not self.is_encoder_decoder:742 # Check issues below for more details743 # 1. https://github.com/huggingface/trl/issues/907744 # 2. https://github.com/EleutherAI/lm-evaluation-harness/pull/531#issuecomment-1595586257745 # 3. https://github.com/LianjiaTech/BELLE/issues/337746 747 if not isinstance(prompt, str):748 raise ValueError(f"prompt should be an str but got {type(prompt)}")749 prompt_tokens = self.processing_class(prompt, add_special_tokens=False)750 prompt_tokens = {f"prompt_{k}": v for k, v in prompt_tokens.items()}751 752 if not isinstance(chosen, str):753 raise ValueError(f"chosen should be an str but got {type(chosen)}")754 chosen_tokens = self.build_tokenized_answer(prompt, chosen)755 756 if not isinstance(rejected, str):757 raise ValueError(f"rejected should be an str but got {type(rejected)}")758 rejected_tokens = self.build_tokenized_answer(prompt, rejected)759 760 # Last prompt token might get merged by tokenizer and761 # it should not be included for generation if that happens762 prompt_len_input_ids = len(prompt_tokens["prompt_input_ids"])763 764 chosen_prompt_len_input_ids = len(chosen_tokens["prompt_input_ids"])765 rejected_prompt_len_input_ids = len(rejected_tokens["prompt_input_ids"])766 prompt_len_input_ids = min(chosen_prompt_len_input_ids, rejected_prompt_len_input_ids)767 768 for k, v in prompt_tokens.items():769 prompt_tokens[k] = v[:prompt_len_input_ids]770 771 # Make sure prompts only have one different token at most an772 # and length only differs by 1 at most773 num_diff_tokens = sum(774 [a != b for a, b in zip(chosen_tokens["prompt_input_ids"], rejected_tokens["prompt_input_ids"])]775 )776 num_diff_len = abs(chosen_prompt_len_input_ids - rejected_prompt_len_input_ids)777 if num_diff_tokens > 1 or num_diff_len > 1:778 raise ValueError(779 "Chosen and rejected prompt_input_ids might only differ on the "780 "last token due to tokenizer merge ops."781 )782 783 # add BOS token to head of prompt. Avoid adding if it's already there784 prompt_tokens, chosen_tokens, rejected_tokens = add_bos_token_if_needed(785 self.processing_class.bos_token_id,786 prompt_len_input_ids,787 prompt_tokens,788 chosen_prompt_len_input_ids,789 chosen_tokens,790 rejected_prompt_len_input_ids,791 rejected_tokens,792 )793 794 # add EOS token to end of answer. Avoid adding if it's already there795 chosen_tokens, rejected_tokens = add_eos_token_if_needed(796 self.processing_class.eos_token_id, chosen_tokens, rejected_tokens797 )798 799 longer_response_length = max(len(chosen_tokens["input_ids"]), len(rejected_tokens["input_ids"]))800 801 # if combined sequence is too long, truncate the prompt802 for answer_tokens in [chosen_tokens, rejected_tokens, prompt_tokens]:803 if len(answer_tokens["prompt_input_ids"]) + longer_response_length > self.max_length:804 if self.truncation_mode == "keep_start":805 for k in ["prompt_input_ids", "prompt_attention_mask"]:806 answer_tokens[k] = answer_tokens[k][: self.max_prompt_length]807 elif self.truncation_mode == "keep_end":808 for k in ["prompt_input_ids", "prompt_attention_mask"]:809 answer_tokens[k] = answer_tokens[k][-self.max_prompt_length :]810 else:811 raise ValueError(f"Unknown truncation mode: {self.truncation_mode}")812 813 # if that's still too long, truncate the response814 for answer_tokens in [chosen_tokens, rejected_tokens]:815 if len(answer_tokens["prompt_input_ids"]) + longer_response_length > self.max_length:816 for k in ["input_ids", "attention_mask"]:817 answer_tokens[k] = answer_tokens[k][: self.max_length - self.max_prompt_length]818 819 # Create labels820 chosen_sequence_tokens = {821 k: chosen_tokens[f"prompt_{k}"] + chosen_tokens[k] for k in ["input_ids", "attention_mask"]822 }823 rejected_sequence_tokens = {824 k: rejected_tokens[f"prompt_{k}"] + rejected_tokens[k] for k in ["input_ids", "attention_mask"]825 }826 chosen_sequence_tokens["labels"] = chosen_sequence_tokens["input_ids"][:]827 chosen_sequence_tokens["labels"][: len(chosen_tokens["prompt_input_ids"])] = [828 self.label_pad_token_id829 ] * len(chosen_tokens["prompt_input_ids"])830 rejected_sequence_tokens["labels"] = rejected_sequence_tokens["input_ids"][:]831 rejected_sequence_tokens["labels"][: len(rejected_tokens["prompt_input_ids"])] = [832 self.label_pad_token_id833 ] * len(rejected_tokens["prompt_input_ids"])834 835 for k, toks in {836 "chosen_": chosen_sequence_tokens,837 "rejected_": rejected_sequence_tokens,838 "": prompt_tokens,839 }.items():840 for type_key, tokens in toks.items():841 if type_key == "token_type_ids":842 continue843 batch[f"{k}{type_key}"] = tokens844 845 else:846 chosen_tokens = self.processing_class(847 chosen, truncation=True, max_length=self.max_completion_length, add_special_tokens=True848 )849 rejected_tokens = self.processing_class(850 rejected, truncation=True, max_length=self.max_completion_length, add_special_tokens=True851 )852 prompt_tokens = self.processing_class(853 prompt, truncation=True, max_length=self.max_prompt_length, add_special_tokens=True854 )855 856 batch["chosen_labels"] = chosen_tokens["input_ids"]857 batch["rejected_labels"] = rejected_tokens["input_ids"]858 batch["prompt_input_ids"] = prompt_tokens["input_ids"]859 batch["prompt_attention_mask"] = prompt_tokens["attention_mask"]860 861 if model is not None and hasattr(model, "prepare_decoder_input_ids_from_labels"):862 batch["rejected_decoder_input_ids"] = model.prepare_decoder_input_ids_from_labels(863 labels=torch.tensor(batch["rejected_labels"])864 )865 batch["chosen_decoder_input_ids"] = model.prepare_decoder_input_ids_from_labels(866 labels=torch.tensor(batch["chosen_labels"])867 )868 869 if is_torch_xla_available():870 # Pad the sequences to global max_length to avoid TorchXLA recompilation871 for k in batch:872 if "labels" in k or self.is_encoder_decoder:873 pad_value = self.label_pad_token_id874 elif k.endswith("_input_ids"):875 pad_value = self.padding_value876 elif k.endswith("_attention_mask"):877 pad_value = 0878 batch[k] = batch[k] + [pad_value] * (self.max_length - len(batch[k]))879 return batch880 881 @staticmethod882 def concatenated_inputs(883 batch: dict[str, Union[list, torch.LongTensor]],884 is_encoder_decoder: bool = False,885 label_pad_token_id: int = -100,886 padding_value: int = 0,887 device: Optional[torch.device] = None,888 ) -> dict[str, torch.LongTensor]:889 """Concatenate the chosen and rejected inputs into a single tensor.890 891 Args:892 batch: A batch of data. Must contain the keys 'chosen_input_ids' and 'rejected_input_ids', which are tensors of shape (batch_size, sequence_length).893 is_encoder_decoder: Whether the model is an encoder-decoder model.894 label_pad_token_id: The label pad token id.895 padding_value: The padding value to use for the concatenated inputs_ids.896 device: The device for the concatenated inputs.897 898 Returns:899 A dictionary containing the concatenated inputs under the key 'concatenated_input_ids'.900 """901 concatenated_batch = {}902 903 if is_encoder_decoder:904 max_length = max(batch["chosen_labels"].shape[1], batch["rejected_labels"].shape[1])905 else:906 max_length = max(batch["chosen_input_ids"].shape[1], batch["rejected_input_ids"].shape[1])907 908 for k in batch:909 if k.startswith("chosen") and isinstance(batch[k], torch.Tensor):910 if "labels" in k or is_encoder_decoder:911 pad_value = label_pad_token_id912 elif k.endswith("_input_ids"):913 pad_value = padding_value914 elif k.endswith("_attention_mask"):915 pad_value = 0916 concatenated_key = k.replace("chosen", "concatenated")917 concatenated_batch[concatenated_key] = pad_to_length(batch[k], max_length, pad_value=pad_value)918 for k in batch:919 if k.startswith("rejected") and isinstance(batch[k], torch.Tensor):920 if "labels" in k or is_encoder_decoder:921 pad_value = label_pad_token_id922 elif k.endswith("_input_ids"):923 pad_value = padding_value924 elif k.endswith("_attention_mask"):925 pad_value = 0926 concatenated_key = k.replace("rejected", "concatenated")927 concatenated_batch[concatenated_key] = torch.cat(928 (929 concatenated_batch[concatenated_key],930 pad_to_length(batch[k], max_length, pad_value=pad_value),931 ),932 dim=0,933 ).to(device=device)934 935 if is_encoder_decoder:936 concatenated_batch["concatenated_input_ids"] = batch["prompt_input_ids"].repeat(2, 1).to(device=device)937 concatenated_batch["concatenated_attention_mask"] = (938 batch["prompt_attention_mask"].repeat(2, 1).to(device=device)939 )940 941 return concatenated_batch942 943 def odds_ratio_loss(944 self,945 policy_chosen_logps: torch.FloatTensor,946 policy_rejected_logps: torch.FloatTensor,947 ) -> tuple[torch.FloatTensor, torch.FloatTensor, torch.FloatTensor, torch.FloatTensor, torch.FloatTensor]:948 """Compute ORPO's odds ratio (OR) loss for a batch of policy and reference model log probabilities.949 950 Args:951 policy_chosen_logps: Log probabilities of the policy model for the chosen responses. Shape: (batch_size,)952 policy_rejected_logps: Log probabilities of the policy model for the rejected responses. Shape: (batch_size,)953 954 Returns:955 A tuple of three tensors: (losses, chosen_rewards, rejected_rewards).956 The losses tensor contains the ORPO loss for each example in the batch.957 The chosen_rewards and rejected_rewards tensors contain the rewards for the chosen and rejected responses, respectively.958 The log odds ratio of the chosen responses over the rejected responses ratio for logging purposes.959 The `log(sigmoid(log_odds_chosen))` for logging purposes.960 """961 962 # Derived from Eqs. (4) and (7) from https://huggingface.co/papers/2403.07691 by using log identities and exp(log(P(y|x)) = P(y|x)963 log_odds = (policy_chosen_logps - policy_rejected_logps) - (964 torch.log1p(-torch.exp(policy_chosen_logps)) - torch.log1p(-torch.exp(policy_rejected_logps))965 )966 ratio = F.logsigmoid(log_odds)967 losses = self.beta * ratio968 969 chosen_rewards = self.beta * (policy_chosen_logps.to(self.accelerator.device)).detach()970 rejected_rewards = self.beta * (policy_rejected_logps.to(self.accelerator.device)).detach()971 972 return losses, chosen_rewards, rejected_rewards, torch.mean(ratio), torch.mean(log_odds)973 974 @staticmethod975 def get_batch_logps(976 logits: torch.FloatTensor,977 labels: torch.LongTensor,978 average_log_prob: bool = False,979 label_pad_token_id: int = -100,980 is_encoder_decoder: bool = False,981 ) -> torch.FloatTensor:982 """Compute the log probabilities of the given labels under the given logits.983 984 Args:985 logits: Logits of the model (unnormalized). Shape: (batch_size, sequence_length, vocab_size)986 labels: Labels for which to compute the log probabilities. Label tokens with a value of label_pad_token_id are ignored. Shape: (batch_size, sequence_length)987 average_log_prob: If True, return the average log probability per (non-masked) token. Otherwise, return the sum of the log probabilities of the (non-masked) tokens.988 label_pad_token_id: The label pad token id.989 is_encoder_decoder: Whether the model is an encoder-decoder model.990 991 Returns:992 A tensor of shape (batch_size,) containing the average/sum log probabilities of the given labels under the given logits.993 """994 if logits.shape[:-1] != labels.shape:995 raise ValueError("Logits (batch and sequence length dim) and labels must have the same shape.")996 997 if not is_encoder_decoder:998 labels = labels[:, 1:].clone()999 logits = logits[:, :-1, :]1000 loss_mask = labels != label_pad_token_id1001 1002 # dummy token; we'll ignore the losses on these tokens later1003 labels = torch.where(labels == label_pad_token_id, 0, labels)1004 1005 per_token_logps = selective_log_softmax(logits, labels)1006 1007 if average_log_prob:1008 return (per_token_logps * loss_mask).sum(-1) / loss_mask.sum(-1)1009 else:1010 return (per_token_logps * loss_mask).sum(-1)1011 1012 def concatenated_forward(1013 self, model: nn.Module, batch: dict[str, Union[list, torch.LongTensor]]1014 ) -> tuple[torch.FloatTensor, torch.FloatTensor, torch.FloatTensor, torch.FloatTensor]:1015 """Run the given model on the given batch of inputs, concatenating the chosen and rejected inputs together.1016 1017 We do this to avoid doing two forward passes, because it's faster for FSDP.1018 """1019 concatenated_batch = self.concatenated_inputs(1020 batch,1021 is_encoder_decoder=self.is_encoder_decoder,1022 label_pad_token_id=self.label_pad_token_id,1023 padding_value=self.padding_value,1024 device=self.accelerator.device,1025 )1026 len_chosen = batch["chosen_labels"].shape[0]1027 1028 model_kwargs = (1029 {1030 "decoder_input_ids": self._shift_right(concatenated_batch["concatenated_labels"]),1031 }1032 if self.is_encoder_decoder1033 else {}1034 )1035 1036 if self.aux_loss_enabled:1037 model_kwargs["output_router_logits"] = True1038 1039 outputs = model(1040 concatenated_batch["concatenated_input_ids"],1041 attention_mask=concatenated_batch["concatenated_attention_mask"],1042 use_cache=False,1043 **model_kwargs,1044 )1045 all_logits = outputs.logits1046 1047 def cross_entropy_loss(logits, labels):1048 if not self.is_encoder_decoder:1049 # Shift so that tokens < n predict n1050 logits = logits[..., :-1, :].contiguous()1051 labels = labels[..., 1:].contiguous()1052 # Flatten the tokens1053 loss_fct = nn.CrossEntropyLoss()1054 logits = logits.view(-1, logits.shape[-1])1055 labels = labels.view(-1)1056 # Enable model parallelism1057 labels = labels.to(logits.device)1058 loss = loss_fct(logits, labels)1059 return loss1060 1061 if self.is_encoder_decoder:1062 labels = concatenated_batch["concatenated_labels"].clone()1063 else:1064 labels = concatenated_batch["concatenated_input_ids"].clone()1065 attention_mask = concatenated_batch["concatenated_attention_mask"]1066 labels = torch.where(attention_mask == 1, labels, self.label_pad_token_id)1067 # orpo chosen nll loss is computed over the full prompt and response1068 chosen_nll_loss = cross_entropy_loss(all_logits[:len_chosen], labels[:len_chosen])1069 1070 all_logps = self.get_batch_logps(1071 all_logits,1072 concatenated_batch["concatenated_labels"],1073 average_log_prob=True,1074 is_encoder_decoder=self.is_encoder_decoder,1075 label_pad_token_id=self.label_pad_token_id,1076 )1077 1078 chosen_logps = all_logps[:len_chosen]1079 rejected_logps = all_logps[len_chosen:]1080 1081 if not self.is_encoder_decoder:1082 chosen_logits = all_logits[:len_chosen, :-1, :]1083 rejected_logits = all_logits[len_chosen:, :-1, :]1084 else:1085 chosen_logits = all_logits[:len_chosen]1086 rejected_logits = all_logits[len_chosen:]1087 1088 if self.aux_loss_enabled:1089 return (chosen_logps, rejected_logps, chosen_logits, rejected_logits, chosen_nll_loss, outputs.aux_loss)1090 1091 return (chosen_logps, rejected_logps, chosen_logits, rejected_logits, chosen_nll_loss)1092 1093 def get_batch_loss_metrics(1094 self,1095 model,1096 batch: dict[str, Union[list, torch.LongTensor]],1097 train_eval: Literal["train", "eval"] = "train",1098 ):1099 """Compute the ORPO loss and other metrics for the given batch of inputs for train or test."""1100 metrics = {}1101 1102 forward_output = self.concatenated_forward(model, batch)1103 (1104 policy_chosen_logps,1105 policy_rejected_logps,1106 policy_chosen_logits,1107 policy_rejected_logits,1108 policy_nll_loss,1109 ) = forward_output[:5]1110 if self.aux_loss_enabled:1111 aux_loss = forward_output[5]1112 1113 losses, chosen_rewards, rejected_rewards, log_odds_ratio, log_odds_chosen = self.odds_ratio_loss(1114 policy_chosen_logps, policy_rejected_logps1115 )1116 # full ORPO loss1117 loss = policy_nll_loss - losses.mean()1118 1119 reward_accuracies = (chosen_rewards > rejected_rewards).float()1120 1121 prefix = "eval_" if train_eval == "eval" else ""1122 metrics[f"{prefix}rewards/chosen"] = self.accelerator.gather_for_metrics(chosen_rewards).mean()1123 metrics[f"{prefix}rewards/rejected"] = self.accelerator.gather_for_metrics(rejected_rewards).mean()1124 metrics[f"{prefix}rewards/accuracies"] = self.accelerator.gather_for_metrics(reward_accuracies).mean()1125 metrics[f"{prefix}rewards/margins"] = self.accelerator.gather_for_metrics(1126 chosen_rewards - rejected_rewards1127 ).mean()1128 metrics[f"{prefix}logps/rejected"] = self.accelerator.gather_for_metrics(policy_rejected_logps).detach().mean()1129 metrics[f"{prefix}logps/chosen"] = self.accelerator.gather_for_metrics(policy_chosen_logps).detach().mean()1130 metrics[f"{prefix}logits/rejected"] = (1131 self.accelerator.gather_for_metrics(policy_rejected_logits).detach().mean()1132 )1133 metrics[f"{prefix}logits/chosen"] = self.accelerator.gather_for_metrics(policy_chosen_logits).detach().mean()1134 metrics[f"{prefix}nll_loss"] = self.accelerator.gather_for_metrics(policy_nll_loss).detach().mean()1135 metrics[f"{prefix}log_odds_ratio"] = self.accelerator.gather_for_metrics(log_odds_ratio).mean()1136 metrics[f"{prefix}log_odds_chosen"] = self.accelerator.gather_for_metrics(log_odds_chosen).mean()1137 if is_torch_xla_available():1138 xm.mark_step() # needed because .item() calls1139 for k, v in metrics.items():1140 metrics[k] = v.item()1141 if self.aux_loss_enabled:1142 loss += self.aux_loss_coef * aux_loss1143 1144 return loss, metrics1145 1146 def compute_loss(1147 self,1148 model: Union[PreTrainedModel, nn.Module],1149 inputs: dict[str, Union[torch.Tensor, Any]],1150 return_outputs=False,1151 num_items_in_batch=None,1152 ) -> Union[torch.Tensor, tuple[torch.Tensor, dict[str, torch.Tensor]]]:1153 compute_loss_context_manager = amp.autocast("cuda") if self._peft_has_been_casted_to_bf16 else nullcontext()1154 1155 with compute_loss_context_manager:1156 loss, metrics = self.get_batch_loss_metrics(model, inputs, train_eval="train")1157 1158 # Make sure to move the loss to the device the original accumulating loss is at back in the `Trainer` class:1159 loss = loss.to(self.args.device)1160 1161 # force log the metrics1162 self.store_metrics(metrics, train_eval="train")1163 1164 if return_outputs:1165 return (loss, metrics)1166 return loss1167 1168 def generate_from_model(self, model, batch: dict[str, torch.LongTensor]) -> str:1169 """Generate samples from the model and reference model for the given batch of inputs."""1170 1171 # If one uses `generate_during_eval` with peft + bf16, we need to explicitly call generate with1172 # the torch cuda amp context manager as some hidden states are silently casted to full precision.1173 generate_context_manager = amp.autocast("cuda") if self._peft_has_been_casted_to_bf16 else nullcontext()1174 1175 with generate_context_manager:1176 policy_output = model.generate(1177 input_ids=batch["prompt_input_ids"],1178 attention_mask=batch["prompt_attention_mask"],1179 max_length=self.max_length,1180 do_sample=True,1181 pad_token_id=self.processing_class.pad_token_id,1182 )1183 1184 policy_output = pad_to_length(policy_output, self.max_length, self.processing_class.pad_token_id)1185 policy_output_decoded = self.processing_class.batch_decode(policy_output, skip_special_tokens=True)1186 1187 return policy_output_decoded1188 1189 def prediction_step(1190 self,1191 model: Union[PreTrainedModel, nn.Module],1192 inputs: dict[str, Union[torch.Tensor, Any]],1193 prediction_loss_only: bool,1194 ignore_keys: Optional[list[str]] = None,1195 ):1196 if not self.use_dpo_data_collator:1197 warnings.warn(1198 "prediction_step is only implemented for DPODataCollatorWithPadding, and you passed a datacollator that is different than "1199 "DPODataCollatorWithPadding - you might see unexpected behavior. Alternatively, you can implement your own prediction_step method if you are using a custom data collator"1200 )