Team Ai
Apppublic

Zwounds/Boolean_Search_Query_Model

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
UnslothKTOTrainer.py1841 linesDownload Raw Back to unsloth_compiled_cache
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"],

Showing the first 1,200 of 1841 lines. Download the file for the rest.