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.kto_trainer import (Any, AutoModelForCausalLM, BaseImageProcessor, Callable, DPODataCollatorWithPadding, DataCollator, DataLoader, Dataset, EvalLoopOutput, F, FeatureExtractionMixin, KTOConfig, KTOTrainer, Literal, Optional, PartialState, PeftModel, PreTrainedModel, PreTrainedModelWrapper, PreTrainedTokenizerBase, ProcessorMixin, SequentialSampler, Trainer, TrainerCallback, TrainingArguments, Union, _get_kl_dataset, _process_tokens, _tokenize, amp, concatenate_datasets, contextmanager, create_reference_model, deepcopy, defaultdict, disable_dropout_in_model, generate_model_card, get_comet_experiment_url, has_length, inspect, is_comet_available, is_peft_available, is_wandb_available, itemgetter, log_table_to_comet_experiment, maybe_apply_chat_template, maybe_extract_prompt, maybe_unpair_preference_dataset, nn, np, nullcontext, os, pad_to_length, pd, peft_module_casting_to_bf16, prepare_model_for_kbit_training, random, textwrap, torch, tqdm, transformers, version, warnings)13 14 15import os16from typing import *17from dataclasses import dataclass, field18from packaging.version import Version19import torch20import numpy as np21from contextlib import nullcontext22from torch.nn import functional as F23from transformers import DataCollatorForSeq2Seq, DataCollatorForLanguageModeling24 25torch_compile_options = {26 "epilogue_fusion" : True,27 "max_autotune" : False,28 "shape_padding" : True,29 "trace.enabled" : False,30 "triton.cudagraphs" : False,31}32 33@torch.compile(dynamic = True, fullgraph = True, options = torch_compile_options,)34def selective_log_softmax(logits, index):35 logits = logits.to(torch.float32)36 selected_logits = torch.gather(logits, dim = -1, index = index.unsqueeze(-1)).squeeze(-1)37 # loop to reduce peak mem consumption38 # logsumexp_values = torch.stack([torch.logsumexp(lg, dim=-1) for lg in logits])39 logsumexp_values = torch.logsumexp(logits, dim = -1)40 per_token_logps = selected_logits - logsumexp_values # log_softmax(x_i) = x_i - logsumexp(x)41 return per_token_logps42@dataclass43class UnslothKTOConfig(KTOConfig):44 """45 46 Configuration class for the [`KTOTrainer`].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 `5e-7`):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.67 loss_type (`str`, *optional*, defaults to `"kto"`):68 Type of loss to use. Possible values are:69 70 - `"kto"`: KTO loss from the [KTO](https://huggingface.co/papers/2402.01306) paper.71 - `"apo_zero_unpaired"`: Unpaired variant of APO-zero loss from the [APO](https://huggingface.co/papers/2408.06266) paper.72 73 desirable_weight (`float`, *optional*, defaults to `1.0`):74 Desirable losses are weighed by this factor to counter unequal number of desirable and undesirable paris.75 undesirable_weight (`float`, *optional*, defaults to `1.0`):76 Undesirable losses are weighed by this factor to counter unequal number of desirable and undesirable pairs.77 label_pad_token_id (`int`, *optional*, defaults to `-100`):78 Label pad token id. This argument is required if you want to use the default data collator.79 padding_value (`int` or `None`, *optional*, defaults to `None`):80 Padding value to use. If `None`, the padding value of the tokenizer is used.81 truncation_mode (`str`, *optional*, defaults to `"keep_end"`):82 Truncation mode to use when the prompt is too long. Possible values are `"keep_end"` or `"keep_start"`.83 This argument is required if you want to use the default data collator.84 generate_during_eval (`bool`, *optional*, defaults to `False`):85 If `True`, generates and logs completions from both the model and the reference model to W&B or Comet during86 evaluation.87 is_encoder_decoder (`bool` or `None`, *optional*, defaults to `None`):88 When using the `model_init` argument (callable) to instantiate the model instead of the `model` argument,89 you need to specify if the model returned by the callable is an encoder-decoder model.90 precompute_ref_log_probs (`bool`, *optional*, defaults to `False`):91 Whether to precompute reference model log probabilities for training and evaluation datasets. This is92 useful when training without the reference model to reduce the total GPU memory needed.93 model_init_kwargs (`dict[str, Any]` or `None`, *optional*, defaults to `None`):94 Keyword arguments to pass to `AutoModelForCausalLM.from_pretrained` when instantiating the model from a95 string.96 ref_model_init_kwargs (`dict[str, Any]` or `None`, *optional*, defaults to `None`):97 Keyword arguments to pass to `AutoModelForCausalLM.from_pretrained` when instantiating the reference model98 from a string.99 dataset_num_proc: (`int` or `None`, *optional*, defaults to `None`):100 Number of processes to use for processing the dataset.101 disable_dropout (`bool`, *optional*, defaults to `True`):102 Whether to disable dropout in the model and reference model.103 104 """105 vllm_sampling_params: Optional[Any] = field(106 default = None,107 metadata = {'help': 'vLLM SamplingParams'},108 )109 unsloth_num_chunks : Optional[int] = field(110 default = -1,111 metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},112 )113 def __init__(114 self,115 output_dir = None,116 overwrite_output_dir = None,117 do_train = False,118 do_eval = False,119 do_predict = False,120 eval_strategy = 'no',121 prediction_loss_only = False,122 per_device_train_batch_size = 4,123 per_device_eval_batch_size = 4,124 per_gpu_train_batch_size = None,125 per_gpu_eval_batch_size = None,126 gradient_accumulation_steps = 2,127 eval_accumulation_steps = 2,128 eval_delay = 0,129 torch_empty_cache_steps = 250,130 learning_rate = 5e-05,131 weight_decay = 0.01,132 adam_beta1 = 0.9,133 adam_beta2 = 0.999,134 adam_epsilon = 1e-08,135 max_grad_norm = 1.0,136 num_train_epochs = 3.0,137 max_steps = -1,138 lr_scheduler_type = 'linear',139 warmup_ratio = 0.1,140 warmup_steps = 0,141 log_level = 'passive',142 log_level_replica = 'warning',143 log_on_each_node = True,144 logging_dir = None,145 logging_strategy = 'steps',146 logging_first_step = False,147 logging_steps = 1,148 logging_nan_inf_filter = False,149 save_strategy = 'steps',150 save_steps = 500,151 save_total_limit = None,152 save_safetensors = True,153 save_on_each_node = False,154 save_only_model = False,155 restore_callback_states_from_checkpoint = False,156 no_cuda = False,157 use_cpu = False,158 use_mps_device = False,159 seed = 3407,160 data_seed = 3407,161 jit_mode_eval = False,162 use_ipex = False,163 bf16 = False,164 fp16 = False,165 fp16_opt_level = 'O1',166 half_precision_backend = 'auto',167 bf16_full_eval = False,168 fp16_full_eval = False,169 tf32 = None,170 local_rank = -1,171 ddp_backend = None,172 tpu_num_cores = None,173 tpu_metrics_debug = False,174 debug = '',175 dataloader_drop_last = False,176 eval_steps = None,177 dataloader_num_workers = 0,178 dataloader_prefetch_factor = None,179 past_index = -1,180 run_name = None,181 disable_tqdm = None,182 remove_unused_columns = True,183 label_names = None,184 load_best_model_at_end = False,185 metric_for_best_model = None,186 greater_is_better = None,187 ignore_data_skip = False,188 fsdp = '',189 fsdp_min_num_params = 0,190 fsdp_config = None,191 tp_size = 0,192 fsdp_transformer_layer_cls_to_wrap = None,193 accelerator_config = None,194 deepspeed = None,195 label_smoothing_factor = 0.0,196 optim = 'adamw_8bit',197 optim_args = None,198 adafactor = False,199 group_by_length = False,200 length_column_name = 'length',201 report_to = None,202 ddp_find_unused_parameters = None,203 ddp_bucket_cap_mb = None,204 ddp_broadcast_buffers = None,205 dataloader_pin_memory = True,206 dataloader_persistent_workers = False,207 skip_memory_metrics = True,208 use_legacy_prediction_loop = False,209 push_to_hub = False,210 resume_from_checkpoint = None,211 hub_model_id = None,212 hub_strategy = 'every_save',213 hub_token = None,214 hub_private_repo = None,215 hub_always_push = False,216 gradient_checkpointing = False,217 gradient_checkpointing_kwargs = None,218 include_inputs_for_metrics = False,219 eval_do_concat_batches = True,220 fp16_backend = 'auto',221 evaluation_strategy = None,222 push_to_hub_model_id = None,223 push_to_hub_organization = None,224 push_to_hub_token = None,225 mp_parameters = '',226 auto_find_batch_size = False,227 full_determinism = False,228 torchdynamo = None,229 ray_scope = 'last',230 ddp_timeout = 1800,231 torch_compile = False,232 torch_compile_backend = None,233 torch_compile_mode = None,234 dispatch_batches = None,235 split_batches = None,236 include_tokens_per_second = False,237 include_num_input_tokens_seen = False,238 neftune_noise_alpha = None,239 optim_target_modules = None,240 batch_eval_metrics = False,241 eval_on_start = False,242 use_liger_kernel = False,243 eval_use_gather_object = False,244 average_tokens_across_devices = False,245 max_length = 1024,246 max_prompt_length = 512,247 max_completion_length = None,248 beta = 0.1,249 loss_type = 'kto',250 desirable_weight = 1.0,251 undesirable_weight = 1.0,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 disable_dropout = True,258 precompute_ref_log_probs = False,259 model_init_kwargs = None,260 ref_model_init_kwargs = None,261 dataset_num_proc = None,262 vllm_sampling_params = None,263 unsloth_num_chunks = -1,264 **kwargs,265 ):266 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!')267 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!')268 if output_dir is None and save_strategy == 'steps' and save_steps == 500:269 output_dir = 'unsloth_training_checkpoints'270 save_strategy = 'no'271 if dataset_num_proc is None:272 from multiprocessing import cpu_count273 dataset_num_proc = cpu_count()274 275 super().__init__(276 output_dir = output_dir,277 overwrite_output_dir = overwrite_output_dir,278 do_train = do_train,279 do_eval = do_eval,280 do_predict = do_predict,281 eval_strategy = eval_strategy,282 prediction_loss_only = prediction_loss_only,283 per_device_train_batch_size = per_device_train_batch_size,284 per_device_eval_batch_size = per_device_eval_batch_size,285 per_gpu_train_batch_size = per_gpu_train_batch_size,286 per_gpu_eval_batch_size = per_gpu_eval_batch_size,287 gradient_accumulation_steps = gradient_accumulation_steps,288 eval_accumulation_steps = eval_accumulation_steps,289 eval_delay = eval_delay,290 torch_empty_cache_steps = torch_empty_cache_steps,291 learning_rate = learning_rate,292 weight_decay = weight_decay,293 adam_beta1 = adam_beta1,294 adam_beta2 = adam_beta2,295 adam_epsilon = adam_epsilon,296 max_grad_norm = max_grad_norm,297 num_train_epochs = num_train_epochs,298 max_steps = max_steps,299 lr_scheduler_type = lr_scheduler_type,300 warmup_ratio = warmup_ratio,301 warmup_steps = warmup_steps,302 log_level = log_level,303 log_level_replica = log_level_replica,304 log_on_each_node = log_on_each_node,305 logging_dir = logging_dir,306 logging_strategy = logging_strategy,307 logging_first_step = logging_first_step,308 logging_steps = logging_steps,309 logging_nan_inf_filter = logging_nan_inf_filter,310 save_strategy = save_strategy,311 save_steps = save_steps,312 save_total_limit = save_total_limit,313 save_safetensors = save_safetensors,314 save_on_each_node = save_on_each_node,315 save_only_model = save_only_model,316 restore_callback_states_from_checkpoint = restore_callback_states_from_checkpoint,317 no_cuda = no_cuda,318 use_cpu = use_cpu,319 use_mps_device = use_mps_device,320 seed = seed,321 data_seed = data_seed,322 jit_mode_eval = jit_mode_eval,323 use_ipex = use_ipex,324 bf16 = bf16,325 fp16 = fp16,326 fp16_opt_level = fp16_opt_level,327 half_precision_backend = half_precision_backend,328 bf16_full_eval = bf16_full_eval,329 fp16_full_eval = fp16_full_eval,330 tf32 = tf32,331 local_rank = local_rank,332 ddp_backend = ddp_backend,333 tpu_num_cores = tpu_num_cores,334 tpu_metrics_debug = tpu_metrics_debug,335 debug = debug,336 dataloader_drop_last = dataloader_drop_last,337 eval_steps = eval_steps,338 dataloader_num_workers = dataloader_num_workers,339 dataloader_prefetch_factor = dataloader_prefetch_factor,340 past_index = past_index,341 run_name = run_name,342 disable_tqdm = disable_tqdm,343 remove_unused_columns = remove_unused_columns,344 label_names = label_names,345 load_best_model_at_end = load_best_model_at_end,346 metric_for_best_model = metric_for_best_model,347 greater_is_better = greater_is_better,348 ignore_data_skip = ignore_data_skip,349 fsdp = fsdp,350 fsdp_min_num_params = fsdp_min_num_params,351 fsdp_config = fsdp_config,352 tp_size = tp_size,353 fsdp_transformer_layer_cls_to_wrap = fsdp_transformer_layer_cls_to_wrap,354 accelerator_config = accelerator_config,355 deepspeed = deepspeed,356 label_smoothing_factor = label_smoothing_factor,357 optim = optim,358 optim_args = optim_args,359 adafactor = adafactor,360 group_by_length = group_by_length,361 length_column_name = length_column_name,362 report_to = report_to,363 ddp_find_unused_parameters = ddp_find_unused_parameters,364 ddp_bucket_cap_mb = ddp_bucket_cap_mb,365 ddp_broadcast_buffers = ddp_broadcast_buffers,366 dataloader_pin_memory = dataloader_pin_memory,367 dataloader_persistent_workers = dataloader_persistent_workers,368 skip_memory_metrics = skip_memory_metrics,369 use_legacy_prediction_loop = use_legacy_prediction_loop,370 push_to_hub = push_to_hub,371 resume_from_checkpoint = resume_from_checkpoint,372 hub_model_id = hub_model_id,373 hub_strategy = hub_strategy,374 hub_token = hub_token,375 hub_private_repo = hub_private_repo,376 hub_always_push = hub_always_push,377 gradient_checkpointing = gradient_checkpointing,378 gradient_checkpointing_kwargs = gradient_checkpointing_kwargs,379 include_inputs_for_metrics = include_inputs_for_metrics,380 eval_do_concat_batches = eval_do_concat_batches,381 fp16_backend = fp16_backend,382 evaluation_strategy = evaluation_strategy,383 push_to_hub_model_id = push_to_hub_model_id,384 push_to_hub_organization = push_to_hub_organization,385 push_to_hub_token = push_to_hub_token,386 mp_parameters = mp_parameters,387 auto_find_batch_size = auto_find_batch_size,388 full_determinism = full_determinism,389 torchdynamo = torchdynamo,390 ray_scope = ray_scope,391 ddp_timeout = ddp_timeout,392 torch_compile = torch_compile,393 torch_compile_backend = torch_compile_backend,394 torch_compile_mode = torch_compile_mode,395 dispatch_batches = dispatch_batches,396 split_batches = split_batches,397 include_tokens_per_second = include_tokens_per_second,398 include_num_input_tokens_seen = include_num_input_tokens_seen,399 neftune_noise_alpha = neftune_noise_alpha,400 optim_target_modules = optim_target_modules,401 batch_eval_metrics = batch_eval_metrics,402 eval_on_start = eval_on_start,403 use_liger_kernel = use_liger_kernel,404 eval_use_gather_object = eval_use_gather_object,405 average_tokens_across_devices = average_tokens_across_devices,406 max_length = max_length,407 max_prompt_length = max_prompt_length,408 max_completion_length = max_completion_length,409 beta = beta,410 loss_type = loss_type,411 desirable_weight = desirable_weight,412 undesirable_weight = undesirable_weight,413 label_pad_token_id = label_pad_token_id,414 padding_value = padding_value,415 truncation_mode = truncation_mode,416 generate_during_eval = generate_during_eval,417 is_encoder_decoder = is_encoder_decoder,418 disable_dropout = disable_dropout,419 precompute_ref_log_probs = precompute_ref_log_probs,420 model_init_kwargs = model_init_kwargs,421 ref_model_init_kwargs = ref_model_init_kwargs,422 dataset_num_proc = dataset_num_proc,**kwargs)423 self.vllm_sampling_params = vllm_sampling_params424 self.unsloth_num_chunks = unsloth_num_chunks425pass426 427class _UnslothKTOTrainer(Trainer):428 r""""""429 430 _tag_names = ["trl", "kto"]431 432 def __init__(433 self,434 model: Union[PreTrainedModel, nn.Module, str] = None,435 ref_model: Optional[Union[PreTrainedModel, nn.Module, str]] = None,436 args: KTOConfig = None,437 train_dataset: Optional[Dataset] = None,438 eval_dataset: Optional[Union[Dataset, dict[str, Dataset]]] = None,439 processing_class: Optional[440 Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin]441 ] = None,442 data_collator: Optional[DataCollator] = None,443 model_init: Optional[Callable[[], PreTrainedModel]] = None,444 callbacks: Optional[list[TrainerCallback]] = None,445 optimizers: tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),446 preprocess_logits_for_metrics: Optional[Callable[[torch.Tensor, torch.Tensor], torch.Tensor]] = None,447 peft_config: Optional[dict] = None,448 compute_metrics: Optional[Callable[[EvalLoopOutput], dict]] = None,449 model_adapter_name: Optional[str] = None,450 ref_adapter_name: Optional[str] = None,451 ):452 if type(args) is TrainingArguments:453 raise ValueError("Please use `KTOConfig` instead TrainingArguments.")454 455 if not isinstance(model, str) and ref_model is model:456 raise ValueError(457 "`model` and `ref_model` cannot be the same object. If you want `ref_model` to be the "458 "same as `model`, you must mass a copy of it, or `None` if you use peft."459 )460 461 if args.model_init_kwargs is None:462 model_init_kwargs = {}463 elif not isinstance(model, str):464 raise ValueError("You passed model_kwargs to the KTOTrainer. But your model is already instantiated.")465 else:466 model_init_kwargs = args.model_init_kwargs467 torch_dtype = model_init_kwargs.get("torch_dtype")468 if torch_dtype is not None:469 # Convert to `torch.dtype` if an str is passed470 if isinstance(torch_dtype, str) and torch_dtype != "auto":471 torch_dtype = getattr(torch, torch_dtype)472 if torch_dtype != "auto" and not isinstance(torch_dtype, torch.dtype):473 raise ValueError(474 f"Invalid `torch_dtype` passed to the KTOConfig. Expected a string with either `torch.dtype` or 'auto', but got {torch_dtype}."475 )476 model_init_kwargs["torch_dtype"] = torch_dtype477 478 if args.ref_model_init_kwargs is None:479 ref_model_init_kwargs = {}480 elif not isinstance(ref_model, str):481 raise ValueError(482 "You passed ref_model_kwargs to the KTOTrainer. But your ref_model is already instantiated."483 )484 else:485 ref_model_init_kwargs = args.ref_model_init_kwargs486 torch_dtype = ref_model_init_kwargs.get("torch_dtype")487 if torch_dtype is not None:488 # Convert to `torch.dtype` if an str is passed489 if isinstance(torch_dtype, str) and torch_dtype != "auto":490 torch_dtype = getattr(torch, torch_dtype)491 if torch_dtype != "auto" and not isinstance(torch_dtype, torch.dtype):492 raise ValueError(493 f"Invalid `torch_dtype` passed to the KTOConfig. Expected a string with either `torch.dtype` or 'auto', but got {torch_dtype}."494 )495 ref_model_init_kwargs["torch_dtype"] = torch_dtype496 497 if isinstance(model, str):498 model = AutoModelForCausalLM.from_pretrained(model, **model_init_kwargs)499 500 if isinstance(ref_model, str):501 ref_model = AutoModelForCausalLM.from_pretrained(ref_model, **ref_model_init_kwargs)502 503 # Initialize this variable to False. This helps tracking the case when `peft_module_casting_to_bf16`504 # has been called in order to properly call autocast if needed.505 self._peft_has_been_casted_to_bf16 = False506 507 if not is_peft_available() and peft_config is not None:508 raise ValueError(509 "PEFT is not installed and you passed a `peft_config` in the trainer's kwargs, please install it with `pip install peft` to use the PEFT models"510 )511 elif is_peft_available() and peft_config is not None:512 # if model is a peft model and we have a peft_config, we merge and unload it first513 if isinstance(model, PeftModel):514 model = model.merge_and_unload()515 516 if getattr(model, "is_loaded_in_8bit", False) or getattr(model, "is_loaded_in_4bit", False):517 _support_gc_kwargs = hasattr(518 args, "gradient_checkpointing_kwargs"519 ) and "gradient_checkpointing_kwargs" in list(520 inspect.signature(prepare_model_for_kbit_training).parameters521 )522 523 prepare_model_kwargs = {"use_gradient_checkpointing": args.gradient_checkpointing}524 525 if _support_gc_kwargs:526 prepare_model_kwargs["gradient_checkpointing_kwargs"] = args.gradient_checkpointing_kwargs527 528 model = prepare_model_for_kbit_training(model, **prepare_model_kwargs)529 elif getattr(args, "gradient_checkpointing", False):530 # For backward compatibility with older versions of transformers531 if hasattr(model, "enable_input_require_grads"):532 model.enable_input_require_grads()533 else:534 535 def make_inputs_require_grad(module, input, output):536 output.requires_grad_(True)537 538 model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)539 540 # get peft model with the given config541 model = model542 if args.bf16 and getattr(model, "is_loaded_in_4bit", False):543 peft_module_casting_to_bf16(model)544 # If args.bf16 we need to explicitly call `generate` with torch amp autocast context manager545 self._peft_has_been_casted_to_bf16 = True546 547 # For models that use gradient_checkpointing, we need to attach a hook that enables input548 # to explicitly have `requires_grad=True`, otherwise training will either silently549 # fail or completely fail.550 elif getattr(args, "gradient_checkpointing", False):551 # For backward compatibility with older versions of transformers552 if hasattr(model, "enable_input_require_grads"):553 model.enable_input_require_grads()554 else:555 556 def make_inputs_require_grad(module, input, output):557 output.requires_grad_(True)558 559 model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)560 561 if args.generate_during_eval and not (is_wandb_available() or is_comet_available()):562 raise ValueError(563 "`generate_during_eval=True` requires Weights and Biases or Comet to be installed."564 " Please install `wandb` or `comet-ml` to resolve."565 )566 567 if model is not None:568 self.is_encoder_decoder = model.config.is_encoder_decoder569 elif args.is_encoder_decoder is None:570 raise ValueError("When no model is provided, you need to pass the parameter is_encoder_decoder.")571 else:572 self.is_encoder_decoder = args.is_encoder_decoder573 574 self.is_peft_model = is_peft_available() and isinstance(model, PeftModel)575 self.model_adapter_name = model_adapter_name576 self.ref_adapter_name = ref_adapter_name577 578 if ref_model:579 self.ref_model = ref_model580 elif self.is_peft_model or args.precompute_ref_log_probs:581 # The `model` with adapters turned off will be used as the reference model582 self.ref_model = None583 else:584 self.ref_model = create_reference_model(model)585 586 if processing_class is None:587 raise ValueError(588 "max_length or a processing_class must be specified when using the default DPODataCollatorWithPadding"589 )590 if args.max_length is None:591 warnings.warn(592 "When using DPODataCollatorWithPadding, you should set `max_length` in the KTOTrainer's init"593 " it will be set to `512` by default, but you should do it yourself in the future.",594 UserWarning,595 )596 max_length = 512597 if args.max_length is not None:598 max_length = args.max_length599 600 if args.max_prompt_length is None:601 warnings.warn(602 "When using DPODataCollatorWithPadding, you should set `max_prompt_length` in the KTOTrainer's init"603 " it will be set to `128` by default, but you should do it yourself in the future.",604 UserWarning,605 )606 max_prompt_length = 128607 if args.max_prompt_length is not None:608 max_prompt_length = args.max_prompt_length609 610 max_completion_length = None611 if args.max_completion_length is None and self.is_encoder_decoder:612 warnings.warn(613 "When using DPODataCollatorWithPadding with an encoder decoder architecture, you should set `max_completion_length` in the KTOTrainer's init"614 " it will be set to `128` by default, but you should do it yourself in the future.",615 UserWarning,616 )617 max_completion_length = 128618 if args.max_completion_length is not None and self.is_encoder_decoder:619 max_completion_length = args.max_completion_length620 621 if data_collator is None:622 data_collator = DPODataCollatorWithPadding(623 pad_token_id=processing_class.pad_token_id,624 label_pad_token_id=args.label_pad_token_id,625 is_encoder_decoder=self.is_encoder_decoder,626 )627 628 if args.remove_unused_columns:629 args.remove_unused_columns = False630 # warn users631 warnings.warn(632 "When using DPODataCollatorWithPadding, you should set `remove_unused_columns=False` in your KTOConfig"633 " we have set it for you, but you should do it yourself in the future.",634 UserWarning,635 )636 637 self.use_dpo_data_collator = True638 else:639 self.use_dpo_data_collator = False640 641 # Disable dropout in the model and reference model642 if args.disable_dropout:643 disable_dropout_in_model(model)644 if self.ref_model is not None:645 disable_dropout_in_model(self.ref_model)646 647 self.loss_type = args.loss_type648 self.max_length = max_length649 self.generate_during_eval = args.generate_during_eval650 self.label_pad_token_id = args.label_pad_token_id651 self.padding_value = args.padding_value if args.padding_value is not None else processing_class.pad_token_id652 self.max_prompt_length = max_prompt_length653 self.truncation_mode = args.truncation_mode654 self.max_completion_length = max_completion_length655 self.processing_class = processing_class656 self.precompute_ref_log_probs = args.precompute_ref_log_probs657 658 # Not all losses require a KL calculation659 self.calculate_KL = True660 if self.loss_type in ["apo_zero_unpaired"]:661 self.calculate_KL = False662 663 # Since ref_logs are precomputed on the first call to get_train/eval_dataloader664 # keep track of first called to avoid computation of future calls665 self._precomputed_train_ref_log_probs = False666 self._precomputed_eval_ref_log_probs = False667 668 # metric669 self._stored_metrics = defaultdict(lambda: defaultdict(list))670 671 # KTO parameter672 self.beta = args.beta673 self.desirable_weight = args.desirable_weight674 self.undesirable_weight = args.undesirable_weight675 self.aux_loss_enabled = getattr(model.config, "output_router_logits", False)676 self.aux_loss_coef = getattr(model.config, "router_aux_loss_coef", 0.0)677 if self.aux_loss_enabled and self.aux_loss_coef == 0.0:678 warnings.warn(679 "You set `output_router_logits` to `True` in the model config, but `router_aux_loss_coef` is set to "680 "`0.0`, meaning the auxiliary loss will not be used. Either set `router_aux_loss_coef` to a value "681 "greater than `0.0`, or set `output_router_logits` to `False` if you don't want to use the auxiliary "682 "loss.",683 UserWarning,684 )685 686 # The trainer estimates the number of FLOPs (floating-point operations) using the number of elements in the687 # input tensor associated with the key "input_ids". However, in KTO, the sampled data does not include the688 # "input_ids" key. Instead, the available keys are "prompt_input_ids" and "completion_input_ids". As a result,689 # the trainer issues the warning: "Could not estimate the number of tokens of the input, floating-point690 # operations will not be computed." To suppress this warning, we set the "estimate_tokens" key in the model's691 # "warnings_issued" dictionary to True. This acts as a flag to indicate that the warning has already been692 # issued.693 model.warnings_issued["estimate_tokens"] = True694 695 # Compute that only on the main process for faster data processing.696 # see: https://github.com/huggingface/trl/pull/1255697 with PartialState().local_main_process_first():698 # Extract the prompt if needed699 train_dataset = train_dataset.map(700 maybe_extract_prompt, num_proc=args.dataset_num_proc, desc="Extracting prompt from train dataset"701 )702 # Unpair the dataset if needed703 train_dataset = maybe_unpair_preference_dataset(704 train_dataset, args.dataset_num_proc, desc="Unpairing train dataset"705 )706 # Apply the chat template if needed707 train_dataset = train_dataset.map(708 maybe_apply_chat_template,709 fn_kwargs={"tokenizer": processing_class},710 num_proc=args.dataset_num_proc,711 desc="Applying chat template to train dataset",712 )713 if eval_dataset is not None:714 eval_dataset = eval_dataset.map(715 maybe_extract_prompt, num_proc=args.dataset_num_proc, desc="Extracting prompt from eval dataset"716 )717 eval_dataset = maybe_unpair_preference_dataset(718 eval_dataset, args.dataset_num_proc, desc="Unpairing eval dataset"719 )720 eval_dataset = eval_dataset.map(721 maybe_apply_chat_template,722 fn_kwargs={"tokenizer": processing_class},723 num_proc=args.dataset_num_proc,724 desc="Applying chat template to eval dataset",725 )726 727 # Tokenize and prepare the training datasets728 train_dataset = train_dataset.map(729 _tokenize,730 batched=True,731 fn_kwargs={"tokenizer": self.processing_class},732 num_proc=args.dataset_num_proc,733 desc="Tokenizing train dataset",734 )735 736 fn_kwargs = {737 "prefix": "",738 "is_encoder_decoder": self.is_encoder_decoder,739 "tokenizer": self.processing_class,740 "max_length": self.max_length,741 "truncation_mode": self.truncation_mode,742 "label_pad_token_id": self.label_pad_token_id,743 "max_prompt_length": self.max_prompt_length,744 "max_completion_length": self.max_completion_length,745 }746 747 train_dataset = train_dataset.map(748 _process_tokens,749 fn_kwargs=fn_kwargs,750 num_proc=args.dataset_num_proc,751 desc="Processing tokenized train dataset",752 )753 754 # Tokenize and prepare the eval datasets755 if eval_dataset is not None:756 eval_dataset = eval_dataset.map(757 _tokenize,758 fn_kwargs={"tokenizer": self.processing_class},759 batched=True,760 num_proc=args.dataset_num_proc,761 desc="Tokenizing eval dataset",762 )763 764 eval_dataset = eval_dataset.map(765 _process_tokens,766 fn_kwargs=fn_kwargs,767 num_proc=args.dataset_num_proc,768 desc="Processing tokenized eval dataset",769 )770 771 # Get KL datasets if needed772 if self.calculate_KL:773 if args.per_device_train_batch_size <= 1:774 raise ValueError(775 "Actual (not effective) batch size must be > 1. KTO will not work properly because the KL term will be equivalent to the implied reward."776 )777 778 # create pairs for estimating the KL term by flipping the matched pairs in each batch of size total_batch_size779 # i.e., (x_1, y_1), ..., (x_n, y_n) --> (x_1, y_n), ..., (x_n, y_1) = (x'_1, y'_1), ..., (x'_n, y'_n)780 train_kl_dataset = train_dataset.map(781 _get_kl_dataset,782 batched=True,783 batch_size=args.per_device_train_batch_size,784 num_proc=args.dataset_num_proc,785 desc="Extracting KL train dataset",786 )787 788 fn_kwargs["prefix"] = "KL_"789 train_kl_dataset = train_kl_dataset.map(790 _process_tokens,791 fn_kwargs=fn_kwargs,792 num_proc=args.dataset_num_proc,793 remove_columns=[c for c in train_kl_dataset.column_names if c in train_dataset.column_names],794 desc="Processing tokenized train KL dataset",795 )796 797 # merge the datasets798 train_dataset = concatenate_datasets([train_dataset, train_kl_dataset], axis=1)799 800 if eval_dataset is not None:801 # Get KL dataset802 eval_kl_dataset = eval_dataset.map(803 _get_kl_dataset,804 batched=True,805 batch_size=args.per_device_train_batch_size,806 num_proc=args.dataset_num_proc,807 desc="Extracting eval KL dataset",808 )809 810 eval_kl_dataset = eval_kl_dataset.map(811 _process_tokens,812 fn_kwargs=fn_kwargs,813 num_proc=args.dataset_num_proc,814 remove_columns=[c for c in eval_kl_dataset.column_names if c in eval_dataset.column_names],815 desc="Processing tokenized eval KL dataset",816 )817 818 # merge the datasets819 eval_dataset = concatenate_datasets([eval_dataset, eval_kl_dataset], axis=1)820 821 # calculate dataset desirability balance822 num_desirable = max(sum(train_dataset["label"]), 1)823 num_undesirable = max(len(train_dataset["label"]) - num_desirable, 1) # "label" is binary824 825 if num_desirable != num_undesirable:826 # The lower and upper bounds come from Eq. (8) of https://huggingface.co/papers/2402.01306827 des_weight_lower_bound = round((num_undesirable * self.undesirable_weight / num_desirable) * 1, 2)828 des_weight_upper_bound = round((num_undesirable * self.undesirable_weight / num_desirable) * 1.33, 2)829 und_weight_lower_bound = round((num_desirable * self.desirable_weight / num_undesirable) / 1.33, 2)830 und_weight_upper_bound = round((num_desirable * self.desirable_weight / num_undesirable) / 1, 2)831 832 des_weight_in_range = des_weight_lower_bound <= self.desirable_weight <= des_weight_upper_bound833 und_weight_in_range = und_weight_lower_bound <= self.undesirable_weight <= und_weight_upper_bound834 835 if not (des_weight_in_range or und_weight_in_range):836 warnings.warn(837 "You have different amounts of desirable/positive and undesirable/negative examples but the "838 "weights on the desirable and undesirable losses don't seem to be in an ideal range. Based "839 f"on your data, we recommend EITHER "840 f"desirable_weight in [{des_weight_lower_bound}, {des_weight_upper_bound}] or "841 f"undesirable_weight in [{und_weight_lower_bound}, {und_weight_upper_bound}] (but NOT BOTH). "842 "See the documentation on how to optimally set these weights.",843 UserWarning,844 )845 846 super().__init__(847 model=model,848 args=args,849 data_collator=data_collator,850 train_dataset=train_dataset,851 eval_dataset=eval_dataset,852 processing_class=processing_class,853 model_init=model_init,854 compute_metrics=compute_metrics,855 callbacks=callbacks,856 optimizers=optimizers,857 preprocess_logits_for_metrics=preprocess_logits_for_metrics,858 )859 860 # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the861 # model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set862 # self.model_accepts_loss_kwargs to False to enable scaling.863 self.model_accepts_loss_kwargs = False864 865 # Add tags for models that have been loaded with the correct transformers version866 if hasattr(self.model, "add_model_tags"):867 self.model.add_model_tags(self._tag_names)868 869 if not hasattr(self, "accelerator"):870 raise AttributeError(871 "Your `Trainer` does not have an `accelerator` object. Consider upgrading `transformers`."872 )873 874 # Deepspeed Zero-3 does not support precompute_ref_log_probs875 if self.is_deepspeed_enabled:876 if self.accelerator.state.deepspeed_plugin.zero_stage == 3 and self.precompute_ref_log_probs:877 raise ValueError(878 "You cannot use `precompute_ref_log_probs=True` with Deepspeed ZeRO-3. Please set `precompute_ref_log_probs=False`."879 )880 881 if self.ref_model is None:882 if not (self.is_peft_model or self.precompute_ref_log_probs):883 raise ValueError(884 "No reference model and model is not a Peft model. Try setting `precompute_ref_log_probs=True`"885 )886 else:887 if self.is_deepspeed_enabled:888 self.ref_model = self._prepare_deepspeed(self.ref_model)889 else:890 self.ref_model = self.accelerator.prepare_model(self.ref_model, evaluation_mode=True)891 892 def _prepare_deepspeed(self, model: PreTrainedModelWrapper):893 # Adapted from accelerate: https://github.com/huggingface/accelerate/blob/739b135f8367becb67ffaada12fe76e3aa60fefd/src/accelerate/accelerator.py#L1473894 deepspeed_plugin = self.accelerator.state.deepspeed_plugin895 config_kwargs = deepcopy(deepspeed_plugin.deepspeed_config)896 897 if model is not None:898 if hasattr(model, "config"):899 hidden_size = (900 max(model.config.hidden_sizes)901 if getattr(model.config, "hidden_sizes", None)902 else getattr(model.config, "hidden_size", None)903 )904 if hidden_size is not None and config_kwargs["zero_optimization"]["stage"] == 3:905 # Note that `stage3_prefetch_bucket_size` can produce DeepSpeed messages like: `Invalidate trace cache @ step 0: expected module 1, but got module 0`906 # This is expected and is not an error, see: https://github.com/microsoft/DeepSpeed/discussions/4081907 config_kwargs.update(908 {909 "zero_optimization.reduce_bucket_size": hidden_size * hidden_size,910 "zero_optimization.stage3_param_persistence_threshold": 10 * hidden_size,911 "zero_optimization.stage3_prefetch_bucket_size": 0.9 * hidden_size * hidden_size,912 }913 )914 915 # If ZeRO-3 is used, we shard both the active and reference model.916 # Otherwise, we assume the reference model fits in memory and is initialized on each device with ZeRO disabled (stage 0)917 if config_kwargs["zero_optimization"]["stage"] != 3:918 config_kwargs["zero_optimization"]["stage"] = 0919 model, *_ = deepspeed.initialize(model=model, config=config_kwargs)920 model.eval()921 return model922 923 @contextmanager924 def null_ref_context(self):925 """Context manager for handling null reference model (that is, peft adapter manipulation)."""926 with (927 self.accelerator.unwrap_model(self.model).disable_adapter()928 if self.is_peft_model and not self.ref_adapter_name929 else nullcontext()930 ):931 if self.ref_adapter_name:932 self.model.set_adapter(self.ref_adapter_name)933 yield934 if self.ref_adapter_name:935 self.model.set_adapter(self.model_adapter_name or "default")936 937 def get_train_dataloader(self) -> DataLoader:938 """939 Returns the training [`~torch.utils.data.DataLoader`].940 941 Subclass of transformers.src.transformers.trainer.get_train_dataloader to precompute `ref_log_probs`.942 """943 944 if self.precompute_ref_log_probs and not self._precomputed_train_ref_log_probs:945 dataloader_params = {946 "batch_size": self.args.per_device_train_batch_size,947 "collate_fn": self.data_collator,948 "num_workers": self.args.dataloader_num_workers,949 "pin_memory": self.args.dataloader_pin_memory,950 "shuffle": False,951 }952 953 # prepare dataloader954 data_loader = self.accelerator.prepare(DataLoader(self.train_dataset, **dataloader_params))955 reference_completion_logps = []956 reference_KL_logps = []957 958 for padded_batch in tqdm(iterable=data_loader, desc="Train dataset reference log probs"):959 reference_completion_logp, reference_KL_logp = self.compute_reference_log_probs(padded_batch)960 961 reference_completion_logp = self.accelerator.gather_for_metrics(reference_completion_logp)962 reference_completion_logps.append(reference_completion_logp.cpu())963 964 if self.calculate_KL:965 reference_KL_logp = self.accelerator.gather_for_metrics(reference_KL_logp)966 reference_KL_logps.append(reference_KL_logp.cpu())967 968 self.train_dataset = self.train_dataset.add_column(969 name="reference_logps", column=torch.cat(reference_completion_logps).float().numpy()970 )971 972 if self.calculate_KL:973 self.train_dataset = self.train_dataset.add_column(974 name="reference_KL_logps", column=torch.cat(reference_KL_logps).float().numpy()975 )976 977 self._precomputed_train_ref_log_probs = True978 979 return super().get_train_dataloader()980 981 def get_eval_dataloader(self, eval_dataset: Optional[Dataset] = None) -> DataLoader:982 """983 Returns the evaluation [`~torch.utils.data.DataLoader`].984 985 Subclass of transformers.src.transformers.trainer.get_eval_dataloader to precompute `ref_log_probs`.986 987 Args:988 eval_dataset (`torch.utils.data.Dataset`, *optional*):989 If provided, will override `self.eval_dataset`. If it is a [`~datasets.Dataset`], columns not accepted990 by the `model.forward()` method are automatically removed. It must implement `__len__`.991 """992 if eval_dataset is None and self.eval_dataset is None:993 raise ValueError("Trainer: evaluation requires an eval_dataset.")994 eval_dataset = eval_dataset if eval_dataset is not None else self.eval_dataset995 996 if self.precompute_ref_log_probs and not self._precomputed_eval_ref_log_probs:997 dataloader_params = {998 "batch_size": self.args.per_device_eval_batch_size,999 "collate_fn": self.data_collator,1000 "num_workers": self.args.dataloader_num_workers,1001 "pin_memory": self.args.dataloader_pin_memory,1002 "shuffle": False,1003 }1004 1005 # prepare dataloader1006 data_loader = self.accelerator.prepare(DataLoader(eval_dataset, **dataloader_params))1007 1008 reference_completion_logps = []1009 reference_KL_logps = []1010 1011 for padded_batch in tqdm(iterable=data_loader, desc="Eval dataset reference log probs"):1012 reference_completion_logp, reference_KL_logp = self.compute_reference_log_probs(padded_batch)1013 1014 reference_completion_logp = self.accelerator.gather_for_metrics(reference_completion_logp)1015 reference_completion_logps.append(reference_completion_logp.cpu())1016 1017 if self.calculate_KL:1018 reference_KL_logp = self.accelerator.gather_for_metrics(reference_KL_logp)1019 reference_KL_logps.append(reference_KL_logp.cpu())1020 1021 eval_dataset = eval_dataset.add_column(1022 name="reference_logps", column=torch.cat(reference_completion_logps).float().numpy()1023 )1024 if self.calculate_KL:1025 eval_dataset = eval_dataset.add_column(1026 name="reference_KL_logps", column=torch.cat(reference_KL_logps).float().numpy()1027 )1028 1029 # Save calculated reference_chosen_logps and reference_rejected_logps to the eval_dataset for subsequent runs1030 if self.eval_dataset is not None:1031 self.eval_dataset = eval_dataset1032 self._precomputed_eval_ref_log_probs = True1033 1034 return super().get_eval_dataloader(eval_dataset=eval_dataset)1035 1036 def compute_reference_log_probs(self, padded_batch: dict) -> dict:1037 """Computes log probabilities of the reference model for a single padded batch of a KTO specific dataset."""1038 with torch.no_grad():1039 if self.ref_model is None:1040 with self.null_ref_context():1041 if self.is_encoder_decoder:1042 completion_logits = self.model(1043 padded_batch["prompt_input_ids"],1044 attention_mask=padded_batch["prompt_attention_mask"],1045 decoder_input_ids=padded_batch.get("completion_decoder_input_ids"),1046 labels=padded_batch["completion_labels"],1047 ).logits1048 1049 if self.calculate_KL:1050 KL_logits = self.model(1051 padded_batch["KL_prompt_input_ids"],1052 attention_mask=padded_batch["KL_prompt_attention_mask"],1053 decoder_input_ids=padded_batch.get("KL_completion_decoder_input_ids"),1054 labels=padded_batch["KL_completion_labels"],1055 ).logits1056 else:1057 completion_logits = self.model(1058 padded_batch["completion_input_ids"],1059 attention_mask=padded_batch["completion_attention_mask"],1060 ).logits1061 1062 if self.calculate_KL:1063 KL_logits = self.model(1064 padded_batch["KL_completion_input_ids"],1065 attention_mask=padded_batch["KL_completion_attention_mask"],1066 ).logits1067 else:1068 if self.is_encoder_decoder:1069 completion_logits = self.ref_model(1070 padded_batch["prompt_input_ids"],1071 attention_mask=padded_batch["prompt_attention_mask"],1072 decoder_input_ids=padded_batch.get("completion_decoder_input_ids"),1073 labels=padded_batch["completion_labels"],1074 ).logits1075 1076 if self.calculate_KL:1077 KL_logits = self.ref_model(1078 padded_batch["KL_prompt_input_ids"],1079 attention_mask=padded_batch["KL_prompt_attention_mask"],1080 decoder_input_ids=padded_batch.get("KL_completion_decoder_input_ids"),1081 labels=padded_batch["KL_completion_labels"],1082 ).logits1083 else:1084 completion_logits = self.ref_model(1085 padded_batch["completion_input_ids"], attention_mask=padded_batch["completion_attention_mask"]1086 ).logits1087 1088 if self.calculate_KL:1089 KL_logits = self.ref_model(1090 padded_batch["KL_completion_input_ids"],1091 attention_mask=padded_batch["KL_completion_attention_mask"],1092 ).logits1093 1094 completion_logps = self.get_batch_logps(1095 completion_logits,1096 padded_batch["completion_labels"],1097 average_log_prob=False,1098 is_encoder_decoder=self.is_encoder_decoder,1099 label_pad_token_id=self.label_pad_token_id,1100 )1101 1102 if self.calculate_KL:1103 KL_logps = self.get_batch_logps(1104 KL_logits,1105 padded_batch["KL_completion_labels"],1106 average_log_prob=False,1107 is_encoder_decoder=self.is_encoder_decoder,1108 label_pad_token_id=self.label_pad_token_id,1109 )1110 else:1111 KL_logps = None1112 1113 return completion_logps, KL_logps1114 1115 @staticmethod1116 def get_batch_logps(1117 logits: torch.FloatTensor,1118 labels: torch.LongTensor,1119 average_log_prob: bool = False,1120 label_pad_token_id: int = -100,1121 is_encoder_decoder: bool = False,1122 ) -> torch.FloatTensor:1123 """Compute the log probabilities of the given labels under the given logits.1124 1125 Args:1126 logits: Logits of the model (unnormalized). Shape: (batch_size, sequence_length, vocab_size)1127 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)1128 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.1129 1130 Returns:1131 A tensor of shape (batch_size,) containing the average/sum log probabilities of the given labels under the given logits.1132 """1133 if logits.shape[:-1] != labels.shape:1134 raise ValueError("Logits (batch and sequence length dim) and labels must have the same shape.")1135 1136 if not is_encoder_decoder:1137 labels = labels[:, 1:].clone()1138 logits = logits[:, :-1, :]1139 else:1140 # Fixes end-dec RuntimeError1141 labels = labels.clone()1142 1143 loss_mask = labels != label_pad_token_id1144 1145 # dummy token; we'll ignore the losses on these tokens later1146 labels[labels == label_pad_token_id] = 01147 1148 per_token_logps = selective_log_softmax(logits, labels)1149 1150 if average_log_prob:1151 return (per_token_logps * loss_mask).sum(-1) / loss_mask.sum(-1)1152 else:1153 return (per_token_logps * loss_mask).sum(-1)1154 1155 def forward(1156 self, model: nn.Module, batch: dict[str, Union[list, torch.LongTensor]]1157 ) -> tuple[torch.FloatTensor, torch.FloatTensor, torch.FloatTensor, torch.FloatTensor]:1158 if self.calculate_KL:1159 KL_logps = None1160 KL_model_kwargs = (1161 {1162 "input_ids": batch["KL_prompt_input_ids"],1163 "attention_mask": batch["KL_prompt_attention_mask"],1164 "labels": batch["KL_completion_labels"],1165 "decoder_input_ids": batch.get("KL_completion_decoder_input_ids"),1166 }1167 if self.is_encoder_decoder1168 else {1169 "input_ids": batch["KL_completion_input_ids"],1170 "attention_mask": batch["KL_completion_attention_mask"],1171 }1172 )1173 with torch.no_grad():1174 KL_logits = model(1175 **KL_model_kwargs,1176 ).logits1177 1178 KL_logps = self.get_batch_logps(1179 KL_logits,1180 batch["KL_completion_labels"],1181 average_log_prob=False,1182 is_encoder_decoder=self.is_encoder_decoder,1183 label_pad_token_id=self.label_pad_token_id,1184 )1185 else:1186 KL_logps = None1187 1188 model_kwargs = (1189 {1190 "labels": batch["completion_labels"],1191 "decoder_input_ids": batch.get("completion_decoder_input_ids"),1192 }1193 if self.is_encoder_decoder1194 else {}1195 )1196 if self.aux_loss_enabled:1197 model_kwargs["output_router_logits"] = True1198 1199 outputs = model(1200 batch["completion_input_ids"],