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.bco_trainer import (Any, AutoModelForCausalLM, BCOConfig, BCOTrainer, BaseImageProcessor, CLF_NAME, Callable, DPODataCollatorWithPadding, DataCollator, DataLoader, Dataset, EvalLoopOutput, F, FeatureExtractionMixin, Literal, Optional, PartialState, PeftModel, PreTrainedModel, PreTrainedModelWrapper, PreTrainedTokenizerBase, ProcessorMixin, RUNNING_NAME, RunningMoments, SequentialSampler, Trainer, TrainerCallback, TrainingArguments, Union, _process_tokens, _tokenize, amp, 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_sklearn_available, is_wandb_available, itemgetter, log_table_to_comet_experiment, maybe_apply_chat_template, 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, F, Optional, PeftModel, PreTrainedModel, Trainer, is_peft_available, os, torch)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 UnslothBCOConfig(BCOConfig):44 """45 46 Configuration class for the [`BCOTrainer`].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 max_length (`int` or `None`, *optional*, defaults to `1024`):54 Maximum length of the sequences (prompt + completion) in the batch. This argument is required if you want55 to use the default data collator.56 max_prompt_length (`int` or `None`, *optional*, defaults to `512`):57 Maximum length of the prompt. This argument is required if you want to use the default data collator.58 max_completion_length (`int` or `None`, *optional*, defaults to `None`):59 Maximum length of the completion. This argument is required if you want to use the default data collator60 and your model is an encoder-decoder.61 beta (`float`, *optional*, defaults to `0.1`):62 Parameter controlling the deviation from the reference model. Higher β means less deviation from the63 reference model.64 label_pad_token_id (`int`, *optional*, defaults to `-100`):65 Label pad token id. This argument is required if you want to use the default data collator.66 padding_value (`int` or `None`, *optional*, defaults to `None`):67 Padding value to use. If `None`, the padding value of the tokenizer is used.68 truncation_mode (`str`, *optional*, defaults to `"keep_end"`):69 Truncation mode to use when the prompt is too long. Possible values are `"keep_end"` or `"keep_start"`.70 This argument is required if you want to use the default data collator.71 disable_dropout (`bool`, *optional*, defaults to `True`):72 Whether to disable dropout in the model and reference model.73 generate_during_eval (`bool`, *optional*, defaults to `False`):74 If `True`, generates and logs completions from both the model and the reference model to W&B or Comet during75 evaluation.76 is_encoder_decoder (`bool` or `None`, *optional*, defaults to `None`):77 When using the `model_init` argument (callable) to instantiate the model instead of the `model` argument,78 you need to specify if the model returned by the callable is an encoder-decoder model.79 precompute_ref_log_probs (`bool`, *optional*, defaults to `False`):80 Whether to precompute reference model log probabilities for training and evaluation datasets. This is81 useful when training without the reference model to reduce the total GPU memory needed.82 model_init_kwargs (`dict[str, Any]` or `None`, *optional*, defaults to `None`):83 Keyword arguments to pass to `AutoModelForCausalLM.from_pretrained` when instantiating the model from a84 string.85 ref_model_init_kwargs (`dict[str, Any]` or `None`, *optional*, defaults to `None`):86 Keyword arguments to pass to `AutoModelForCausalLM.from_pretrained` when instantiating the reference model87 from a string.88 dataset_num_proc (`int` or `None`, *optional*, defaults to `None`):89 Number of processes to use for processing the dataset.90 prompt_sample_size (`int`, *optional*, defaults to `1024`):91 Number of prompts that are fed to density ratio classifier.92 min_density_ratio (`float`, *optional*, defaults to `0.5`):93 Minimum value of the density ratio. The estimated density ratio is clamped to this value.94 max_density_ratio (`float`, *optional*, defaults to `10.0`):95 Maximum value of the density ratio. The estimated density ratio is clamped to this value.96 97 """98 vllm_sampling_params: Optional[Any] = field(99 default = None,100 metadata = {'help': 'vLLM SamplingParams'},101 )102 unsloth_num_chunks : Optional[int] = field(103 default = -1,104 metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},105 )106 def __init__(107 self,108 output_dir = None,109 overwrite_output_dir = None,110 do_train = False,111 do_eval = False,112 do_predict = False,113 eval_strategy = 'no',114 prediction_loss_only = False,115 per_device_train_batch_size = 4,116 per_device_eval_batch_size = 4,117 per_gpu_train_batch_size = None,118 per_gpu_eval_batch_size = None,119 gradient_accumulation_steps = 2,120 eval_accumulation_steps = 2,121 eval_delay = 0,122 torch_empty_cache_steps = 250,123 learning_rate = 5e-05,124 weight_decay = 0.01,125 adam_beta1 = 0.9,126 adam_beta2 = 0.999,127 adam_epsilon = 1e-08,128 max_grad_norm = 1.0,129 num_train_epochs = 3.0,130 max_steps = -1,131 lr_scheduler_type = 'linear',132 warmup_ratio = 0.1,133 warmup_steps = 0,134 log_level = 'passive',135 log_level_replica = 'warning',136 log_on_each_node = True,137 logging_dir = None,138 logging_strategy = 'steps',139 logging_first_step = False,140 logging_steps = 1,141 logging_nan_inf_filter = False,142 save_strategy = 'steps',143 save_steps = 500,144 save_total_limit = None,145 save_safetensors = True,146 save_on_each_node = False,147 save_only_model = False,148 restore_callback_states_from_checkpoint = False,149 no_cuda = False,150 use_cpu = False,151 use_mps_device = False,152 seed = 3407,153 data_seed = 3407,154 jit_mode_eval = False,155 use_ipex = False,156 bf16 = False,157 fp16 = False,158 fp16_opt_level = 'O1',159 half_precision_backend = 'auto',160 bf16_full_eval = False,161 fp16_full_eval = False,162 tf32 = None,163 local_rank = -1,164 ddp_backend = None,165 tpu_num_cores = None,166 tpu_metrics_debug = False,167 debug = '',168 dataloader_drop_last = False,169 eval_steps = None,170 dataloader_num_workers = 0,171 dataloader_prefetch_factor = None,172 past_index = -1,173 run_name = None,174 disable_tqdm = None,175 remove_unused_columns = True,176 label_names = None,177 load_best_model_at_end = False,178 metric_for_best_model = None,179 greater_is_better = None,180 ignore_data_skip = False,181 fsdp = '',182 fsdp_min_num_params = 0,183 fsdp_config = None,184 tp_size = 0,185 fsdp_transformer_layer_cls_to_wrap = None,186 accelerator_config = None,187 deepspeed = None,188 label_smoothing_factor = 0.0,189 optim = 'adamw_8bit',190 optim_args = None,191 adafactor = False,192 group_by_length = False,193 length_column_name = 'length',194 report_to = None,195 ddp_find_unused_parameters = None,196 ddp_bucket_cap_mb = None,197 ddp_broadcast_buffers = None,198 dataloader_pin_memory = True,199 dataloader_persistent_workers = False,200 skip_memory_metrics = True,201 use_legacy_prediction_loop = False,202 push_to_hub = False,203 resume_from_checkpoint = None,204 hub_model_id = None,205 hub_strategy = 'every_save',206 hub_token = None,207 hub_private_repo = None,208 hub_always_push = False,209 gradient_checkpointing = False,210 gradient_checkpointing_kwargs = None,211 include_inputs_for_metrics = False,212 eval_do_concat_batches = True,213 fp16_backend = 'auto',214 evaluation_strategy = None,215 push_to_hub_model_id = None,216 push_to_hub_organization = None,217 push_to_hub_token = None,218 mp_parameters = '',219 auto_find_batch_size = False,220 full_determinism = False,221 torchdynamo = None,222 ray_scope = 'last',223 ddp_timeout = 1800,224 torch_compile = False,225 torch_compile_backend = None,226 torch_compile_mode = None,227 dispatch_batches = None,228 split_batches = None,229 include_tokens_per_second = False,230 include_num_input_tokens_seen = False,231 neftune_noise_alpha = None,232 optim_target_modules = None,233 batch_eval_metrics = False,234 eval_on_start = False,235 use_liger_kernel = False,236 eval_use_gather_object = False,237 average_tokens_across_devices = False,238 max_length = 1024,239 max_prompt_length = 512,240 max_completion_length = None,241 beta = 0.1,242 label_pad_token_id = -100,243 padding_value = None,244 truncation_mode = 'keep_end',245 disable_dropout = True,246 generate_during_eval = False,247 is_encoder_decoder = None,248 precompute_ref_log_probs = False,249 model_init_kwargs = None,250 ref_model_init_kwargs = None,251 dataset_num_proc = None,252 prompt_sample_size = 1024,253 min_density_ratio = 0.5,254 max_density_ratio = 10.0,255 vllm_sampling_params = None,256 unsloth_num_chunks = -1,257 **kwargs,258 ):259 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!')260 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!')261 if output_dir is None and save_strategy == 'steps' and save_steps == 500:262 output_dir = 'unsloth_training_checkpoints'263 save_strategy = 'no'264 if dataset_num_proc is None:265 from multiprocessing import cpu_count266 dataset_num_proc = cpu_count()267 268 super().__init__(269 output_dir = output_dir,270 overwrite_output_dir = overwrite_output_dir,271 do_train = do_train,272 do_eval = do_eval,273 do_predict = do_predict,274 eval_strategy = eval_strategy,275 prediction_loss_only = prediction_loss_only,276 per_device_train_batch_size = per_device_train_batch_size,277 per_device_eval_batch_size = per_device_eval_batch_size,278 per_gpu_train_batch_size = per_gpu_train_batch_size,279 per_gpu_eval_batch_size = per_gpu_eval_batch_size,280 gradient_accumulation_steps = gradient_accumulation_steps,281 eval_accumulation_steps = eval_accumulation_steps,282 eval_delay = eval_delay,283 torch_empty_cache_steps = torch_empty_cache_steps,284 learning_rate = learning_rate,285 weight_decay = weight_decay,286 adam_beta1 = adam_beta1,287 adam_beta2 = adam_beta2,288 adam_epsilon = adam_epsilon,289 max_grad_norm = max_grad_norm,290 num_train_epochs = num_train_epochs,291 max_steps = max_steps,292 lr_scheduler_type = lr_scheduler_type,293 warmup_ratio = warmup_ratio,294 warmup_steps = warmup_steps,295 log_level = log_level,296 log_level_replica = log_level_replica,297 log_on_each_node = log_on_each_node,298 logging_dir = logging_dir,299 logging_strategy = logging_strategy,300 logging_first_step = logging_first_step,301 logging_steps = logging_steps,302 logging_nan_inf_filter = logging_nan_inf_filter,303 save_strategy = save_strategy,304 save_steps = save_steps,305 save_total_limit = save_total_limit,306 save_safetensors = save_safetensors,307 save_on_each_node = save_on_each_node,308 save_only_model = save_only_model,309 restore_callback_states_from_checkpoint = restore_callback_states_from_checkpoint,310 no_cuda = no_cuda,311 use_cpu = use_cpu,312 use_mps_device = use_mps_device,313 seed = seed,314 data_seed = data_seed,315 jit_mode_eval = jit_mode_eval,316 use_ipex = use_ipex,317 bf16 = bf16,318 fp16 = fp16,319 fp16_opt_level = fp16_opt_level,320 half_precision_backend = half_precision_backend,321 bf16_full_eval = bf16_full_eval,322 fp16_full_eval = fp16_full_eval,323 tf32 = tf32,324 local_rank = local_rank,325 ddp_backend = ddp_backend,326 tpu_num_cores = tpu_num_cores,327 tpu_metrics_debug = tpu_metrics_debug,328 debug = debug,329 dataloader_drop_last = dataloader_drop_last,330 eval_steps = eval_steps,331 dataloader_num_workers = dataloader_num_workers,332 dataloader_prefetch_factor = dataloader_prefetch_factor,333 past_index = past_index,334 run_name = run_name,335 disable_tqdm = disable_tqdm,336 remove_unused_columns = remove_unused_columns,337 label_names = label_names,338 load_best_model_at_end = load_best_model_at_end,339 metric_for_best_model = metric_for_best_model,340 greater_is_better = greater_is_better,341 ignore_data_skip = ignore_data_skip,342 fsdp = fsdp,343 fsdp_min_num_params = fsdp_min_num_params,344 fsdp_config = fsdp_config,345 tp_size = tp_size,346 fsdp_transformer_layer_cls_to_wrap = fsdp_transformer_layer_cls_to_wrap,347 accelerator_config = accelerator_config,348 deepspeed = deepspeed,349 label_smoothing_factor = label_smoothing_factor,350 optim = optim,351 optim_args = optim_args,352 adafactor = adafactor,353 group_by_length = group_by_length,354 length_column_name = length_column_name,355 report_to = report_to,356 ddp_find_unused_parameters = ddp_find_unused_parameters,357 ddp_bucket_cap_mb = ddp_bucket_cap_mb,358 ddp_broadcast_buffers = ddp_broadcast_buffers,359 dataloader_pin_memory = dataloader_pin_memory,360 dataloader_persistent_workers = dataloader_persistent_workers,361 skip_memory_metrics = skip_memory_metrics,362 use_legacy_prediction_loop = use_legacy_prediction_loop,363 push_to_hub = push_to_hub,364 resume_from_checkpoint = resume_from_checkpoint,365 hub_model_id = hub_model_id,366 hub_strategy = hub_strategy,367 hub_token = hub_token,368 hub_private_repo = hub_private_repo,369 hub_always_push = hub_always_push,370 gradient_checkpointing = gradient_checkpointing,371 gradient_checkpointing_kwargs = gradient_checkpointing_kwargs,372 include_inputs_for_metrics = include_inputs_for_metrics,373 eval_do_concat_batches = eval_do_concat_batches,374 fp16_backend = fp16_backend,375 evaluation_strategy = evaluation_strategy,376 push_to_hub_model_id = push_to_hub_model_id,377 push_to_hub_organization = push_to_hub_organization,378 push_to_hub_token = push_to_hub_token,379 mp_parameters = mp_parameters,380 auto_find_batch_size = auto_find_batch_size,381 full_determinism = full_determinism,382 torchdynamo = torchdynamo,383 ray_scope = ray_scope,384 ddp_timeout = ddp_timeout,385 torch_compile = torch_compile,386 torch_compile_backend = torch_compile_backend,387 torch_compile_mode = torch_compile_mode,388 dispatch_batches = dispatch_batches,389 split_batches = split_batches,390 include_tokens_per_second = include_tokens_per_second,391 include_num_input_tokens_seen = include_num_input_tokens_seen,392 neftune_noise_alpha = neftune_noise_alpha,393 optim_target_modules = optim_target_modules,394 batch_eval_metrics = batch_eval_metrics,395 eval_on_start = eval_on_start,396 use_liger_kernel = use_liger_kernel,397 eval_use_gather_object = eval_use_gather_object,398 average_tokens_across_devices = average_tokens_across_devices,399 max_length = max_length,400 max_prompt_length = max_prompt_length,401 max_completion_length = max_completion_length,402 beta = beta,403 label_pad_token_id = label_pad_token_id,404 padding_value = padding_value,405 truncation_mode = truncation_mode,406 disable_dropout = disable_dropout,407 generate_during_eval = generate_during_eval,408 is_encoder_decoder = is_encoder_decoder,409 precompute_ref_log_probs = precompute_ref_log_probs,410 model_init_kwargs = model_init_kwargs,411 ref_model_init_kwargs = ref_model_init_kwargs,412 dataset_num_proc = dataset_num_proc,413 prompt_sample_size = prompt_sample_size,414 min_density_ratio = min_density_ratio,415 max_density_ratio = max_density_ratio,**kwargs)416 self.vllm_sampling_params = vllm_sampling_params417 self.unsloth_num_chunks = unsloth_num_chunks418pass419 420class _UnslothBCOTrainer(Trainer):421 r""""""422 423 _tag_names = ["trl", "bco"]424 425 def __init__(426 self,427 model: Union[PreTrainedModel, nn.Module, str] = None,428 ref_model: Optional[Union[PreTrainedModel, nn.Module, str]] = None,429 args: BCOConfig = None,430 train_dataset: Optional[Dataset] = None,431 eval_dataset: Optional[Union[Dataset, dict[str, Dataset]]] = None,432 processing_class: Optional[433 Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin]434 ] = None,435 data_collator: Optional[DataCollator] = None,436 model_init: Optional[Callable[[], PreTrainedModel]] = None,437 callbacks: Optional[list[TrainerCallback]] = None,438 optimizers: tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),439 preprocess_logits_for_metrics: Optional[Callable[[torch.Tensor, torch.Tensor], torch.Tensor]] = None,440 peft_config: Optional[dict] = None,441 compute_metrics: Optional[Callable[[EvalLoopOutput], dict]] = None,442 model_adapter_name: Optional[str] = None,443 ref_adapter_name: Optional[str] = None,444 embedding_func: Optional[Callable] = None,445 embedding_tokenizer: Optional[PreTrainedTokenizerBase] = None,446 ):447 if not is_sklearn_available():448 raise ImportError(449 "BCOTrainer requires the scikit-learn library. Please install it with `pip install scikit-learn`."450 )451 452 if type(args) is TrainingArguments:453 raise ValueError("Please use `BCOConfig` 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 BCOTrainer. 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 BCOConfig. 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 BCOTrainer. 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 BCOConfig. 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 `BCOConfig`. "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 `BCOConfig`. "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 BCOTrainer'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 BCOConfig"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.max_length = max_length648 self.generate_during_eval = args.generate_during_eval649 self.label_pad_token_id = args.label_pad_token_id650 self.padding_value = args.padding_value if args.padding_value is not None else processing_class.pad_token_id651 self.max_prompt_length = max_prompt_length652 self.truncation_mode = args.truncation_mode653 self.max_completion_length = max_completion_length654 self.precompute_ref_log_probs = args.precompute_ref_log_probs655 656 # Since ref_logs are precomputed on the first call to get_train/eval_dataloader657 # keep track of first called to avoid computation of future calls658 self._precomputed_train_ref_log_probs = False659 self._precomputed_eval_ref_log_probs = False660 661 # metric662 self._stored_metrics = defaultdict(lambda: defaultdict(list))663 664 # BCO parameter665 self.beta = args.beta666 self.aux_loss_enabled = getattr(model.config, "output_router_logits", False)667 self.aux_loss_coef = getattr(model.config, "router_aux_loss_coef", 0.0)668 if self.aux_loss_enabled and self.aux_loss_coef == 0.0:669 warnings.warn(670 "You set `output_router_logits` to `True` in the model config, but `router_aux_loss_coef` is set to "671 "`0.0`, meaning the auxiliary loss will not be used. Either set `router_aux_loss_coef` to a value "672 "greater than `0.0`, or set `output_router_logits` to `False` if you don't want to use the auxiliary "673 "loss.",674 UserWarning,675 )676 677 # Underlying Distribution Matching argument678 self.embedding_func = embedding_func679 self.embedding_tokenizer = embedding_tokenizer680 681 # The trainer estimates the number of FLOPs (floating-point operations) using the number of elements in the682 # input tensor associated with the key "input_ids". However, in BCO, the sampled data does not include the683 # "input_ids" key. Instead, the available keys are "prompt_input_ids" and "completion_input_ids". As a result,684 # the trainer issues the warning: "Could not estimate the number of tokens of the input, floating-point685 # operations will not be computed." To suppress this warning, we set the "estimate_tokens" key in the model's686 # "warnings_issued" dictionary to True. This acts as a flag to indicate that the warning has already been687 # issued.688 model.warnings_issued["estimate_tokens"] = True689 690 with PartialState().local_main_process_first():691 # Apply the chat template if needed692 train_dataset = train_dataset.map(693 maybe_apply_chat_template, fn_kwargs={"tokenizer": processing_class}, num_proc=args.dataset_num_proc694 )695 if eval_dataset is not None:696 eval_dataset = eval_dataset.map(697 maybe_apply_chat_template,698 fn_kwargs={"tokenizer": processing_class},699 num_proc=args.dataset_num_proc,700 )701 # Shuffle the datasets702 train_dataset = train_dataset.shuffle(seed=args.data_seed)703 if eval_dataset is not None:704 eval_dataset = eval_dataset.shuffle(seed=args.data_seed)705 # Tokenize and prepare the training datasets706 train_dataset = train_dataset.map(707 _tokenize,708 batched=True,709 fn_kwargs={"tokenizer": processing_class, "embedding_tokenizer": self.embedding_tokenizer},710 num_proc=args.dataset_num_proc,711 desc="Tokenizing train dataset",712 )713 714 # Prepare the datasets715 fn_kwargs = {716 "prefix": "",717 "is_encoder_decoder": self.is_encoder_decoder,718 "tokenizer": processing_class,719 "max_length": self.max_length,720 "truncation_mode": self.truncation_mode,721 "label_pad_token_id": self.label_pad_token_id,722 "max_prompt_length": self.max_prompt_length,723 "max_completion_length": self.max_completion_length,724 }725 train_dataset = train_dataset.map(726 _process_tokens,727 fn_kwargs=fn_kwargs,728 num_proc=args.dataset_num_proc,729 desc="Processing tokenized train dataset",730 )731 732 if eval_dataset is not None:733 # Tokenize734 eval_dataset = eval_dataset.map(735 _tokenize,736 fn_kwargs={"tokenizer": processing_class, "embedding_tokenizer": self.embedding_tokenizer},737 batched=True,738 num_proc=args.dataset_num_proc,739 desc="Tokenizing eval dataset",740 )741 742 # Process743 fn_kwargs = {744 "prefix": "",745 "is_encoder_decoder": self.is_encoder_decoder,746 "tokenizer": processing_class,747 "max_length": self.max_length,748 "truncation_mode": self.truncation_mode,749 "label_pad_token_id": self.label_pad_token_id,750 "max_prompt_length": self.max_prompt_length,751 "max_completion_length": self.max_completion_length,752 }753 eval_dataset = eval_dataset.map(754 _process_tokens,755 fn_kwargs=fn_kwargs,756 num_proc=args.dataset_num_proc,757 desc="Processing tokenized eval dataset",758 )759 760 desirable = train_dataset.filter(761 lambda x: x["label"], num_proc=args.dataset_num_proc, desc="Filtering desirable examples"762 )763 undesirable = train_dataset.filter(764 lambda x: not x["label"], num_proc=args.dataset_num_proc, desc="Filtering undesirable examples"765 )766 767 desirable = desirable.shuffle(seed=args.data_seed)768 undesirable = undesirable.shuffle(seed=args.data_seed)769 770 super().__init__(771 model=model,772 args=args,773 data_collator=data_collator,774 train_dataset=train_dataset,775 eval_dataset=eval_dataset,776 processing_class=processing_class,777 model_init=model_init,778 compute_metrics=compute_metrics,779 callbacks=callbacks,780 optimizers=optimizers,781 preprocess_logits_for_metrics=preprocess_logits_for_metrics,782 )783 784 # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the785 # model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set786 # self.model_accepts_loss_kwargs to False to enable scaling.787 self.model_accepts_loss_kwargs = False788 789 # Add tags for models that have been loaded with the correct transformers version790 if hasattr(self.model, "add_model_tags"):791 self.model.add_model_tags(self._tag_names)792 793 if not hasattr(self, "accelerator"):794 raise AttributeError(795 "Your `Trainer` does not have an `accelerator` object. Consider upgrading `transformers`."796 )797 798 # Deepspeed Zero-3 does not support precompute_ref_log_probs799 if self.is_deepspeed_enabled:800 if self.accelerator.state.deepspeed_plugin.zero_stage == 3 and self.precompute_ref_log_probs:801 raise ValueError(802 "You cannot use `precompute_ref_log_probs=True` with Deepspeed ZeRO-3. Please set `precompute_ref_log_probs=False`."803 )804 805 if self.ref_model is None:806 if not (self.is_peft_model or self.precompute_ref_log_probs):807 raise ValueError(808 "No reference model and model is not a Peft model. Try setting `precompute_ref_log_probs=True`"809 )810 else:811 if self.is_deepspeed_enabled:812 self.ref_model = self._prepare_deepspeed(self.ref_model)813 else:814 self.ref_model = self.accelerator.prepare_model(self.ref_model, evaluation_mode=True)815 816 self.running = RunningMoments(accelerator=self.accelerator)817 818 if self.embedding_func is None:819 return820 821 chosen_embeddings = self._get_sample_prompt_embeddings(desirable, sample_size=self.args.prompt_sample_size)822 rejected_embeddings = self._get_sample_prompt_embeddings(undesirable, sample_size=self.args.prompt_sample_size)823 824 embeddings = torch.cat((chosen_embeddings, rejected_embeddings), dim=0)825 labels = torch.cat(826 (torch.ones_like(chosen_embeddings[:, 0]), torch.zeros_like(rejected_embeddings[:, 0])), dim=0827 )828 829 self.clf = LogisticRegression(class_weight="balanced").fit(830 embeddings.cpu().float().numpy(), labels.cpu().numpy()831 )832 833 @property834 def match_underlying_distribution(self):835 return self.embedding_func is not None and self.embedding_tokenizer is not None836 837 def _get_chosen_prob(self, prompt_embeddings: torch.FloatTensor) -> torch.FloatTensor:838 """839 Calculates the probability if the given prompt embedding is from desirable dataset.840 This function calculates the probability in the process and ensemble across processes.841 """842 dtype = prompt_embeddings.dtype843 device = prompt_embeddings.device844 rank = self.accelerator.process_index845 846 padded_prompt_embeddings = self.accelerator.pad_across_processes(847 prompt_embeddings, pad_index=self.embedding_tokenizer.pad_token_id848 )849 sample_size = padded_prompt_embeddings.shape[0]850 nonzero = padded_prompt_embeddings.mean(dim=1) != self.embedding_tokenizer.pad_token_id851 prompt_embeddings = self.accelerator.gather(padded_prompt_embeddings)852 853 # cannot predict for all empty values854 if prompt_embeddings.shape[0] == 0:855 return torch.tensor([], device=device, dtype=dtype)856 857 prob = self.clf.predict_proba(prompt_embeddings.cpu().float().numpy())[:, 1]858 prob = torch.as_tensor(prob, dtype=dtype, device=device)859 prob = self.accelerator.reduce(prob, reduction="mean")860 861 prob = prob[sample_size * rank : sample_size * (rank + 1)]862 prob = prob[nonzero]863 864 return prob865 866 def _vectorize_prompt(self, input_ids: torch.LongTensor, attention_mask: torch.LongTensor) -> torch.FloatTensor:867 """868 Replaces processing_class.pad_token_id to embedding_tokenizer.pad_token_id869 and applies self.embedding_func870 """871 input_ids = torch.where(872 input_ids == self.processing_class.pad_token_id,873 self.embedding_tokenizer.pad_token_id,874 input_ids,875 )876 877 with torch.no_grad():878 embeddings = self.embedding_func(879 input_ids=input_ids,880 attention_mask=attention_mask,881 )882 883 return embeddings884 885 def _get_prompt_embeddings(886 self, batch: dict[str, Union[list, torch.LongTensor]]887 ) -> tuple[torch.FloatTensor, torch.FloatTensor]:888 """Extract embeddings from frozen embedding model"""889 890 if not self.match_underlying_distribution:891 return None, None892 893 embeddings = self._vectorize_prompt(894 input_ids=batch["embedding_input_ids"],895 attention_mask=batch["embedding_attention_mask"],896 )897 898 chosen_idx = [i for i in range(len(batch["label"])) if batch["label"][i] is True]899 rejected_idx = [i for i in range(len(batch["label"])) if batch["label"][i] is False]900 901 chosen_embeddings = embeddings[chosen_idx, ...]902 rejected_embeddings = embeddings[rejected_idx, ...]903 904 return (chosen_embeddings, rejected_embeddings)905 906 def _get_sample_prompt_embeddings(self, dataset: Dataset, sample_size: int = 512) -> torch.FloatTensor:907 """908 Sample instances from dataset and get prompt embeddings.909 Used for density ratio classifier training.910 """911 n_samples = min(len(dataset), sample_size)912 rand_indices = np.random.choice(len(dataset), size=(n_samples,))913 914 embedding_dataset = dataset.select(rand_indices)915 916 dataloader_params = {917 "batch_size": self.args.per_device_train_batch_size,918 "collate_fn": self.data_collator,919 "num_workers": self.args.dataloader_num_workers,920 "pin_memory": self.args.dataloader_pin_memory,921 "shuffle": False,922 }923 924 # prepare dataloader925 data_loader = self.accelerator.prepare(DataLoader(embedding_dataset, **dataloader_params))926 927 with torch.no_grad():928 all_embeddings = torch.empty(0)929 for padded_batch in tqdm(iterable=data_loader, desc="Building sample prompt embeddings"):930 embeddings = self._vectorize_prompt(931 input_ids=padded_batch["embedding_input_ids"],932 attention_mask=padded_batch["embedding_attention_mask"],933 )934 embeddings = self.accelerator.gather_for_metrics(embeddings)935 all_embeddings = torch.cat((all_embeddings, embeddings.cpu()))936 937 return all_embeddings938 939 def _prepare_deepspeed(self, model: PreTrainedModelWrapper):940 # Adapted from accelerate: https://github.com/huggingface/accelerate/blob/739b135f8367becb67ffaada12fe76e3aa60fefd/src/accelerate/accelerator.py#L1473941 deepspeed_plugin = self.accelerator.state.deepspeed_plugin942 config_kwargs = deepcopy(deepspeed_plugin.deepspeed_config)943 944 if model is not None:945 if hasattr(model, "config"):946 hidden_size = (947 max(model.config.hidden_sizes)948 if getattr(model.config, "hidden_sizes", None)949 else getattr(model.config, "hidden_size", None)950 )951 if hidden_size is not None and config_kwargs["zero_optimization"]["stage"] == 3:952 # Note that `stage3_prefetch_bucket_size` can produce DeepSpeed messages like: `Invalidate trace cache @ step 0: expected module 1, but got module 0`953 # This is expected and is not an error, see: https://github.com/microsoft/DeepSpeed/discussions/4081954 config_kwargs.update(955 {956 "zero_optimization.reduce_bucket_size": hidden_size * hidden_size,957 "zero_optimization.stage3_param_persistence_threshold": 10 * hidden_size,958 "zero_optimization.stage3_prefetch_bucket_size": 0.9 * hidden_size * hidden_size,959 }960 )961 962 # If ZeRO-3 is used, we shard both the active and reference model.963 # Otherwise, we assume the reference model fits in memory and is initialized on each device with ZeRO disabled (stage 0)964 if config_kwargs["zero_optimization"]["stage"] != 3:965 config_kwargs["zero_optimization"]["stage"] = 0966 model, *_ = deepspeed.initialize(model=model, config=config_kwargs)967 model.eval()968 return model969 970 def _save_optimizer_and_scheduler(self, output_dir):971 super()._save_optimizer_and_scheduler(output_dir)972 973 # When saving optimizer and scheduler to checkpoint, save also the running delta object.974 output_dir = output_dir if output_dir is not None else self.args.output_dir975 976 self.running.save_to_json(os.path.join(output_dir, RUNNING_NAME))977 978 if self.match_underlying_distribution:979 torch.save(self.clf.get_params(), os.path.join(output_dir, CLF_NAME))980 981 def _load_optimizer_and_scheduler(self, checkpoint):982 super()._load_optimizer_and_scheduler(checkpoint)983 984 if checkpoint is None:985 return986 # when loading optimizer and scheduler from checkpoint, also load the running delta object.987 running_file = os.path.join(checkpoint, RUNNING_NAME)988 if os.path.isfile(running_file):989 self.running = RunningMoments.load_from_json(self.accelerator, running_file)990 991 if self.match_underlying_distribution:992 clf_file = os.path.join(checkpoint, CLF_NAME)993 if os.path.isfile(running_file):994 self.clf.set_params(**torch.load(clf_file, weights_only=True, map_location="cpu"))995 996 @contextmanager997 def null_ref_context(self):998 """Context manager for handling null reference model (that is, peft adapter manipulation)."""999 with (1000 self.accelerator.unwrap_model(self.model).disable_adapter()1001 if self.is_peft_model and not self.ref_adapter_name1002 else nullcontext()1003 ):1004 if self.ref_adapter_name:1005 self.model.set_adapter(self.ref_adapter_name)1006 yield1007 if self.ref_adapter_name:1008 self.model.set_adapter(self.model_adapter_name or "default")1009 1010 def get_train_dataloader(self) -> DataLoader:1011 """1012 Returns the training [`~torch.utils.data.DataLoader`].1013 1014 Subclass of transformers.src.transformers.trainer.get_train_dataloader to precompute `ref_log_probs`.1015 """1016 1017 if self.precompute_ref_log_probs and not self._precomputed_train_ref_log_probs:1018 dataloader_params = {1019 "batch_size": self.args.per_device_train_batch_size,1020 "collate_fn": self.data_collator,1021 "num_workers": self.args.dataloader_num_workers,1022 "pin_memory": self.args.dataloader_pin_memory,1023 "shuffle": False,1024 }1025 1026 # prepare dataloader1027 data_loader = self.accelerator.prepare(DataLoader(self.train_dataset, **dataloader_params))1028 reference_completion_logps = []1029 1030 for padded_batch in tqdm(iterable=data_loader, desc="Train dataset reference log probs"):1031 reference_completion_logp = self.compute_reference_log_probs(padded_batch)1032 1033 reference_completion_logp = self.accelerator.gather_for_metrics(reference_completion_logp)1034 reference_completion_logps.append(reference_completion_logp.cpu())1035 1036 self.train_dataset = self.train_dataset.add_column(1037 name="reference_logps", column=torch.cat(reference_completion_logps).float().numpy()1038 )1039 1040 self._precomputed_train_ref_log_probs = True1041 1042 return super().get_train_dataloader()1043 1044 def get_eval_dataloader(self, eval_dataset: Optional[Dataset] = None) -> DataLoader:1045 """1046 Returns the evaluation [`~torch.utils.data.DataLoader`].1047 1048 Subclass of transformers.src.transformers.trainer.get_eval_dataloader to precompute `ref_log_probs`.1049 1050 Args:1051 eval_dataset (`torch.utils.data.Dataset`, *optional*):1052 If provided, will override `self.eval_dataset`. If it is a [`~datasets.Dataset`], columns not accepted1053 by the `model.forward()` method are automatically removed. It must implement `__len__`.1054 """1055 if eval_dataset is None and self.eval_dataset is None:1056 raise ValueError("Trainer: evaluation requires an eval_dataset.")1057 eval_dataset = eval_dataset if eval_dataset is not None else self.eval_dataset1058 1059 if self.precompute_ref_log_probs and not self._precomputed_eval_ref_log_probs:1060 dataloader_params = {1061 "batch_size": self.args.per_device_eval_batch_size,1062 "collate_fn": self.data_collator,1063 "num_workers": self.args.dataloader_num_workers,1064 "pin_memory": self.args.dataloader_pin_memory,1065 "shuffle": False,1066 }1067 1068 # prepare dataloader1069 data_loader = self.accelerator.prepare(DataLoader(eval_dataset, **dataloader_params))1070 1071 reference_completion_logps = []1072 1073 for padded_batch in tqdm(iterable=data_loader, desc="Eval dataset reference log probs"):1074 reference_completion_logp = self.compute_reference_log_probs(padded_batch)1075 1076 reference_completion_logp = self.accelerator.gather_for_metrics(reference_completion_logp)1077 reference_completion_logps.append(reference_completion_logp.cpu())1078 1079 eval_dataset = eval_dataset.add_column(1080 name="reference_logps", column=torch.cat(reference_completion_logps).float().numpy()1081 )1082 1083 # Save calculated reference_chosen_logps and reference_rejected_logps to the eval_dataset for subsequent runs1084 if self.eval_dataset is not None:1085 self.eval_dataset = eval_dataset1086 self._precomputed_eval_ref_log_probs = True1087 1088 return super().get_eval_dataloader(eval_dataset=eval_dataset)1089 1090 def compute_reference_log_probs(self, padded_batch: dict) -> dict:1091 """Computes log probabilities of the reference model for a single padded batch of a BCO specific dataset."""1092 with torch.no_grad():1093 if self.ref_model is None:1094 with self.null_ref_context():1095 if self.is_encoder_decoder:1096 completion_logits = self.model(1097 padded_batch["prompt_input_ids"],1098 attention_mask=padded_batch["prompt_attention_mask"],1099 decoder_input_ids=padded_batch.get("completion_decoder_input_ids"),1100 labels=padded_batch["completion_labels"],1101 ).logits1102 1103 else:1104 completion_logits = self.model(1105 padded_batch["completion_input_ids"],1106 attention_mask=padded_batch["completion_attention_mask"],1107 ).logits1108 1109 else:1110 if self.is_encoder_decoder:1111 completion_logits = self.ref_model(1112 padded_batch["prompt_input_ids"],1113 attention_mask=padded_batch["prompt_attention_mask"],1114 decoder_input_ids=padded_batch.get("completion_decoder_input_ids"),1115 labels=padded_batch["completion_labels"],1116 ).logits1117 1118 else:1119 completion_logits = self.ref_model(1120 padded_batch["completion_input_ids"], attention_mask=padded_batch["completion_attention_mask"]1121 ).logits1122 1123 completion_logps = self.get_batch_logps(1124 completion_logits,1125 padded_batch["completion_labels"],1126 average_log_prob=False,1127 is_encoder_decoder=self.is_encoder_decoder,1128 label_pad_token_id=self.label_pad_token_id,1129 )1130 1131 return completion_logps1132 1133 @staticmethod1134 def get_batch_logps(1135 logits: torch.FloatTensor,1136 labels: torch.LongTensor,1137 average_log_prob: bool = False,1138 label_pad_token_id: int = -100,1139 is_encoder_decoder: bool = False,1140 ) -> torch.FloatTensor:1141 """Compute the log probabilities of the given labels under the given logits.1142 1143 Args:1144 logits: Logits of the model (unnormalized). Shape: (batch_size, sequence_length, vocab_size)1145 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)1146 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.1147 1148 Returns:1149 A tensor of shape (batch_size,) containing the average/sum log probabilities of the given labels under the given logits.1150 """1151 if logits.shape[:-1] != labels.shape:1152 raise ValueError("Logits (batch and sequence length dim) and labels must have the same shape.")1153 1154 if not is_encoder_decoder:1155 labels = labels[:, 1:].clone()1156 logits = logits[:, :-1, :]1157 else:1158 # Fixes end-dec RuntimeError1159 labels = labels.clone()1160 1161 loss_mask = labels != label_pad_token_id1162 1163 # dummy token; we'll ignore the losses on these tokens later1164 labels[labels == label_pad_token_id] = 01165 1166 per_token_logps = selective_log_softmax(logits, labels)1167 1168 if average_log_prob:1169 return (per_token_logps * loss_mask).sum(-1) / loss_mask.sum(-1)1170 else:1171 return (per_token_logps * loss_mask).sum(-1)1172 1173 def forward(1174 self, model: nn.Module, batch: dict[str, Union[list, torch.LongTensor]]1175 ) -> tuple[torch.FloatTensor, torch.FloatTensor, torch.FloatTensor, torch.FloatTensor]:1176 model_kwargs = (1177 {1178 "labels": batch["completion_labels"],1179 "decoder_input_ids": batch.get("completion_decoder_input_ids"),1180 }1181 if self.is_encoder_decoder1182 else {}1183 )1184 if self.aux_loss_enabled:1185 model_kwargs["output_router_logits"] = True1186 1187 outputs = model(1188 batch["completion_input_ids"],1189 attention_mask=batch["completion_attention_mask"],1190 **model_kwargs,1191 )1192 completion_logits = outputs.logits1193 1194 completion_logps = self.get_batch_logps(1195 completion_logits,1196 batch["completion_labels"],1197 average_log_prob=False,1198 is_encoder_decoder=self.is_encoder_decoder,1199 label_pad_token_id=self.label_pad_token_id,1200 )