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.cpo_trainer import (Any, AutoModelForCausalLM, BaseImageProcessor, CPOConfig, CPOTrainer, Callable, DPODataCollatorWithPadding, DataCollator, DataLoader, Dataset, EvalLoopOutput, F, FeatureExtractionMixin, Literal, Optional, PartialState, PeftModel, PreTrainedModel, PreTrainedTokenizerBase, ProcessorMixin, Trainer, TrainerCallback, Union, add_bos_token_if_needed, add_eos_token_if_needed, amp, defaultdict, disable_dropout_in_model, generate_model_card, get_comet_experiment_url, inspect, is_comet_available, is_peft_available, is_torch_fx_proxy, 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 UnslothCPOConfig(CPOConfig):44 """45 46 Configuration class for the [`CPOTrainer`].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 deviation from the reference model. Higher β means less deviation from the66 reference model. For the IPO loss (`loss_type="ipo"`), β is the regularization parameter denoted by τ in67 the [paper](https://huggingface.co/papers/2310.12036).68 label_smoothing (`float`, *optional*, defaults to `0.0`):69 Label smoothing factor. This argument is required if you want to use the default data collator.70 loss_type (`str`, *optional*, defaults to `"sigmoid"`):71 Type of loss to use. Possible values are:72 73 - `"sigmoid"`: sigmoid loss from the original [DPO](https://huggingface.co/papers/2305.18290) paper.74 - `"hinge"`: hinge loss on the normalized likelihood from the [SLiC](https://huggingface.co/papers/2305.10425) paper.75 - `"ipo"`: IPO loss from the [IPO](https://huggingface.co/papers/2310.12036) paper.76 - `"simpo"`: SimPO loss from the [SimPO](https://huggingface.co/papers/2405.14734) paper.77 78 disable_dropout (`bool`, *optional*, defaults to `True`):79 Whether to disable dropout in the model.80 cpo_alpha (`float`, *optional*, defaults to `1.0`):81 Weight of the BC regularizer in CPO training.82 simpo_gamma (`float`, *optional*, defaults to `0.5`):83 Target reward margin for the SimPO loss, used only when the `loss_type="simpo"`.84 label_pad_token_id (`int`, *optional*, defaults to `-100`):85 Label pad token id. This argument is required if you want to use the default data collator.86 padding_value (`int` or `None`, *optional*, defaults to `None`):87 Padding value to use. If `None`, the padding value of the tokenizer is used.88 truncation_mode (`str`,*optional*, defaults to `"keep_end"`):89 Truncation mode to use when the prompt is too long. Possible values are `"keep_end"` or `"keep_start"`.90 This argument is required if you want to use the default data collator.91 generate_during_eval (`bool`, *optional*, defaults to `False`):92 If `True`, generates and logs completions from the model to W&B or Comet during evaluation.93 is_encoder_decoder (`bool` or `None`, *optional*, defaults to `None`):94 When using the `model_init` argument (callable) to instantiate the model instead of the `model` argument,95 you need to specify if the model returned by the callable is an encoder-decoder model.96 model_init_kwargs (`dict[str, Any]` or `None`, *optional*, defaults to `None`):97 Keyword arguments to pass to `AutoModelForCausalLM.from_pretrained` when instantiating the model from a98 string.99 dataset_num_proc (`int` or `None`, *optional*, defaults to `None`):100 Number of processes to use for processing the dataset.101 102 """103 vllm_sampling_params: Optional[Any] = field(104 default = None,105 metadata = {'help': 'vLLM SamplingParams'},106 )107 unsloth_num_chunks : Optional[int] = field(108 default = -1,109 metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},110 )111 def __init__(112 self,113 output_dir = None,114 overwrite_output_dir = None,115 do_train = False,116 do_eval = False,117 do_predict = False,118 eval_strategy = 'no',119 prediction_loss_only = False,120 per_device_train_batch_size = 4,121 per_device_eval_batch_size = 4,122 per_gpu_train_batch_size = None,123 per_gpu_eval_batch_size = None,124 gradient_accumulation_steps = 2,125 eval_accumulation_steps = 2,126 eval_delay = 0,127 torch_empty_cache_steps = 250,128 learning_rate = 5e-05,129 weight_decay = 0.01,130 adam_beta1 = 0.9,131 adam_beta2 = 0.999,132 adam_epsilon = 1e-08,133 max_grad_norm = 1.0,134 num_train_epochs = 3.0,135 max_steps = -1,136 lr_scheduler_type = 'linear',137 warmup_ratio = 0.1,138 warmup_steps = 0,139 log_level = 'passive',140 log_level_replica = 'warning',141 log_on_each_node = True,142 logging_dir = None,143 logging_strategy = 'steps',144 logging_first_step = False,145 logging_steps = 1,146 logging_nan_inf_filter = False,147 save_strategy = 'steps',148 save_steps = 500,149 save_total_limit = None,150 save_safetensors = True,151 save_on_each_node = False,152 save_only_model = False,153 restore_callback_states_from_checkpoint = False,154 no_cuda = False,155 use_cpu = False,156 use_mps_device = False,157 seed = 3407,158 data_seed = 3407,159 jit_mode_eval = False,160 use_ipex = False,161 bf16 = False,162 fp16 = False,163 fp16_opt_level = 'O1',164 half_precision_backend = 'auto',165 bf16_full_eval = False,166 fp16_full_eval = False,167 tf32 = None,168 local_rank = -1,169 ddp_backend = None,170 tpu_num_cores = None,171 tpu_metrics_debug = False,172 debug = '',173 dataloader_drop_last = False,174 eval_steps = None,175 dataloader_num_workers = 0,176 dataloader_prefetch_factor = None,177 past_index = -1,178 run_name = None,179 disable_tqdm = None,180 remove_unused_columns = True,181 label_names = None,182 load_best_model_at_end = False,183 metric_for_best_model = None,184 greater_is_better = None,185 ignore_data_skip = False,186 fsdp = '',187 fsdp_min_num_params = 0,188 fsdp_config = None,189 tp_size = 0,190 fsdp_transformer_layer_cls_to_wrap = None,191 accelerator_config = None,192 deepspeed = None,193 label_smoothing_factor = 0.0,194 optim = 'adamw_8bit',195 optim_args = None,196 adafactor = False,197 group_by_length = False,198 length_column_name = 'length',199 report_to = None,200 ddp_find_unused_parameters = None,201 ddp_bucket_cap_mb = None,202 ddp_broadcast_buffers = None,203 dataloader_pin_memory = True,204 dataloader_persistent_workers = False,205 skip_memory_metrics = True,206 use_legacy_prediction_loop = False,207 push_to_hub = False,208 resume_from_checkpoint = None,209 hub_model_id = None,210 hub_strategy = 'every_save',211 hub_token = None,212 hub_private_repo = None,213 hub_always_push = False,214 gradient_checkpointing = False,215 gradient_checkpointing_kwargs = None,216 include_inputs_for_metrics = False,217 eval_do_concat_batches = True,218 fp16_backend = 'auto',219 evaluation_strategy = None,220 push_to_hub_model_id = None,221 push_to_hub_organization = None,222 push_to_hub_token = None,223 mp_parameters = '',224 auto_find_batch_size = False,225 full_determinism = False,226 torchdynamo = None,227 ray_scope = 'last',228 ddp_timeout = 1800,229 torch_compile = False,230 torch_compile_backend = None,231 torch_compile_mode = None,232 dispatch_batches = None,233 split_batches = None,234 include_tokens_per_second = False,235 include_num_input_tokens_seen = False,236 neftune_noise_alpha = None,237 optim_target_modules = None,238 batch_eval_metrics = False,239 eval_on_start = False,240 use_liger_kernel = False,241 eval_use_gather_object = False,242 average_tokens_across_devices = False,243 max_length = 1024,244 max_prompt_length = 512,245 max_completion_length = None,246 beta = 0.1,247 label_smoothing = 0.0,248 loss_type = 'sigmoid',249 disable_dropout = True,250 cpo_alpha = 1.0,251 simpo_gamma = 0.5,252 label_pad_token_id = -100,253 padding_value = None,254 truncation_mode = 'keep_end',255 generate_during_eval = False,256 is_encoder_decoder = None,257 model_init_kwargs = None,258 dataset_num_proc = None,259 vllm_sampling_params = None,260 unsloth_num_chunks = -1,261 **kwargs,262 ):263 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!')264 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!')265 if output_dir is None and save_strategy == 'steps' and save_steps == 500:266 output_dir = 'unsloth_training_checkpoints'267 save_strategy = 'no'268 if dataset_num_proc is None:269 from multiprocessing import cpu_count270 dataset_num_proc = cpu_count()271 272 super().__init__(273 output_dir = output_dir,274 overwrite_output_dir = overwrite_output_dir,275 do_train = do_train,276 do_eval = do_eval,277 do_predict = do_predict,278 eval_strategy = eval_strategy,279 prediction_loss_only = prediction_loss_only,280 per_device_train_batch_size = per_device_train_batch_size,281 per_device_eval_batch_size = per_device_eval_batch_size,282 per_gpu_train_batch_size = per_gpu_train_batch_size,283 per_gpu_eval_batch_size = per_gpu_eval_batch_size,284 gradient_accumulation_steps = gradient_accumulation_steps,285 eval_accumulation_steps = eval_accumulation_steps,286 eval_delay = eval_delay,287 torch_empty_cache_steps = torch_empty_cache_steps,288 learning_rate = learning_rate,289 weight_decay = weight_decay,290 adam_beta1 = adam_beta1,291 adam_beta2 = adam_beta2,292 adam_epsilon = adam_epsilon,293 max_grad_norm = max_grad_norm,294 num_train_epochs = num_train_epochs,295 max_steps = max_steps,296 lr_scheduler_type = lr_scheduler_type,297 warmup_ratio = warmup_ratio,298 warmup_steps = warmup_steps,299 log_level = log_level,300 log_level_replica = log_level_replica,301 log_on_each_node = log_on_each_node,302 logging_dir = logging_dir,303 logging_strategy = logging_strategy,304 logging_first_step = logging_first_step,305 logging_steps = logging_steps,306 logging_nan_inf_filter = logging_nan_inf_filter,307 save_strategy = save_strategy,308 save_steps = save_steps,309 save_total_limit = save_total_limit,310 save_safetensors = save_safetensors,311 save_on_each_node = save_on_each_node,312 save_only_model = save_only_model,313 restore_callback_states_from_checkpoint = restore_callback_states_from_checkpoint,314 no_cuda = no_cuda,315 use_cpu = use_cpu,316 use_mps_device = use_mps_device,317 seed = seed,318 data_seed = data_seed,319 jit_mode_eval = jit_mode_eval,320 use_ipex = use_ipex,321 bf16 = bf16,322 fp16 = fp16,323 fp16_opt_level = fp16_opt_level,324 half_precision_backend = half_precision_backend,325 bf16_full_eval = bf16_full_eval,326 fp16_full_eval = fp16_full_eval,327 tf32 = tf32,328 local_rank = local_rank,329 ddp_backend = ddp_backend,330 tpu_num_cores = tpu_num_cores,331 tpu_metrics_debug = tpu_metrics_debug,332 debug = debug,333 dataloader_drop_last = dataloader_drop_last,334 eval_steps = eval_steps,335 dataloader_num_workers = dataloader_num_workers,336 dataloader_prefetch_factor = dataloader_prefetch_factor,337 past_index = past_index,338 run_name = run_name,339 disable_tqdm = disable_tqdm,340 remove_unused_columns = remove_unused_columns,341 label_names = label_names,342 load_best_model_at_end = load_best_model_at_end,343 metric_for_best_model = metric_for_best_model,344 greater_is_better = greater_is_better,345 ignore_data_skip = ignore_data_skip,346 fsdp = fsdp,347 fsdp_min_num_params = fsdp_min_num_params,348 fsdp_config = fsdp_config,349 tp_size = tp_size,350 fsdp_transformer_layer_cls_to_wrap = fsdp_transformer_layer_cls_to_wrap,351 accelerator_config = accelerator_config,352 deepspeed = deepspeed,353 label_smoothing_factor = label_smoothing_factor,354 optim = optim,355 optim_args = optim_args,356 adafactor = adafactor,357 group_by_length = group_by_length,358 length_column_name = length_column_name,359 report_to = report_to,360 ddp_find_unused_parameters = ddp_find_unused_parameters,361 ddp_bucket_cap_mb = ddp_bucket_cap_mb,362 ddp_broadcast_buffers = ddp_broadcast_buffers,363 dataloader_pin_memory = dataloader_pin_memory,364 dataloader_persistent_workers = dataloader_persistent_workers,365 skip_memory_metrics = skip_memory_metrics,366 use_legacy_prediction_loop = use_legacy_prediction_loop,367 push_to_hub = push_to_hub,368 resume_from_checkpoint = resume_from_checkpoint,369 hub_model_id = hub_model_id,370 hub_strategy = hub_strategy,371 hub_token = hub_token,372 hub_private_repo = hub_private_repo,373 hub_always_push = hub_always_push,374 gradient_checkpointing = gradient_checkpointing,375 gradient_checkpointing_kwargs = gradient_checkpointing_kwargs,376 include_inputs_for_metrics = include_inputs_for_metrics,377 eval_do_concat_batches = eval_do_concat_batches,378 fp16_backend = fp16_backend,379 evaluation_strategy = evaluation_strategy,380 push_to_hub_model_id = push_to_hub_model_id,381 push_to_hub_organization = push_to_hub_organization,382 push_to_hub_token = push_to_hub_token,383 mp_parameters = mp_parameters,384 auto_find_batch_size = auto_find_batch_size,385 full_determinism = full_determinism,386 torchdynamo = torchdynamo,387 ray_scope = ray_scope,388 ddp_timeout = ddp_timeout,389 torch_compile = torch_compile,390 torch_compile_backend = torch_compile_backend,391 torch_compile_mode = torch_compile_mode,392 dispatch_batches = dispatch_batches,393 split_batches = split_batches,394 include_tokens_per_second = include_tokens_per_second,395 include_num_input_tokens_seen = include_num_input_tokens_seen,396 neftune_noise_alpha = neftune_noise_alpha,397 optim_target_modules = optim_target_modules,398 batch_eval_metrics = batch_eval_metrics,399 eval_on_start = eval_on_start,400 use_liger_kernel = use_liger_kernel,401 eval_use_gather_object = eval_use_gather_object,402 average_tokens_across_devices = average_tokens_across_devices,403 max_length = max_length,404 max_prompt_length = max_prompt_length,405 max_completion_length = max_completion_length,406 beta = beta,407 label_smoothing = label_smoothing,408 loss_type = loss_type,409 disable_dropout = disable_dropout,410 cpo_alpha = cpo_alpha,411 simpo_gamma = simpo_gamma,412 label_pad_token_id = label_pad_token_id,413 padding_value = padding_value,414 truncation_mode = truncation_mode,415 generate_during_eval = generate_during_eval,416 is_encoder_decoder = is_encoder_decoder,417 model_init_kwargs = model_init_kwargs,418 dataset_num_proc = dataset_num_proc,**kwargs)419 self.vllm_sampling_params = vllm_sampling_params420 self.unsloth_num_chunks = unsloth_num_chunks421pass422 423class _UnslothCPOTrainer(Trainer):424 r""""""425 426 _tag_names = ["trl", "cpo"]427 428 def __init__(429 self,430 model: Optional[Union[PreTrainedModel, nn.Module, str]] = None,431 args: Optional[CPOConfig] = None,432 data_collator: Optional[DataCollator] = None,433 train_dataset: Optional[Dataset] = None,434 eval_dataset: Optional[Union[Dataset, dict[str, Dataset]]] = None,435 processing_class: Optional[436 Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin]437 ] = None,438 model_init: Optional[Callable[[], PreTrainedModel]] = None,439 callbacks: Optional[list[TrainerCallback]] = None,440 optimizers: tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),441 preprocess_logits_for_metrics: Optional[Callable[[torch.Tensor, torch.Tensor], torch.Tensor]] = None,442 peft_config: Optional[dict] = None,443 compute_metrics: Optional[Callable[[EvalLoopOutput], dict]] = None,444 ):445 if args.model_init_kwargs is None:446 model_init_kwargs = {}447 elif not isinstance(model, str):448 raise ValueError("You passed model_kwargs to the CPOTrainer. But your model is already instantiated.")449 else:450 model_init_kwargs = args.model_init_kwargs451 torch_dtype = model_init_kwargs.get("torch_dtype")452 if torch_dtype is not None:453 # Convert to `torch.dtype` if an str is passed454 if isinstance(torch_dtype, str) and torch_dtype != "auto":455 torch_dtype = getattr(torch, torch_dtype)456 if torch_dtype != "auto" and not isinstance(torch_dtype, torch.dtype):457 raise ValueError(458 f"Invalid `torch_dtype` passed to the CPOConfig. Expected a string with either `torch.dtype` or 'auto', but got {torch_dtype}."459 )460 model_init_kwargs["torch_dtype"] = torch_dtype461 462 if isinstance(model, str):463 model = AutoModelForCausalLM.from_pretrained(model, **model_init_kwargs)464 465 # Initialize this variable to False. This helps tracking the case when `peft_module_casting_to_bf16`466 # has been called in order to properly call autocast if needed.467 self._peft_has_been_casted_to_bf16 = False468 469 if not is_peft_available() and peft_config is not None:470 raise ValueError(471 "PEFT is not installed and you passed a `peft_config` in the trainer's kwargs, please install it to use the PEFT models"472 )473 elif is_peft_available() and peft_config is not None:474 # if model is a peft model and we have a peft_config, we merge and unload it first475 if isinstance(model, PeftModel):476 model = model.merge_and_unload()477 478 if getattr(model, "is_loaded_in_8bit", False) or getattr(model, "is_loaded_in_4bit", False):479 _support_gc_kwargs = hasattr(480 args, "gradient_checkpointing_kwargs"481 ) and "gradient_checkpointing_kwargs" in list(482 inspect.signature(prepare_model_for_kbit_training).parameters483 )484 485 prepare_model_kwargs = {"use_gradient_checkpointing": args.gradient_checkpointing}486 487 if _support_gc_kwargs:488 prepare_model_kwargs["gradient_checkpointing_kwargs"] = args.gradient_checkpointing_kwargs489 490 model = prepare_model_for_kbit_training(model, **prepare_model_kwargs)491 elif getattr(args, "gradient_checkpointing", False):492 # For backward compatibility with older versions of transformers493 if hasattr(model, "enable_input_require_grads"):494 model.enable_input_require_grads()495 else:496 497 def make_inputs_require_grad(module, input, output):498 output.requires_grad_(True)499 500 model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)501 502 # get peft model with the given config503 model = model504 if args.bf16 and getattr(model, "is_loaded_in_4bit", False):505 peft_module_casting_to_bf16(model)506 # If args.bf16 we need to explicitly call `generate` with torch amp autocast context manager507 self._peft_has_been_casted_to_bf16 = True508 509 # For models that use gradient_checkpointing, we need to attach a hook that enables input510 # to explicitly have `requires_grad=True`, otherwise training will either silently511 # fail or completely fail.512 elif getattr(args, "gradient_checkpointing", False):513 # For backward compatibility with older versions of transformers514 if hasattr(model, "enable_input_require_grads"):515 model.enable_input_require_grads()516 else:517 518 def make_inputs_require_grad(module, input, output):519 output.requires_grad_(True)520 521 model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)522 523 if args.generate_during_eval and not (is_wandb_available() or is_comet_available()):524 raise ValueError(525 "`generate_during_eval=True` requires Weights and Biases or Comet to be installed."526 " Please install `wandb` or `comet-ml` to resolve."527 )528 529 if model is not None:530 self.is_encoder_decoder = model.config.is_encoder_decoder531 elif args.is_encoder_decoder is None:532 raise ValueError("When no model is provided, you need to pass the parameter is_encoder_decoder.")533 else:534 self.is_encoder_decoder = args.is_encoder_decoder535 536 if self.is_encoder_decoder:537 self.decoder_start_token_id = model.config.decoder_start_token_id538 self.pad_token_id = model.config.pad_token_id539 540 if processing_class is None:541 raise ValueError("processing_class must be specified to tokenize a CPO dataset.")542 if args.max_length is None:543 warnings.warn(544 "`max_length` is not set in the CPOConfig's init"545 " it will default to `512` by default, but you should do it yourself in the future.",546 UserWarning,547 )548 max_length = 512549 else:550 max_length = args.max_length551 if args.max_prompt_length is None:552 warnings.warn(553 "`max_prompt_length` is not set in the CPOConfig's init"554 " it will default to `128` by default, but you should do it yourself in the future.",555 UserWarning,556 )557 max_prompt_length = 128558 else:559 max_prompt_length = args.max_prompt_length560 561 if args.max_completion_length is None and self.is_encoder_decoder:562 warnings.warn(563 "When using an encoder decoder architecture, you should set `max_completion_length` in the CPOConfig's init"564 " it will default to `128` by default, but you should do it yourself in the future.",565 UserWarning,566 )567 max_completion_length = 128568 else:569 max_completion_length = args.max_completion_length570 571 if data_collator is None:572 data_collator = DPODataCollatorWithPadding(573 pad_token_id=processing_class.pad_token_id,574 label_pad_token_id=args.label_pad_token_id,575 is_encoder_decoder=self.is_encoder_decoder,576 )577 578 if args.remove_unused_columns:579 args.remove_unused_columns = False580 # warn users581 warnings.warn(582 "When using DPODataCollatorWithPadding, you should set `remove_unused_columns=False` in your TrainingArguments"583 " we have set it for you, but you should do it yourself in the future.",584 UserWarning,585 )586 587 self.use_dpo_data_collator = True588 else:589 self.use_dpo_data_collator = False590 591 # Disable dropout in the model592 if args.disable_dropout:593 disable_dropout_in_model(model)594 595 self.max_length = max_length596 self.generate_during_eval = args.generate_during_eval597 self.label_pad_token_id = args.label_pad_token_id598 self.padding_value = args.padding_value if args.padding_value is not None else processing_class.pad_token_id599 self.max_prompt_length = max_prompt_length600 self.truncation_mode = args.truncation_mode601 self.max_completion_length = max_completion_length602 self.processing_class = processing_class603 604 if args.loss_type in ["hinge", "ipo"] and args.label_smoothing > 0:605 warnings.warn(606 f"You are using the {args.loss_type} loss type that does not support label smoothing. The "607 "`label_smoothing` parameter will be ignored. Set `label_smoothing` to `0.0` to remove this warning.",608 UserWarning,609 )610 if args.loss_type == "kto_pair":611 raise ValueError("Support for kto_pair has been removed in CPOTrainer. Please use KTOTrainer.")612 613 self.beta = args.beta614 self.label_smoothing = args.label_smoothing615 self.loss_type = args.loss_type616 self.cpo_alpha = args.cpo_alpha617 self.aux_loss_enabled = getattr(model.config, "output_router_logits", False)618 self.aux_loss_coef = getattr(model.config, "router_aux_loss_coef", 0.0)619 if self.aux_loss_enabled and self.aux_loss_coef == 0.0:620 warnings.warn(621 "You set `output_router_logits` to `True` in the model config, but `router_aux_loss_coef` is set to "622 "`0.0`, meaning the auxiliary loss will not be used. Either set `router_aux_loss_coef` to a value "623 "greater than `0.0`, or set `output_router_logits` to `False` if you don't want to use the auxiliary "624 "loss.",625 UserWarning,626 )627 628 if args.loss_type == "simpo":629 self.simpo_gamma = args.simpo_gamma630 631 self._stored_metrics = defaultdict(lambda: defaultdict(list))632 633 # The trainer estimates the number of FLOPs (floating-point operations) using the number of elements in the634 # input tensor associated with the key "input_ids". However, in CPO, the sampled data does not include the635 # "input_ids" key. Instead, the available keys are "prompt_input_ids", "chosen_input_ids", and636 # "rejected_input_ids". As a result, the trainer issues the warning: "Could not estimate the number of tokens637 # of the input, floating-point operations will not be computed." To suppress this warning, we set the638 # "estimate_tokens" key in the model's "warnings_issued" dictionary to True. This acts as a flag to indicate639 # that the warning has already been issued.640 model.warnings_issued["estimate_tokens"] = True641 642 # Compute that only on the main process for faster data processing.643 # see: https://github.com/huggingface/trl/pull/1255644 with PartialState().local_main_process_first():645 # Extract the prompt if needed, and apply the chat template if needed646 train_dataset = train_dataset.map(maybe_extract_prompt, num_proc=args.dataset_num_proc)647 train_dataset = train_dataset.map(648 maybe_apply_chat_template, fn_kwargs={"tokenizer": processing_class}, num_proc=args.dataset_num_proc649 )650 if eval_dataset is not None:651 eval_dataset = eval_dataset.map(maybe_extract_prompt, num_proc=args.dataset_num_proc)652 eval_dataset = eval_dataset.map(653 maybe_apply_chat_template,654 fn_kwargs={"tokenizer": processing_class},655 num_proc=args.dataset_num_proc,656 )657 658 # tokenize the dataset659 train_dataset = train_dataset.map(self.tokenize_row, num_proc=args.dataset_num_proc)660 if eval_dataset is not None:661 eval_dataset = eval_dataset.map(self.tokenize_row, num_proc=args.dataset_num_proc)662 663 super().__init__(664 model=model,665 args=args,666 data_collator=data_collator,667 train_dataset=train_dataset,668 eval_dataset=eval_dataset,669 processing_class=processing_class,670 model_init=model_init,671 compute_metrics=compute_metrics,672 callbacks=callbacks,673 optimizers=optimizers,674 preprocess_logits_for_metrics=preprocess_logits_for_metrics,675 )676 677 # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the678 # model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set679 # self.model_accepts_loss_kwargs to False to enable scaling.680 self.model_accepts_loss_kwargs = False681 682 # Add tags for models that have been loaded with the correct transformers version683 if hasattr(self.model, "add_model_tags"):684 self.model.add_model_tags(self._tag_names)685 686 if not hasattr(self, "accelerator"):687 raise AttributeError(688 "Your `Trainer` does not have an `accelerator` object. Consider upgrading `transformers`."689 )690 691 def build_tokenized_answer(self, prompt, answer):692 """693 Llama tokenizer does satisfy `enc(a + b) = enc(a) + enc(b)`.694 It does ensure `enc(a + b) = enc(a) + enc(a + b)[len(enc(a)):]`.695 Reference:696 https://github.com/EleutherAI/lm-evaluation-harness/pull/531#issuecomment-1595586257697 """698 699 full_tokenized = self.processing_class(prompt + answer, add_special_tokens=False)700 prompt_input_ids = self.processing_class(prompt, add_special_tokens=False)["input_ids"]701 702 answer_input_ids = full_tokenized["input_ids"][len(prompt_input_ids) :]703 answer_attention_mask = full_tokenized["attention_mask"][len(prompt_input_ids) :]704 705 # Concat tokens to form `enc(a) + enc(a + b)[len(enc(a)):]`706 full_concat_input_ids = np.concatenate([prompt_input_ids, answer_input_ids])707 708 # Prepare input tokens for token by token comparison709 full_input_ids = np.array(full_tokenized["input_ids"])710 711 if len(full_input_ids) != len(full_concat_input_ids):712 raise ValueError("Prompt input ids and answer input ids should have the same length.")713 714 # On some tokenizers, like Llama-2 tokenizer, there are occasions where tokens715 # can be merged together when tokenizing prompt+answer. This could result716 # on the last token from the prompt being different when tokenized on its own717 # vs when done as prompt+answer.718 response_token_ids_start_idx = len(prompt_input_ids)719 720 # If tokenized prompt is different than both prompt+answer, then it means the721 # last token has changed due to merging.722 if prompt_input_ids != full_tokenized["input_ids"][:response_token_ids_start_idx]:723 response_token_ids_start_idx -= 1724 725 prompt_input_ids = full_tokenized["input_ids"][:response_token_ids_start_idx]726 prompt_attention_mask = full_tokenized["attention_mask"][:response_token_ids_start_idx]727 728 if len(prompt_input_ids) != len(prompt_attention_mask):729 raise ValueError("Prompt input ids and attention mask should have the same length.")730 731 answer_input_ids = full_tokenized["input_ids"][response_token_ids_start_idx:]732 answer_attention_mask = full_tokenized["attention_mask"][response_token_ids_start_idx:]733 734 return dict(735 prompt_input_ids=prompt_input_ids,736 prompt_attention_mask=prompt_attention_mask,737 input_ids=answer_input_ids,738 attention_mask=answer_attention_mask,739 )740 741 def tokenize_row(self, feature, model: Optional[Union[PreTrainedModel, nn.Module]] = None) -> dict:742 """Tokenize a single row from a CPO specific dataset.743 744 At this stage, we don't convert to PyTorch tensors yet; we just handle the truncation745 in case the prompt + chosen or prompt + rejected responses is/are too long. First746 we truncate the prompt; if we're still too long, we truncate the chosen/rejected.747 748 We also create the labels for the chosen/rejected responses, which are of length equal to749 the sum of the length of the prompt and the chosen/rejected response, with750 label_pad_token_id for the prompt tokens.751 """752 batch = {}753 prompt = feature["prompt"]754 chosen = feature["chosen"]755 rejected = feature["rejected"]756 757 if not self.is_encoder_decoder:758 # Check issues below for more details759 # 1. https://github.com/huggingface/trl/issues/907760 # 2. https://github.com/EleutherAI/lm-evaluation-harness/pull/531#issuecomment-1595586257761 # 3. https://github.com/LianjiaTech/BELLE/issues/337762 763 if not isinstance(prompt, str):764 raise ValueError(f"prompt should be an str but got {type(prompt)}")765 prompt_tokens = self.processing_class(prompt, add_special_tokens=False)766 prompt_tokens = {f"prompt_{k}": v for k, v in prompt_tokens.items()}767 768 if not isinstance(chosen, str):769 raise ValueError(f"chosen should be an str but got {type(chosen)}")770 chosen_tokens = self.build_tokenized_answer(prompt, chosen)771 772 if not isinstance(rejected, str):773 raise ValueError(f"rejected should be an str but got {type(rejected)}")774 rejected_tokens = self.build_tokenized_answer(prompt, rejected)775 776 # Last prompt token might get merged by tokenizer and777 # it should not be included for generation if that happens778 prompt_len_input_ids = len(prompt_tokens["prompt_input_ids"])779 780 chosen_prompt_len_input_ids = len(chosen_tokens["prompt_input_ids"])781 rejected_prompt_len_input_ids = len(rejected_tokens["prompt_input_ids"])782 prompt_len_input_ids = min(chosen_prompt_len_input_ids, rejected_prompt_len_input_ids)783 784 for k, v in prompt_tokens.items():785 prompt_tokens[k] = v[:prompt_len_input_ids]786 787 # Make sure prompts only have one different token at most an788 # and length only differs by 1 at most789 num_diff_tokens = sum(790 [a != b for a, b in zip(chosen_tokens["prompt_input_ids"], rejected_tokens["prompt_input_ids"])]791 )792 num_diff_len = abs(chosen_prompt_len_input_ids - rejected_prompt_len_input_ids)793 if num_diff_tokens > 1 or num_diff_len > 1:794 raise ValueError(795 "Chosen and rejected prompt_input_ids might only differ on the "796 "last token due to tokenizer merge ops."797 )798 799 # add BOS token to head of prompt. Avoid adding if it's already there800 prompt_tokens, chosen_tokens, rejected_tokens = add_bos_token_if_needed(801 self.processing_class.bos_token_id,802 prompt_len_input_ids,803 prompt_tokens,804 chosen_prompt_len_input_ids,805 chosen_tokens,806 rejected_prompt_len_input_ids,807 rejected_tokens,808 )809 810 # add EOS token to end of answer. Avoid adding if it's already there811 chosen_tokens, rejected_tokens = add_eos_token_if_needed(812 self.processing_class.eos_token_id, chosen_tokens, rejected_tokens813 )814 815 longer_response_length = max(len(chosen_tokens["input_ids"]), len(rejected_tokens["input_ids"]))816 817 # if combined sequence is too long, truncate the prompt818 for answer_tokens in [chosen_tokens, rejected_tokens, prompt_tokens]:819 if len(answer_tokens["prompt_input_ids"]) + longer_response_length > self.max_length:820 if self.truncation_mode == "keep_start":821 for k in ["prompt_input_ids", "prompt_attention_mask"]:822 answer_tokens[k] = answer_tokens[k][: self.max_prompt_length]823 elif self.truncation_mode == "keep_end":824 for k in ["prompt_input_ids", "prompt_attention_mask"]:825 answer_tokens[k] = answer_tokens[k][-self.max_prompt_length :]826 else:827 raise ValueError(f"Unknown truncation mode: {self.truncation_mode}")828 829 # if that's still too long, truncate the response830 for answer_tokens in [chosen_tokens, rejected_tokens]:831 if len(answer_tokens["prompt_input_ids"]) + longer_response_length > self.max_length:832 for k in ["input_ids", "attention_mask"]:833 answer_tokens[k] = answer_tokens[k][: self.max_length - self.max_prompt_length]834 835 # Create labels836 chosen_sequence_tokens = {837 k: chosen_tokens[f"prompt_{k}"] + chosen_tokens[k] for k in ["input_ids", "attention_mask"]838 }839 rejected_sequence_tokens = {840 k: rejected_tokens[f"prompt_{k}"] + rejected_tokens[k] for k in ["input_ids", "attention_mask"]841 }842 chosen_sequence_tokens["labels"] = chosen_sequence_tokens["input_ids"][:]843 chosen_sequence_tokens["labels"][: len(chosen_tokens["prompt_input_ids"])] = [844 self.label_pad_token_id845 ] * len(chosen_tokens["prompt_input_ids"])846 rejected_sequence_tokens["labels"] = rejected_sequence_tokens["input_ids"][:]847 rejected_sequence_tokens["labels"][: len(rejected_tokens["prompt_input_ids"])] = [848 self.label_pad_token_id849 ] * len(rejected_tokens["prompt_input_ids"])850 851 for k, toks in {852 "chosen_": chosen_sequence_tokens,853 "rejected_": rejected_sequence_tokens,854 "": prompt_tokens,855 }.items():856 for type_key, tokens in toks.items():857 if type_key == "token_type_ids":858 continue859 batch[f"{k}{type_key}"] = tokens860 861 else:862 chosen_tokens = self.processing_class(863 chosen, truncation=True, max_length=self.max_completion_length, add_special_tokens=True864 )865 rejected_tokens = self.processing_class(866 rejected, truncation=True, max_length=self.max_completion_length, add_special_tokens=True867 )868 prompt_tokens = self.processing_class(869 prompt, truncation=True, max_length=self.max_prompt_length, add_special_tokens=True870 )871 872 batch["chosen_labels"] = chosen_tokens["input_ids"]873 batch["rejected_labels"] = rejected_tokens["input_ids"]874 batch["prompt_input_ids"] = prompt_tokens["input_ids"]875 batch["prompt_attention_mask"] = prompt_tokens["attention_mask"]876 877 if model is not None and hasattr(model, "prepare_decoder_input_ids_from_labels"):878 batch["rejected_decoder_input_ids"] = model.prepare_decoder_input_ids_from_labels(879 labels=torch.tensor(batch["rejected_labels"])880 )881 batch["chosen_decoder_input_ids"] = model.prepare_decoder_input_ids_from_labels(882 labels=torch.tensor(batch["chosen_labels"])883 )884 885 return batch886 887 @staticmethod888 def concatenated_inputs(889 batch: dict[str, Union[list, torch.LongTensor]],890 is_encoder_decoder: bool = False,891 label_pad_token_id: int = -100,892 padding_value: int = 0,893 device: Optional[torch.device] = None,894 ) -> dict[str, torch.LongTensor]:895 """Concatenate the chosen and rejected inputs into a single tensor.896 897 Args:898 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).899 is_encoder_decoder: Whether the model is an encoder-decoder model.900 label_pad_token_id: The label pad token id.901 padding_value: The padding value to use for the concatenated inputs_ids.902 device: The device for the concatenated inputs.903 904 Returns:905 A dictionary containing the concatenated inputs under the key 'concatenated_input_ids'.906 """907 concatenated_batch = {}908 909 if is_encoder_decoder:910 max_length = max(batch["chosen_labels"].shape[1], batch["rejected_labels"].shape[1])911 else:912 max_length = max(batch["chosen_input_ids"].shape[1], batch["rejected_input_ids"].shape[1])913 914 for k in batch:915 if k.startswith("chosen") and isinstance(batch[k], torch.Tensor):916 if "labels" in k or is_encoder_decoder:917 pad_value = label_pad_token_id918 elif k.endswith("_input_ids"):919 pad_value = padding_value920 elif k.endswith("_attention_mask"):921 pad_value = 0922 concatenated_key = k.replace("chosen", "concatenated")923 concatenated_batch[concatenated_key] = pad_to_length(batch[k], max_length, pad_value=pad_value)924 for k in batch:925 if k.startswith("rejected") and isinstance(batch[k], torch.Tensor):926 if "labels" in k or is_encoder_decoder:927 pad_value = label_pad_token_id928 elif k.endswith("_input_ids"):929 pad_value = padding_value930 elif k.endswith("_attention_mask"):931 pad_value = 0932 concatenated_key = k.replace("rejected", "concatenated")933 concatenated_batch[concatenated_key] = torch.cat(934 (935 concatenated_batch[concatenated_key],936 pad_to_length(batch[k], max_length, pad_value=pad_value),937 ),938 dim=0,939 ).to(device=device)940 941 if is_encoder_decoder:942 concatenated_batch["concatenated_input_ids"] = batch["prompt_input_ids"].repeat(2, 1).to(device=device)943 concatenated_batch["concatenated_attention_mask"] = (944 batch["prompt_attention_mask"].repeat(2, 1).to(device=device)945 )946 947 return concatenated_batch948 949 def cpo_loss(950 self,951 policy_chosen_logps: torch.FloatTensor,952 policy_rejected_logps: torch.FloatTensor,953 ) -> tuple[torch.FloatTensor, torch.FloatTensor, torch.FloatTensor]:954 """Compute the CPO loss for a batch of policy and reference model log probabilities.955 956 Args:957 policy_chosen_logps: Log probabilities of the policy model for the chosen responses. Shape: (batch_size,)958 policy_rejected_logps: Log probabilities of the policy model for the rejected responses. Shape: (batch_size,)959 960 Returns:961 A tuple of three tensors: (losses, chosen_rewards, rejected_rewards).962 The losses tensor contains the CPO loss for each example in the batch.963 The chosen_rewards and rejected_rewards tensors contain the rewards for the chosen and rejected responses, respectively.964 """965 logits = (policy_chosen_logps - policy_rejected_logps).to(self.accelerator.device)966 967 # The beta is a temperature parameter for the CPO loss, typically something in the range of 0.1 to 0.5.968 # We ignore the reference model as beta -> 0. The label_smoothing parameter encodes our uncertainty about the labels and969 # calculates a conservative CPO loss.970 971 if self.loss_type == "simpo":972 gamma_logratios = self.simpo_gamma / self.beta973 logits = logits - gamma_logratios974 # This reduces to Equation 3 from the CPO paper when label_smoothing -> 0.975 losses = (976 -F.logsigmoid(self.beta * logits) * (1 - self.label_smoothing)977 - F.logsigmoid(-self.beta * logits) * self.label_smoothing978 )979 elif self.loss_type == "sigmoid":980 # This reduces to Equation 3 from the CPO paper when label_smoothing -> 0.981 losses = (982 -F.logsigmoid(self.beta * logits) * (1 - self.label_smoothing)983 - F.logsigmoid(-self.beta * logits) * self.label_smoothing984 )985 elif self.loss_type == "hinge":986 losses = torch.relu(1 - self.beta * logits)987 elif self.loss_type == "ipo":988 # eqn (17) of the paper where beta is the regularization parameter for the IPO loss, denoted by tau in the paper.989 losses = (logits - 1 / (2 * self.beta)) ** 2990 else:991 raise ValueError(992 f"Unknown loss type: {self.loss_type}. Should be one of ['sigmoid', 'hinge', 'ipo', 'simpo']"993 )994 995 chosen_rewards = self.beta * (policy_chosen_logps.to(self.accelerator.device)).detach()996 rejected_rewards = self.beta * (policy_rejected_logps.to(self.accelerator.device)).detach()997 998 return losses, chosen_rewards, rejected_rewards999 1000 @staticmethod1001 def get_batch_logps(1002 logits: torch.FloatTensor,1003 labels: torch.LongTensor,1004 average_log_prob: bool = False,1005 label_pad_token_id: int = -100,1006 is_encoder_decoder: bool = False,1007 ) -> torch.FloatTensor:1008 """Compute the log probabilities of the given labels under the given logits.1009 1010 Args:1011 logits: Logits of the model (unnormalized). Shape: (batch_size, sequence_length, vocab_size)1012 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)1013 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.1014 label_pad_token_id: The label pad token id.1015 is_encoder_decoder: Whether the model is an encoder-decoder model.1016 1017 Returns:1018 A tensor of shape (batch_size,) containing the average/sum log probabilities of the given labels under the given logits.1019 """1020 if logits.shape[:-1] != labels.shape:1021 raise ValueError("Logits (batch and sequence length dim) and labels must have the same shape.")1022 1023 if not is_encoder_decoder:1024 labels = labels[:, 1:].clone()1025 logits = logits[:, :-1, :]1026 loss_mask = labels != label_pad_token_id1027 1028 # dummy token; we'll ignore the losses on these tokens later1029 labels[labels == label_pad_token_id] = 01030 1031 per_token_logps = selective_log_softmax(logits, labels)1032 1033 if average_log_prob:1034 return (per_token_logps * loss_mask).sum(-1) / loss_mask.sum(-1)1035 else:1036 return (per_token_logps * loss_mask).sum(-1)1037 1038 def concatenated_forward(1039 self, model: nn.Module, batch: dict[str, Union[list, torch.LongTensor]]1040 ) -> tuple[torch.FloatTensor, torch.FloatTensor, torch.FloatTensor, torch.FloatTensor]:1041 """Run the given model on the given batch of inputs, concatenating the chosen and rejected inputs together.1042 1043 We do this to avoid doing two forward passes, because it's faster for FSDP.1044 """1045 concatenated_batch = self.concatenated_inputs(1046 batch,1047 is_encoder_decoder=self.is_encoder_decoder,1048 label_pad_token_id=self.label_pad_token_id,1049 padding_value=self.padding_value,1050 device=self.accelerator.device,1051 )1052 len_chosen = batch["chosen_labels"].shape[0]1053 1054 model_kwargs = (1055 {1056 "decoder_input_ids": self._shift_right(concatenated_batch["concatenated_labels"]),1057 }1058 if self.is_encoder_decoder1059 else {}1060 )1061 1062 if self.aux_loss_enabled:1063 model_kwargs["output_router_logits"] = True1064 1065 outputs = model(1066 concatenated_batch["concatenated_input_ids"],1067 attention_mask=concatenated_batch["concatenated_attention_mask"],1068 use_cache=False,1069 **model_kwargs,1070 )1071 all_logits = outputs.logits1072 1073 def cross_entropy_loss(logits, labels):1074 if not self.is_encoder_decoder:1075 # Shift so that tokens < n predict n1076 logits = logits[..., :-1, :].contiguous()1077 labels = labels[..., 1:].contiguous()1078 # Flatten the tokens1079 loss_fct = nn.CrossEntropyLoss()1080 logits = logits.view(-1, logits.shape[-1])1081 labels = labels.view(-1)1082 # Enable model parallelism1083 labels = labels.to(logits.device)1084 loss = loss_fct(logits, labels)1085 return loss1086 1087 labels = concatenated_batch["concatenated_labels"].clone()1088 1089 if self.cpo_alpha == 0:1090 nll_loss = torch.tensor(0.0).to(self.accelerator.device)1091 else:1092 nll_loss = cross_entropy_loss(all_logits[:len_chosen], labels[:len_chosen])1093 1094 all_logps = self.get_batch_logps(1095 all_logits,1096 concatenated_batch["concatenated_labels"],1097 average_log_prob=self.loss_type in ["ipo", "simpo"],1098 is_encoder_decoder=self.is_encoder_decoder,1099 label_pad_token_id=self.label_pad_token_id,1100 )1101 1102 chosen_logps = all_logps[:len_chosen]1103 rejected_logps = all_logps[len_chosen:]1104 1105 chosen_logits = all_logits[:len_chosen]1106 rejected_logits = all_logits[len_chosen:]1107 1108 if self.aux_loss_enabled:1109 return (chosen_logps, rejected_logps, chosen_logits, rejected_logits, nll_loss, outputs.aux_loss)1110 1111 return (chosen_logps, rejected_logps, chosen_logits, rejected_logits, nll_loss)1112 1113 def get_batch_loss_metrics(1114 self,1115 model,1116 batch: dict[str, Union[list, torch.LongTensor]],1117 train_eval: Literal["train", "eval"] = "train",1118 ):1119 """Compute the CPO loss and other metrics for the given batch of inputs for train or test."""1120 metrics = {}1121 1122 forward_output = self.concatenated_forward(model, batch)1123 (1124 policy_chosen_logps,1125 policy_rejected_logps,1126 policy_chosen_logits,1127 policy_rejected_logits,1128 policy_nll_loss,1129 ) = forward_output[:5]1130 if self.aux_loss_enabled:1131 aux_loss = forward_output[5]1132 1133 losses, chosen_rewards, rejected_rewards = self.cpo_loss(1134 policy_chosen_logps,1135 policy_rejected_logps,1136 )1137 1138 loss = losses.mean() + self.cpo_alpha * policy_nll_loss1139 reward_accuracies = (chosen_rewards > rejected_rewards).float()1140 1141 prefix = "eval_" if train_eval == "eval" else ""1142 metrics[f"{prefix}rewards/chosen"] = self.accelerator.gather_for_metrics(chosen_rewards).mean().item()1143 metrics[f"{prefix}rewards/rejected"] = self.accelerator.gather_for_metrics(rejected_rewards).mean().item()1144 metrics[f"{prefix}rewards/accuracies"] = self.accelerator.gather_for_metrics(reward_accuracies).mean().item()1145 metrics[f"{prefix}rewards/margins"] = (1146 self.accelerator.gather_for_metrics(chosen_rewards - rejected_rewards).mean().item()1147 )1148 metrics[f"{prefix}logps/rejected"] = (1149 self.accelerator.gather_for_metrics(policy_rejected_logps).detach().mean().item()1150 )1151 metrics[f"{prefix}logps/chosen"] = (1152 self.accelerator.gather_for_metrics(policy_chosen_logps).detach().mean().item()1153 )1154 metrics[f"{prefix}logits/rejected"] = (1155 self.accelerator.gather_for_metrics(policy_rejected_logits).detach().mean().item()1156 )1157 metrics[f"{prefix}logits/chosen"] = (1158 self.accelerator.gather_for_metrics(policy_chosen_logits).detach().mean().item()1159 )1160 metrics[f"{prefix}nll_loss"] = self.accelerator.gather_for_metrics(policy_nll_loss).detach().mean().item()1161 1162 if self.aux_loss_enabled:1163 loss += self.aux_loss_coef * aux_loss1164 1165 return loss, metrics1166 1167 def compute_loss(1168 self,1169 model: Union[PreTrainedModel, nn.Module],1170 inputs: dict[str, Union[torch.Tensor, Any]],1171 return_outputs=False,1172 num_items_in_batch=None,1173 ) -> Union[torch.Tensor, tuple[torch.Tensor, dict[str, torch.Tensor]]]:1174 compute_loss_context_manager = amp.autocast("cuda") if self._peft_has_been_casted_to_bf16 else nullcontext()1175 1176 with compute_loss_context_manager:1177 loss, metrics = self.get_batch_loss_metrics(model, inputs, train_eval="train")1178 1179 # force log the metrics1180 self.store_metrics(metrics, train_eval="train")1181 1182 if return_outputs:1183 return (loss, metrics)1184 return loss1185 1186 def generate_from_model(self, model, batch: dict[str, torch.LongTensor]) -> str:1187 """Generate samples from the model and reference model for the given batch of inputs."""1188 1189 # If one uses `generate_during_eval` with peft + bf16, we need to explicitly call generate with1190 # the torch cuda amp context manager as some hidden states are silently casted to full precision.1191 generate_context_manager = amp.autocast("cuda") if self._peft_has_been_casted_to_bf16 else nullcontext()1192 1193 with generate_context_manager:1194 policy_output = model.generate(1195 input_ids=batch["prompt_input_ids"],1196 attention_mask=batch["prompt_attention_mask"],1197 max_length=self.max_length,1198 do_sample=True,1199 pad_token_id=self.processing_class.pad_token_id,1200 )