Team Ai
Apppublic

Zwounds/Boolean_Search_Query_Model

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
UnslothDPOTrainer.py2090 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.dpo_trainer import (Any, AutoModelForCausalLM, BaseImageProcessor, Callable, DPOConfig, DPOTrainer, DataCollator, DataCollatorForPreference, DataLoader, Dataset, EvalLoopOutput, F, FDivergenceConstants, FDivergenceType, FeatureExtractionMixin, IterableDataset, Literal, MODEL_FOR_VISION_2_SEQ_MAPPING_NAMES, Optional, PartialState, PeftModel, PreTrainedModel, PreTrainedModelWrapper, PreTrainedTokenizerBase, ProcessorMixin, RunningMoments, SyncRefModelCallback, Trainer, TrainerCallback, Union, amp, cap_exp, contextmanager, create_reference_model, dataclass, deepcopy, defaultdict, deprecate_kwarg, disable_dropout_in_model, empty_cache, flush_left, generate_model_card, get_comet_experiment_url, inspect, is_comet_available, is_peft_available, is_torch_xpu_available, is_wandb_available, log_table_to_comet_experiment, maybe_apply_chat_template, maybe_extract_prompt, nn, nullcontext, os, pad, 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 UnslothDPOConfig(DPOConfig):44    """45    46    Configuration class for the [`DPOTrainer`].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        > Parameters that control the model and reference model54 55        model_init_kwargs (`dict[str, Any]` or `None`, *optional*, defaults to `None`):56            Keyword arguments for `AutoModelForCausalLM.from_pretrained`, used when the `model` argument of the57            [`DPOTrainer`] is provided as a string.58        ref_model_init_kwargs (`dict[str, Any]` or `None`, *optional*, defaults to `None`):59            Keyword arguments for `AutoModelForCausalLM.from_pretrained`, used when the `ref_model` argument of the60            [`DPOTrainer`] is provided as a string.61        model_adapter_name (`str` or `None`, *optional*, defaults to `None`):62            Name of the train target PEFT adapter, when using LoRA with multiple adapters.63        ref_adapter_name (`str` or `None`, *optional*, defaults to `None`):64            Name of the reference PEFT adapter, when using LoRA with multiple adapters.65        force_use_ref_model (`bool`, *optional*, defaults to `False`):66            If you provide a PEFT model as the active model and wish to use a different model for the `ref_model`, set67            this flag to `True`.68        disable_dropout (`bool`, *optional*, defaults to `True`):69            Whether to disable dropout in the model and reference model.70        use_logits_to_keep (`bool`, *optional*, defaults to `False`):71            If `True`, only a specified number of logits are computed in the forward pass. This can be useful for72            saving memory and speeding up training by not computing the logits for all tokens, especially in73            scenarios when working with very long prompts where labels are ignored (-100).74 75        > Parameters that control the data preprocessing76 77        dataset_num_proc (`int` or `None`, *optional*, defaults to `None`):78            Number of processes to use for processing the dataset.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        label_pad_token_id (`int`, *optional*, defaults to `-100`):82            Padding value to use for labels.83        max_prompt_length (`int` or `None`, *optional*, defaults to `512`):84            Maximum length of the prompt.85        max_completion_length (`int` or `None`, *optional*, defaults to `None`):86            Maximum length of the completion.87        max_length (`int` or `None`, *optional*, defaults to `1024`):88            Maximum length of the full sequence (prompt + completion).89        truncation_mode (`str`, *optional*, defaults to `"keep_end"`):90            Truncation mode to use when the sequence exceeds `max_length`. Possible values are `"keep_end"` and91            `"keep_start"`.92        padding_free (`bool`, *optional*, defaults to `False`):93            Whether forward passes are performed without padding by flattening all sequences in the batch94            into a single continuous sequence. This approach requires associating a `position_ids` vector to track95            positional information. Currently, this is only supported with the `flash_attention_2` mechanism, as it96            can handle the flattened batch structure.97        precompute_ref_log_probs (`bool`, *optional*, defaults to `False`):98            Whether to precompute the log probabilities from the reference model. Setting this to `True` allows99            training without needing the reference model during training, which can help reduce GPU memory usage. If100            set to `False` (default), the reference model will be used during training to compute log probabilities101            on-the-fly.102        precompute_ref_batch_size (`int` or `None`, *optional*, defaults to `None`):103            Batch size to use when precomputing reference model log probabilities. This can be set higher than the104            training batch size to speed up preprocessing. If `None`, defaults to `per_device_train_batch_size` for105            training and `per_device_eval_batch_size` for evaluation.106        tools (`Optional[list[Union[dict, Callable]]]`, *optional*, defaults to `None`):107            List of tools (callable functions) that will be accessible to the model.108            If the template does not support function calling, this argument will have no effect.109 110        > Parameters that control the training111 112        learning_rate (`float`, *optional*, defaults to `1e-6`):113            Initial learning rate for [`AdamW`] optimizer. The default value replaces that of114            [`~transformers.TrainingArguments`].115        loss_type (`str`, *optional*, defaults to `"sigmoid"`):116            Type of loss to use. Possible values are:117 118                - `"sigmoid"`: sigmoid loss from the original [DPO](https://huggingface.co/papers/2305.18290) paper.119                - `"hinge"`: hinge loss on the normalized likelihood from the [SLiC](https://huggingface.co/papers/2305.10425) paper.120                - `"ipo"`: IPO loss from the [IPO](https://huggingface.co/papers/2310.12036) paper.121                - `"exo_pair"`: pairwise EXO loss from the [EXO](https://huggingface.co/papers/2402.00856) paper.122                - `"nca_pair"`: pairwise NCA loss from the [NCA](https://huggingface.co/papers/2402.05369) paper.123                - `"robust"`: unbiased estimate of the DPO loss that is robust to preference noise from the [Robust DPO](https://huggingface.co/papers/2403.00409) paper.124                - `"bco_pair"`: pairwise BCO loss from the [BCO](https://huggingface.co/papers/2404.04656) paper.125                - `"sppo_hard"`: SPPO loss with hard label from the [SPPO](https://huggingface.co/papers/2405.00675) paper.126                - `"aot"`: AOT loss for paired datasets from the [AOT](https://huggingface.co/papers/2406.05882) paper.127                - `"aot_pair"`: AOT loss for unpaired datasets from the [AOT](https://huggingface.co/papers/2406.05882) paper.128                - `"discopop"`: DiscoPOP (a.k.a Log-Ratio Modulated Loss, LRML) loss from the [DiscoPOP](https://huggingface.co/papers/2406.08414) paper.129                - `"apo_zero"`: APO-zero loss from the [APO](https://huggingface.co/papers/2408.06266) paper.130                - `"apo_down"`: APO-down loss from the [APO](https://huggingface.co/papers/2408.06266) paper.131 132        beta (`float`, *optional*, defaults to `0.1`):133            Parameter controlling the deviation from the reference model. Higher β means less deviation from the134            reference model. For the IPO loss (`loss_type="ipo"`), β is the regularization parameter denoted by τ in135            the [paper](https://huggingface.co/papers/2310.12036).136        f_divergence_type (`str`, *optional*, defaults to `FDivergenceType.REVERSE_KL`):137            Type of f-divergence regularization function to compute divergence between policy and reference model.138        f_alpha_divergence_coef (`float`, *optional*, defaults to `1.0`):139            α coefficient in the α-divergence u^-α regularization function for DPO loss.140        reference_free (`bool`, *optional*, defaults to `False`):141            Whether to ignore the provided reference model and implicitly use a reference model that assigns equal142            probability to all responses.143        label_smoothing (`float`, *optional*, defaults to `0.0`):144            Robust DPO label smoothing parameter from the [cDPO](https://ericmitchell.ai/cdpo.pdf) report and145            [Robust DPO](https://huggingface.co/papers/2403.00409) paper that should be between `0.0` and `0.5`.146        use_weighting (`bool`, *optional*, defaults to `False`):147            Whether to weight the loss as done in the [WPO](https://huggingface.co/papers/2406.11827) paper.148        rpo_alpha (`float`, *optional*, defaults to `None`):149            α parameter from the [RPO](https://huggingface.co/papers/2404.19733) paper (v3), which controls the150            weighting of the NLL term in the loss. If `None`, no weighting is applied and the loss is the same as the151            DPO loss. The paper recommends `rpo_alpha=1.0`.152        discopop_tau (`float`, *optional*, defaults to `0.05`):153            τ/temperature parameter from the [DiscoPOP](https://huggingface.co/papers/2406.08414) paper, which controls154            the shape of log ratio modulated loss. The paper recommends the default value `discopop_tau=0.05`.155        sync_ref_model (`bool`, *optional*, defaults to `False`):156            Whether to synchronize the reference model with the active model every `ref_model_sync_steps` steps, using157            the `ref_model_mixup_alpha` parameter. This synchronization originites from the158            [TR-DPO](https://huggingface.co/papers/2404.09656) paper.159        ref_model_mixup_alpha (`float`, *optional*, defaults to `0.9`):160            α parameter from the [TR-DPO](https://huggingface.co/papers/2404.09656) paper, which controls the mix161            between the current policy and the previous reference policy during updates. The reference policy is162            updated according to the equation: `π_ref = α * π_θ + (1 - α) * π_ref_prev`. To use this parameter, you163            must set `sync_ref_model=True`.164        ref_model_sync_steps (`int`, *optional*, defaults to `64`):165            τ parameter from the [TR-DPO](https://huggingface.co/papers/2404.09656) paper, which determines how166            frequently the current policy is synchronized with the reference policy. To use this parameter, you must167            set `sync_ref_model=True`.168 169        > Parameters that control the logging170 171        generate_during_eval (`bool`, *optional*, defaults to `False`):172            Whether to generate and log completions from both the model and the reference model to W&B or Comet during173            evaluation.174    175    """176    vllm_sampling_params: Optional[Any] = field(177        default = None,178        metadata = {'help': 'vLLM SamplingParams'},179    )180    unsloth_num_chunks : Optional[int] = field(181        default = -1,182        metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},183    )184    def __init__(185        self,186        output_dir = None,187        overwrite_output_dir = None,188        do_train = False,189        do_eval = False,190        do_predict = False,191        eval_strategy = 'no',192        prediction_loss_only = False,193        per_device_train_batch_size = 4,194        per_device_eval_batch_size = 4,195        per_gpu_train_batch_size = None,196        per_gpu_eval_batch_size = None,197        gradient_accumulation_steps = 2,198        eval_accumulation_steps = 2,199        eval_delay = 0,200        torch_empty_cache_steps = 250,201        learning_rate = 5e-05,202        weight_decay = 0.01,203        adam_beta1 = 0.9,204        adam_beta2 = 0.999,205        adam_epsilon = 1e-08,206        max_grad_norm = 1.0,207        num_train_epochs = 3.0,208        max_steps = -1,209        lr_scheduler_type = 'linear',210        warmup_ratio = 0.1,211        warmup_steps = 0,212        log_level = 'passive',213        log_level_replica = 'warning',214        log_on_each_node = True,215        logging_dir = None,216        logging_strategy = 'steps',217        logging_first_step = False,218        logging_steps = 1,219        logging_nan_inf_filter = False,220        save_strategy = 'steps',221        save_steps = 500,222        save_total_limit = None,223        save_safetensors = True,224        save_on_each_node = False,225        save_only_model = False,226        restore_callback_states_from_checkpoint = False,227        no_cuda = False,228        use_cpu = False,229        use_mps_device = False,230        seed = 3407,231        data_seed = 3407,232        jit_mode_eval = False,233        use_ipex = False,234        bf16 = False,235        fp16 = False,236        fp16_opt_level = 'O1',237        half_precision_backend = 'auto',238        bf16_full_eval = False,239        fp16_full_eval = False,240        tf32 = None,241        local_rank = -1,242        ddp_backend = None,243        tpu_num_cores = None,244        tpu_metrics_debug = False,245        debug = '',246        dataloader_drop_last = False,247        eval_steps = None,248        dataloader_num_workers = 0,249        dataloader_prefetch_factor = None,250        past_index = -1,251        run_name = None,252        disable_tqdm = None,253        remove_unused_columns = True,254        label_names = None,255        load_best_model_at_end = False,256        metric_for_best_model = None,257        greater_is_better = None,258        ignore_data_skip = False,259        fsdp = '',260        fsdp_min_num_params = 0,261        fsdp_config = None,262        tp_size = 0,263        fsdp_transformer_layer_cls_to_wrap = None,264        accelerator_config = None,265        deepspeed = None,266        label_smoothing_factor = 0.0,267        optim = 'adamw_8bit',268        optim_args = None,269        adafactor = False,270        group_by_length = False,271        length_column_name = 'length',272        report_to = None,273        ddp_find_unused_parameters = None,274        ddp_bucket_cap_mb = None,275        ddp_broadcast_buffers = None,276        dataloader_pin_memory = True,277        dataloader_persistent_workers = False,278        skip_memory_metrics = True,279        use_legacy_prediction_loop = False,280        push_to_hub = False,281        resume_from_checkpoint = None,282        hub_model_id = None,283        hub_strategy = 'every_save',284        hub_token = None,285        hub_private_repo = None,286        hub_always_push = False,287        gradient_checkpointing = False,288        gradient_checkpointing_kwargs = None,289        include_inputs_for_metrics = False,290        eval_do_concat_batches = True,291        fp16_backend = 'auto',292        evaluation_strategy = None,293        push_to_hub_model_id = None,294        push_to_hub_organization = None,295        push_to_hub_token = None,296        mp_parameters = '',297        auto_find_batch_size = False,298        full_determinism = False,299        torchdynamo = None,300        ray_scope = 'last',301        ddp_timeout = 1800,302        torch_compile = False,303        torch_compile_backend = None,304        torch_compile_mode = None,305        dispatch_batches = None,306        split_batches = None,307        include_tokens_per_second = False,308        include_num_input_tokens_seen = False,309        neftune_noise_alpha = None,310        optim_target_modules = None,311        batch_eval_metrics = False,312        eval_on_start = False,313        use_liger_kernel = False,314        eval_use_gather_object = False,315        average_tokens_across_devices = False,316        model_init_kwargs = None,317        ref_model_init_kwargs = None,318        model_adapter_name = None,319        ref_adapter_name = None,320        force_use_ref_model = False,321        disable_dropout = True,322        use_logits_to_keep = False,323        dataset_num_proc = None,324        padding_value = None,325        label_pad_token_id = -100,326        max_prompt_length = 512,327        max_completion_length = None,328        max_length = 1024,329        truncation_mode = 'keep_end',330        padding_free = False,331        precompute_ref_log_probs = False,332        precompute_ref_batch_size = None,333        tools = None,334        loss_type = 'sigmoid',335        beta = 0.1,336        f_alpha_divergence_coef = 1.0,337        reference_free = False,338        label_smoothing = 0.0,339        use_weighting = False,340        rpo_alpha = None,341        discopop_tau = 0.05,342        sync_ref_model = False,343        ref_model_mixup_alpha = 0.9,344        ref_model_sync_steps = 64,345        generate_during_eval = False,346        use_num_logits_to_keep = False,347        vllm_sampling_params = None,348        unsloth_num_chunks = -1,349        **kwargs,350    ):351        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!')352        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!')353        if output_dir is None and save_strategy == 'steps' and save_steps == 500:354            output_dir = 'unsloth_training_checkpoints'355            save_strategy = 'no'356        if dataset_num_proc is None:357            from multiprocessing import cpu_count358            dataset_num_proc = cpu_count()359        360        super().__init__(361            output_dir = output_dir,362            overwrite_output_dir = overwrite_output_dir,363            do_train = do_train,364            do_eval = do_eval,365            do_predict = do_predict,366            eval_strategy = eval_strategy,367            prediction_loss_only = prediction_loss_only,368            per_device_train_batch_size = per_device_train_batch_size,369            per_device_eval_batch_size = per_device_eval_batch_size,370            per_gpu_train_batch_size = per_gpu_train_batch_size,371            per_gpu_eval_batch_size = per_gpu_eval_batch_size,372            gradient_accumulation_steps = gradient_accumulation_steps,373            eval_accumulation_steps = eval_accumulation_steps,374            eval_delay = eval_delay,375            torch_empty_cache_steps = torch_empty_cache_steps,376            learning_rate = learning_rate,377            weight_decay = weight_decay,378            adam_beta1 = adam_beta1,379            adam_beta2 = adam_beta2,380            adam_epsilon = adam_epsilon,381            max_grad_norm = max_grad_norm,382            num_train_epochs = num_train_epochs,383            max_steps = max_steps,384            lr_scheduler_type = lr_scheduler_type,385            warmup_ratio = warmup_ratio,386            warmup_steps = warmup_steps,387            log_level = log_level,388            log_level_replica = log_level_replica,389            log_on_each_node = log_on_each_node,390            logging_dir = logging_dir,391            logging_strategy = logging_strategy,392            logging_first_step = logging_first_step,393            logging_steps = logging_steps,394            logging_nan_inf_filter = logging_nan_inf_filter,395            save_strategy = save_strategy,396            save_steps = save_steps,397            save_total_limit = save_total_limit,398            save_safetensors = save_safetensors,399            save_on_each_node = save_on_each_node,400            save_only_model = save_only_model,401            restore_callback_states_from_checkpoint = restore_callback_states_from_checkpoint,402            no_cuda = no_cuda,403            use_cpu = use_cpu,404            use_mps_device = use_mps_device,405            seed = seed,406            data_seed = data_seed,407            jit_mode_eval = jit_mode_eval,408            use_ipex = use_ipex,409            bf16 = bf16,410            fp16 = fp16,411            fp16_opt_level = fp16_opt_level,412            half_precision_backend = half_precision_backend,413            bf16_full_eval = bf16_full_eval,414            fp16_full_eval = fp16_full_eval,415            tf32 = tf32,416            local_rank = local_rank,417            ddp_backend = ddp_backend,418            tpu_num_cores = tpu_num_cores,419            tpu_metrics_debug = tpu_metrics_debug,420            debug = debug,421            dataloader_drop_last = dataloader_drop_last,422            eval_steps = eval_steps,423            dataloader_num_workers = dataloader_num_workers,424            dataloader_prefetch_factor = dataloader_prefetch_factor,425            past_index = past_index,426            run_name = run_name,427            disable_tqdm = disable_tqdm,428            remove_unused_columns = remove_unused_columns,429            label_names = label_names,430            load_best_model_at_end = load_best_model_at_end,431            metric_for_best_model = metric_for_best_model,432            greater_is_better = greater_is_better,433            ignore_data_skip = ignore_data_skip,434            fsdp = fsdp,435            fsdp_min_num_params = fsdp_min_num_params,436            fsdp_config = fsdp_config,437            tp_size = tp_size,438            fsdp_transformer_layer_cls_to_wrap = fsdp_transformer_layer_cls_to_wrap,439            accelerator_config = accelerator_config,440            deepspeed = deepspeed,441            label_smoothing_factor = label_smoothing_factor,442            optim = optim,443            optim_args = optim_args,444            adafactor = adafactor,445            group_by_length = group_by_length,446            length_column_name = length_column_name,447            report_to = report_to,448            ddp_find_unused_parameters = ddp_find_unused_parameters,449            ddp_bucket_cap_mb = ddp_bucket_cap_mb,450            ddp_broadcast_buffers = ddp_broadcast_buffers,451            dataloader_pin_memory = dataloader_pin_memory,452            dataloader_persistent_workers = dataloader_persistent_workers,453            skip_memory_metrics = skip_memory_metrics,454            use_legacy_prediction_loop = use_legacy_prediction_loop,455            push_to_hub = push_to_hub,456            resume_from_checkpoint = resume_from_checkpoint,457            hub_model_id = hub_model_id,458            hub_strategy = hub_strategy,459            hub_token = hub_token,460            hub_private_repo = hub_private_repo,461            hub_always_push = hub_always_push,462            gradient_checkpointing = gradient_checkpointing,463            gradient_checkpointing_kwargs = gradient_checkpointing_kwargs,464            include_inputs_for_metrics = include_inputs_for_metrics,465            eval_do_concat_batches = eval_do_concat_batches,466            fp16_backend = fp16_backend,467            evaluation_strategy = evaluation_strategy,468            push_to_hub_model_id = push_to_hub_model_id,469            push_to_hub_organization = push_to_hub_organization,470            push_to_hub_token = push_to_hub_token,471            mp_parameters = mp_parameters,472            auto_find_batch_size = auto_find_batch_size,473            full_determinism = full_determinism,474            torchdynamo = torchdynamo,475            ray_scope = ray_scope,476            ddp_timeout = ddp_timeout,477            torch_compile = torch_compile,478            torch_compile_backend = torch_compile_backend,479            torch_compile_mode = torch_compile_mode,480            dispatch_batches = dispatch_batches,481            split_batches = split_batches,482            include_tokens_per_second = include_tokens_per_second,483            include_num_input_tokens_seen = include_num_input_tokens_seen,484            neftune_noise_alpha = neftune_noise_alpha,485            optim_target_modules = optim_target_modules,486            batch_eval_metrics = batch_eval_metrics,487            eval_on_start = eval_on_start,488            use_liger_kernel = use_liger_kernel,489            eval_use_gather_object = eval_use_gather_object,490            average_tokens_across_devices = average_tokens_across_devices,491            model_init_kwargs = model_init_kwargs,492            ref_model_init_kwargs = ref_model_init_kwargs,493            model_adapter_name = model_adapter_name,494            ref_adapter_name = ref_adapter_name,495            force_use_ref_model = force_use_ref_model,496            disable_dropout = disable_dropout,497            use_logits_to_keep = use_logits_to_keep,498            dataset_num_proc = dataset_num_proc,499            padding_value = padding_value,500            label_pad_token_id = label_pad_token_id,501            max_prompt_length = max_prompt_length,502            max_completion_length = max_completion_length,503            max_length = max_length,504            truncation_mode = truncation_mode,505            padding_free = padding_free,506            precompute_ref_log_probs = precompute_ref_log_probs,507            precompute_ref_batch_size = precompute_ref_batch_size,508            tools = tools,509            loss_type = loss_type,510            beta = beta,511            f_alpha_divergence_coef = f_alpha_divergence_coef,512            reference_free = reference_free,513            label_smoothing = label_smoothing,514            use_weighting = use_weighting,515            rpo_alpha = rpo_alpha,516            discopop_tau = discopop_tau,517            sync_ref_model = sync_ref_model,518            ref_model_mixup_alpha = ref_model_mixup_alpha,519            ref_model_sync_steps = ref_model_sync_steps,520            generate_during_eval = generate_during_eval,521            use_num_logits_to_keep = use_num_logits_to_keep,**kwargs)522        self.vllm_sampling_params = vllm_sampling_params523        self.unsloth_num_chunks = unsloth_num_chunks524pass525 526class _UnslothDPOTrainer(Trainer):527    r""""""528 529    _tag_names = ["trl", "dpo"]530 531    @deprecate_kwarg(532        "tokenizer", "0.16.0", "processing_class", warn_if_greater_or_equal_version=True, raise_if_both_names=True533    )534    def __init__(535        self,536        model: Optional[Union[PreTrainedModel, nn.Module, str]] = None,537        ref_model: Optional[Union[PreTrainedModel, nn.Module, str]] = None,538        args: Optional[DPOConfig] = None,539        data_collator: Optional[DataCollator] = None,540        train_dataset: Optional[Dataset] = None,541        eval_dataset: Optional[Union[Dataset, dict[str, Dataset]]] = None,542        processing_class: Optional[543            Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin]544        ] = None,545        model_init: Optional[Callable[[], PreTrainedModel]] = None,546        compute_metrics: Optional[Callable[[EvalLoopOutput], dict]] = None,547        callbacks: Optional[list[TrainerCallback]] = None,548        optimizers: tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),549        preprocess_logits_for_metrics: Optional[Callable[[torch.Tensor, torch.Tensor], torch.Tensor]] = None,550        peft_config: Optional[dict] = None,551    ):552        if model is None:553            raise ValueError("No model provided. Please provide a model to train.")554 555        if not isinstance(model, str) and ref_model is model:556            raise ValueError(557                "`model` and `ref_model` cannot be the same object. If you want `ref_model` to be the "558                "same as `model`, you must mass a copy of it, or `None` if you use peft."559            )560 561        if args.model_init_kwargs is None:562            model_init_kwargs = {}563        elif not isinstance(model, str):564            raise ValueError(565                "You passed model_init_kwargs to the DPOTrainer/DPOConfig, but your model is already instantiated."566            )567        else:568            model_init_kwargs = args.model_init_kwargs569            torch_dtype = model_init_kwargs.get("torch_dtype")570            if torch_dtype is not None:571                # Convert to `torch.dtype` if an str is passed572                if isinstance(torch_dtype, str) and torch_dtype != "auto":573                    torch_dtype = getattr(torch, torch_dtype)574                if torch_dtype != "auto" and not isinstance(torch_dtype, torch.dtype):575                    raise ValueError(576                        f"Invalid `torch_dtype` passed to the DPOConfig. Expected a string with either `torch.dtype` or 'auto', but got {torch_dtype}."577                    )578                model_init_kwargs["torch_dtype"] = torch_dtype579 580        if args.ref_model_init_kwargs is None:581            ref_model_init_kwargs = {}582        elif not isinstance(ref_model, str):583            raise ValueError(584                "You passed ref_model_init_kwargs to the DPOTrainer/DPOConfig, but your ref_model is already instantiated."585            )586        else:587            ref_model_init_kwargs = args.ref_model_init_kwargs588            torch_dtype = ref_model_init_kwargs.get("torch_dtype")589            if torch_dtype is not None:590                # Convert to `torch.dtype` if an str is passed591                if isinstance(torch_dtype, str) and torch_dtype != "auto":592                    torch_dtype = getattr(torch, torch_dtype)593                if torch_dtype != "auto" and not isinstance(torch_dtype, torch.dtype):594                    raise ValueError(595                        f"Invalid `torch_dtype` passed to the DPOConfig. Expected a string with either `torch.dtype` or 'auto', but got {torch_dtype}."596                    )597                ref_model_init_kwargs["torch_dtype"] = torch_dtype598 599        if isinstance(model, str):600            model = AutoModelForCausalLM.from_pretrained(model, **model_init_kwargs)601 602        if isinstance(ref_model, str):603            ref_model = AutoModelForCausalLM.from_pretrained(ref_model, **ref_model_init_kwargs)604 605        # Initialize this variable to False. This helps tracking the case when `peft_module_casting_to_bf16`606        # has been called in order to properly call autocast if needed.607        self._peft_has_been_casted_to_bf16 = False608 609        if not is_peft_available() and peft_config is not None:610            raise ValueError(611                "PEFT is not installed and you passed a `peft_config` in the trainer's kwargs, please install it to use the PEFT models"612            )613        elif is_peft_available() and peft_config is not None:614            # if model is a peft model and we have a peft_config, we merge and unload it first615            if isinstance(model, PeftModel):616                model = model.merge_and_unload()617 618            if ref_model is not None and not args.force_use_ref_model:619                raise ValueError(620                    "You passed both a ref_model and a peft_config. For training PEFT adapters with DPO there is no need to pass a reference"621                    " model. Please pass `ref_model=None` in case you want to train PEFT adapters, or pass a ref_model with `force_use_ref_model=True` in DPOTrainer's init."622                    " if you want to use a different ref_model."623                )624 625            if getattr(model, "is_loaded_in_8bit", False) or getattr(model, "is_loaded_in_4bit", False):626                _support_gc_kwargs = hasattr(627                    args, "gradient_checkpointing_kwargs"628                ) and "gradient_checkpointing_kwargs" in list(629                    inspect.signature(prepare_model_for_kbit_training).parameters630                )631 632                prepare_model_kwargs = {"use_gradient_checkpointing": args.gradient_checkpointing}633 634                if _support_gc_kwargs:635                    prepare_model_kwargs["gradient_checkpointing_kwargs"] = args.gradient_checkpointing_kwargs636 637                model = prepare_model_for_kbit_training(model, **prepare_model_kwargs)638            elif getattr(args, "gradient_checkpointing", False):639                # For backward compatibility with older versions of transformers640                if hasattr(model, "enable_input_require_grads"):641                    model.enable_input_require_grads()642                else:643 644                    def make_inputs_require_grad(module, input, output):645                        output.requires_grad_(True)646 647                    model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)648 649            # get peft model with the given config650            model = model651            if args.bf16 and getattr(model, "is_loaded_in_4bit", False):652                peft_module_casting_to_bf16(model)653                # If args.bf16 we need to explicitly call `generate` with torch amp autocast context manager654                self._peft_has_been_casted_to_bf16 = True655 656        # For models that use gradient_checkpointing, we need to attach a hook that enables input657        # to explicitly have `requires_grad=True`, otherwise training will either silently658        # fail or completely fail.659        elif getattr(args, "gradient_checkpointing", False):660            # For backward compatibility with older versions of transformers661            if hasattr(model, "enable_input_require_grads"):662                model.enable_input_require_grads()663            else:664 665                def make_inputs_require_grad(module, input, output):666                    output.requires_grad_(True)667 668                model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)669 670        if args.generate_during_eval and not (is_wandb_available() or is_comet_available()):671            raise ValueError(672                "`generate_during_eval=True` requires Weights and Biases or Comet to be installed."673                " Please install `wandb` or `comet-ml` to resolve."674            )675 676        self.is_encoder_decoder = model.config.is_encoder_decoder677        self.is_vision_model = model.config.model_type in MODEL_FOR_VISION_2_SEQ_MAPPING_NAMES.keys()678        self.is_peft_model = is_peft_available() and isinstance(model, PeftModel)679        self.model_adapter_name = args.model_adapter_name680        self.ref_adapter_name = args.ref_adapter_name681        self.reference_free = args.reference_free682 683        if ref_model:684            self.ref_model = ref_model685        elif self.is_peft_model or args.precompute_ref_log_probs:686            # The `model` with adapters turned off will be used as the reference model687            self.ref_model = None688        else:689            self.ref_model = create_reference_model(model)690 691        if processing_class is None:692            raise ValueError("processing_class must be specified to tokenize a DPO dataset.")693 694        if args.padding_value is not None:695            self.padding_value = args.padding_value696        else:697            if hasattr(processing_class, "pad_token_id") and processing_class.pad_token_id is not None:698                self.padding_value = processing_class.pad_token_id699            elif hasattr(processing_class, "tokenizer") and processing_class.tokenizer.pad_token_id is not None:700                self.padding_value = processing_class.tokenizer.pad_token_id701            else:702                raise ValueError(703                    "`padding_value` is not specified in `DPOConfig`, and `pad_token_id` is missing in the "704                    "`processing_class`. Please either set the `padding_value` argument in `DPOConfig`, or set "705                    "`tokenizer.pad_token` (e.g., `tokenizer.pad_token = tokenizer.eos_token`) before instantiating "706                    "the trainer."707                )708 709        if data_collator is None:710            data_collator = DataCollatorForPreference(pad_token_id=self.padding_value)711 712        # Disable dropout in the model and reference model713        if args.disable_dropout:714            disable_dropout_in_model(model)715            if self.ref_model is not None:716                disable_dropout_in_model(self.ref_model)717 718        self.generate_during_eval = args.generate_during_eval719        self.label_pad_token_id = args.label_pad_token_id720        self.max_prompt_length = args.max_prompt_length721        self.max_completion_length = args.max_completion_length722        self.max_length = args.max_length723        self.truncation_mode = args.truncation_mode724        self.precompute_ref_log_probs = args.precompute_ref_log_probs725        self.use_logits_to_keep = args.use_logits_to_keep726 727        if args.padding_free:728            if model.config._attn_implementation != "flash_attention_2":729                warnings.warn(730                    "Padding-free training is enabled, but the attention implementation is not set to "731                    "'flash_attention_2'. Padding-free training flattens batches into a single sequence, and "732                    "'flash_attention_2' is the only known attention mechanism that reliably supports this. Using "733                    "other implementations may lead to unexpected behavior. To ensure compatibility, set "734                    "`attn_implementation='flash_attention_2'` in the model configuration, or verify that your "735                    "attention mechanism can handle flattened sequences."736                )737        self.padding_free = args.padding_free738 739        # Since ref_logs are precomputed on the first call to get_train/eval_dataloader740        # keep track of first called to avoid computation of future calls741        self._precomputed_train_ref_log_probs = False742        self._precomputed_eval_ref_log_probs = False743 744        if (745            args.loss_type in ["hinge", "ipo", "bco_pair", "sppo_hard", "nca_pair", "apo_zero", "apo_down"]746            and args.label_smoothing > 0747        ):748            warnings.warn(749                f"You are using the {args.loss_type} loss type that does not support label smoothing. The "750                "`label_smoothing` parameter will be ignored. Set `label_smoothing` to `0.0` to remove this warning.",751                UserWarning,752            )753        if args.loss_type == "kto_pair":754            raise ValueError("Support for kto_pair has been removed in DPOTrainer. Please use KTOTrainer.")755 756        self.beta = args.beta757        self.label_smoothing = args.label_smoothing758        self.loss_type = args.loss_type759        self.aux_loss_enabled = getattr(model.config, "output_router_logits", False)760        self.use_weighting = args.use_weighting761        self.aux_loss_coef = getattr(model.config, "router_aux_loss_coef", 0.0)762        if self.aux_loss_enabled and self.aux_loss_coef == 0.0:763            warnings.warn(764                "You set `output_router_logits` to `True` in the model config, but `router_aux_loss_coef` is set to "765                "`0.0`, meaning the auxiliary loss will not be used. Either set `router_aux_loss_coef` to a value "766                "greater than `0.0`, or set `output_router_logits` to `False` if you don't want to use the auxiliary "767                "loss.",768                UserWarning,769            )770 771        self._stored_metrics = defaultdict(lambda: defaultdict(list))772        self.f_divergence_type = args.f_divergence_type773        self.f_divergence_params = {FDivergenceConstants.ALPHA_DIVERGENCE_COEF_KEY: args.f_alpha_divergence_coef}774        self.dataset_num_proc = args.dataset_num_proc775 776        # The trainer estimates the number of FLOPs (floating-point operations) using the number of elements in the777        # input tensor associated with the key "input_ids". However, in DPO, the sampled data does not include the778        # "input_ids" key. Instead, the available keys are "prompt_input_ids", "chosen_input_ids", and779        # "rejected_input_ids". As a result, the trainer issues the warning: "Could not estimate the number of tokens780        # of the input, floating-point operations will not be computed." To suppress this warning, we set the781        # "estimate_tokens" key in the model's "warnings_issued" dictionary to True. This acts as a flag to indicate782        # that the warning has already been issued.783        model.warnings_issued["estimate_tokens"] = True784 785        # Dataset preparation786        train_dataset = self._prepare_dataset(train_dataset, processing_class, args, "train")787        if eval_dataset is not None:788            if isinstance(eval_dataset, dict):789                eval_dataset = {790                    key: self._prepare_dataset(dataset, processing_class, args, key)791                    for key, dataset in eval_dataset.items()792                }793            else:794                eval_dataset = self._prepare_dataset(eval_dataset, processing_class, args, "eval")795 796        super().__init__(797            model=model,798            args=args,799            data_collator=data_collator,800            train_dataset=train_dataset,801            eval_dataset=eval_dataset,802            processing_class=processing_class,803            model_init=model_init,804            compute_metrics=compute_metrics,805            callbacks=callbacks,806            optimizers=optimizers,807            preprocess_logits_for_metrics=preprocess_logits_for_metrics,808        )809 810        # Gradient accumulation requires scaled loss. Normally, loss scaling in the parent class depends on whether the811        # model accepts loss-related kwargs. Since we compute our own loss, this check is irrelevant. We set812        # self.model_accepts_loss_kwargs to False to enable scaling.813        self.model_accepts_loss_kwargs = False814 815        # Add tags for models that have been loaded with the correct transformers version816        if hasattr(self.model, "add_model_tags"):817            self.model.add_model_tags(self._tag_names)818 819        if not hasattr(self, "accelerator"):820            raise AttributeError(821                "Your `Trainer` does not have an `accelerator` object. Consider upgrading `transformers`."822            )823 824        # Deepspeed Zero-3 does not support precompute_ref_log_probs825        if self.is_deepspeed_enabled:826            if self.accelerator.state.deepspeed_plugin.zero_stage == 3 and self.precompute_ref_log_probs:827                raise ValueError(828                    "You cannot use `precompute_ref_log_probs=True` with Deepspeed ZeRO-3. Please set `precompute_ref_log_probs=False`."829                )830 831        if self.ref_model is None:832            if not (self.is_peft_model or self.precompute_ref_log_probs):833                raise ValueError(834                    "No reference model and model is not a Peft model. Try setting `precompute_ref_log_probs=True`"835                )836            if args.sync_ref_model:837                raise ValueError(838                    "You currently cannot use `ref_model=None` with TR-DPO method. Please provide `ref_model`."839                )840        else:841            if self.is_deepspeed_enabled:842                self.ref_model = self._prepare_deepspeed(self.ref_model)843            else:844                self.ref_model = self.accelerator.prepare_model(self.ref_model, evaluation_mode=True)845 846        if args.sync_ref_model:847            if self.precompute_ref_log_probs:848                raise ValueError(849                    "You cannot use `precompute_ref_log_probs=True` with TR-DPO method. Please set `precompute_ref_log_probs=False`."850                )851 852            self.add_callback(SyncRefModelCallback(ref_model=self.ref_model, accelerator=self.accelerator))853 854        if self.loss_type == "bco_pair":855            self.running = RunningMoments(self.accelerator)856 857    def _prepare_dataset(858        self,859        dataset: Union[Dataset, IterableDataset],860        processing_class: Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin],861        args: DPOConfig,862        dataset_name: str,863    ) -> Union[Dataset, IterableDataset]:864        # Build the kwargs for the `map` function865        map_kwargs = {"writer_batch_size": 10}866        if isinstance(dataset, Dataset):  # IterableDataset does not support num_proc867            map_kwargs["num_proc"] = args.dataset_num_proc868 869        with PartialState().local_main_process_first():870            # Extract prompt if needed871            if isinstance(dataset, Dataset):  # `IterableDataset.map` does not support `desc`872                map_kwargs["desc"] = f"Extracting prompt in {dataset_name} dataset"873            dataset = dataset.map(maybe_extract_prompt, **map_kwargs)874 875            # Apply the chat template if needed876            if isinstance(dataset, Dataset):  # `IterableDataset.map` does not support `desc`877                map_kwargs["desc"] = f"Applying chat template to {dataset_name} dataset"878            dataset = dataset.map(879                maybe_apply_chat_template, fn_kwargs={"tokenizer": processing_class, "tools": args.tools}, **map_kwargs880            )881 882            # Tokenize the dataset883            if isinstance(dataset, Dataset):  # `IterableDataset.map` does not support `desc`884                map_kwargs["desc"] = f"Tokenizing {dataset_name} dataset"885 886            dataset = dataset.map(887                self.tokenize_row if not self.is_vision_model else self.process_row,888                remove_columns=["prompt", "chosen", "rejected"],889                fn_kwargs={890                    "processing_class": processing_class,891                    "max_prompt_length": args.max_prompt_length,892                    "max_completion_length": args.max_completion_length,893                    # for enc-dec, we add the special tokens ([bos_token] + prompt + [eos_token]; completion + [eos_token])894                    "add_special_tokens": False,895                },896                **map_kwargs,897            )898 899        return dataset900 901    @staticmethod902    def tokenize_row(features, processing_class, max_prompt_length, max_completion_length, add_special_tokens):903        """904        Tokenize a row of the dataset.905 906        Args:907            features (`dict[str, str]`):908                Row of the dataset, should contain the keys `"prompt"`, `"chosen"`, and `"rejected"`.909            processing_class (`PreTrainedTokenizerBase`):910                Processing class used to process the data.911            max_prompt_length (`int` or `None`):912                Maximum length of the prompt sequence. If `None`, the prompt sequence is not truncated.913            max_completion_length (`int` or `None`):914                Maximum length of the completion sequences. If `None`, the completion sequences are not truncated.915            add_special_tokens (`bool`):916                Whether to add special tokens to the sequences. Typically used for encoder-decoder models. If `True`,917                the prompt sequence will have a bos token prepended and an eos token appended. In any case, the918                completion sequences will have an eos token appended.919 920        Returns:921            `dict[str, list[int]]`:922                Tokenized sequences with the keys `"prompt_input_ids"`, `"chosen_input_ids"`, and923                `"rejected_input_ids".924 925        Example:926        ```python927        >>> from transformers import GPT2Tokenizer928        >>> tokenizer = GPT2Tokenizer.from_pretrained("gpt2")929        >>> features = {"prompt": "The sky is", "chosen": " blue", "rejected": " green"}930        >>> DPOTrainer.tokenize_row(931        ...     features, tokenizer, max_prompt_length=3, max_completion_length=3, add_special_tokens=False932        ... )933        {'prompt_input_ids': [464, 6766, 318], 'chosen_input_ids': [4171, 50256], 'rejected_input_ids': [4077, 50256]}934        ```935        """936        tokenizer = processing_class  # the processing class is a tokenizer937        prompt_input_ids = tokenizer(features["prompt"], add_special_tokens=False)["input_ids"]938        chosen_input_ids = tokenizer(features["chosen"], add_special_tokens=False)["input_ids"]939        rejected_input_ids = tokenizer(features["rejected"], add_special_tokens=False)["input_ids"]940 941        # Add special tokens (typically for encoder-decoder models)942        if add_special_tokens:943            if tokenizer.bos_token_id is not None:944                prompt_input_ids = [tokenizer.bos_token_id] + prompt_input_ids945            if tokenizer.eos_token_id is not None:946                prompt_input_ids = prompt_input_ids + [tokenizer.eos_token_id]947        chosen_input_ids = chosen_input_ids + [tokenizer.eos_token_id]948        rejected_input_ids = rejected_input_ids + [tokenizer.eos_token_id]949 950        # Truncate prompt and completion sequences951        if max_prompt_length is not None:952            prompt_input_ids = prompt_input_ids[-max_prompt_length:]953        if max_completion_length is not None:954            chosen_input_ids = chosen_input_ids[:max_completion_length]955            rejected_input_ids = rejected_input_ids[:max_completion_length]956 957        return {958            "prompt_input_ids": prompt_input_ids,959            "chosen_input_ids": chosen_input_ids,960            "rejected_input_ids": rejected_input_ids,961        }962 963    @staticmethod964    def process_row(features, processing_class, max_prompt_length, max_completion_length, add_special_tokens):965        """966        Same as `tokenize_row` but for vision models. Please refer to `tokenize_row` for more information.967        """968        processor, tokenizer = processing_class, processing_class.tokenizer  # the processing class is a processor969        processed_features = processor(images=features["images"], text=features["prompt"], add_special_tokens=False)970 971        prompt_input_ids = processed_features["input_ids"][0]972        pixel_values = processed_features["pixel_values"][0]973        chosen_input_ids = tokenizer(features["chosen"], add_special_tokens=False)["input_ids"]974        rejected_input_ids = tokenizer(features["rejected"], add_special_tokens=False)["input_ids"]975 976        # Add special tokens (typically for encoder-decoder models)977        if add_special_tokens:978            if tokenizer.bos_token_id is not None:979                prompt_input_ids = [tokenizer.bos_token_id] + prompt_input_ids980            if tokenizer.eos_token_id is not None:981                prompt_input_ids = prompt_input_ids + [tokenizer.eos_token_id]982        chosen_input_ids = chosen_input_ids + [tokenizer.eos_token_id]983        rejected_input_ids = rejected_input_ids + [tokenizer.eos_token_id]984 985        # Truncate prompt and completion sequences986        if max_prompt_length is not None:987            prompt_input_ids = prompt_input_ids[-max_prompt_length:]988        if max_completion_length is not None:989            chosen_input_ids = chosen_input_ids[:max_completion_length]990            rejected_input_ids = rejected_input_ids[:max_completion_length]991 992        output = {993            "prompt_input_ids": prompt_input_ids,994            "pixel_values": pixel_values,995            "chosen_input_ids": chosen_input_ids,996            "rejected_input_ids": rejected_input_ids,997        }998 999        if "pixel_attention_mask" in processed_features:1000            output["pixel_attention_mask"] = processed_features["pixel_attention_mask"][0]1001        if "image_sizes" in processed_features:1002            output["image_sizes"] = processed_features["image_sizes"][0]1003 1004        return output1005 1006    def _prepare_deepspeed(self, model: PreTrainedModelWrapper):1007        # Adapted from accelerate: https://github.com/huggingface/accelerate/blob/739b135f8367becb67ffaada12fe76e3aa60fefd/src/accelerate/accelerator.py#L14731008        deepspeed_plugin = self.accelerator.state.deepspeed_plugin1009        config_kwargs = deepcopy(deepspeed_plugin.deepspeed_config)1010 1011        if model is not None:1012            if hasattr(model, "config"):1013                hidden_size = (1014                    max(model.config.hidden_sizes)1015                    if getattr(model.config, "hidden_sizes", None)1016                    else getattr(model.config, "hidden_size", None)1017                )1018                if hidden_size is not None and config_kwargs["zero_optimization"]["stage"] == 3:1019                    # Note that `stage3_prefetch_bucket_size` can produce DeepSpeed messages like: `Invalidate trace cache @ step 0: expected module 1, but got module 0`1020                    # This is expected and is not an error, see: https://github.com/microsoft/DeepSpeed/discussions/40811021                    config_kwargs.update(1022                        {1023                            "zero_optimization.reduce_bucket_size": hidden_size * hidden_size,1024                            "zero_optimization.stage3_param_persistence_threshold": 10 * hidden_size,1025                            "zero_optimization.stage3_prefetch_bucket_size": 0.9 * hidden_size * hidden_size,1026                        }1027                    )1028 1029        # If ZeRO-3 is used, we shard both the active and reference model.1030        # Otherwise, we assume the reference model fits in memory and is initialized on each device with ZeRO disabled (stage 0)1031        if config_kwargs["zero_optimization"]["stage"] != 3:1032            config_kwargs["zero_optimization"]["stage"] = 01033        model, *_ = deepspeed.initialize(model=model, config=config_kwargs)1034        model.eval()1035        return model1036 1037    def _set_signature_columns_if_needed(self):1038        # If `self.args.remove_unused_columns` is True, non-signature columns are removed.1039        # By default, this method sets `self._signature_columns` to the model's expected inputs.1040        # In DPOTrainer, we preprocess data, so using the model's signature columns doesn't work.1041        # Instead, we set them to the columns expected by `DataCollatorForPreference`, hence the override.1042        if self._signature_columns is None:1043            self._signature_columns = [1044                "prompt_input_ids",1045                "chosen_input_ids",1046                "rejected_input_ids",1047                "image_sizes",1048                "ref_chosen_logps",1049                "ref_rejected_logps",1050            ]1051 1052    def get_train_dataloader(self) -> DataLoader:1053        """1054        Returns the training [`~torch.utils.data.DataLoader`].1055 1056        Subclass of transformers.src.transformers.trainer.get_train_dataloader to precompute `ref_log_probs`.1057        """1058 1059        if self.precompute_ref_log_probs and not self._precomputed_train_ref_log_probs:1060            batch_size = self.args.precompute_ref_batch_size or self.args.per_device_train_batch_size1061            dataloader_params = {1062                "batch_size": batch_size,1063                "collate_fn": self.data_collator,1064                "num_workers": self.args.dataloader_num_workers,1065                "pin_memory": self.args.dataloader_pin_memory,1066                "shuffle": False,1067            }1068 1069            # prepare dataloader1070            data_loader = self.accelerator.prepare(DataLoader(self.train_dataset, **dataloader_params))1071 1072            ref_chosen_logps = []1073            ref_rejected_logps = []1074            for padded_batch in tqdm(iterable=data_loader, desc="Train dataset reference log probs"):1075                ref_chosen_logp, ref_rejected_logp = self.compute_ref_log_probs(padded_batch)1076                ref_chosen_logp, ref_rejected_logp = self.accelerator.gather_for_metrics(1077                    (ref_chosen_logp, ref_rejected_logp)1078                )1079                ref_chosen_logps.append(ref_chosen_logp.cpu())1080                ref_rejected_logps.append(ref_rejected_logp.cpu())1081 1082                # Unnecessary cache clearing to avoid OOM1083                empty_cache()1084                self.accelerator.free_memory()1085 1086            all_ref_chosen_logps = torch.cat(ref_chosen_logps).float().numpy()1087            all_ref_rejected_logps = torch.cat(ref_rejected_logps).float().numpy()1088 1089            self.train_dataset = self.train_dataset.add_column(name="ref_chosen_logps", column=all_ref_chosen_logps)1090            self.train_dataset = self.train_dataset.add_column(1091                name="ref_rejected_logps", column=all_ref_rejected_logps1092            )1093 1094            self._precomputed_train_ref_log_probs = True1095 1096        return super().get_train_dataloader()1097 1098    def get_eval_dataloader(self, eval_dataset: Optional[Dataset] = None) -> DataLoader:1099        """1100        Returns the evaluation [`~torch.utils.data.DataLoader`].1101 1102        Subclass of transformers.src.transformers.trainer.get_eval_dataloader to precompute `ref_log_probs`.1103 1104        Args:1105            eval_dataset (`torch.utils.data.Dataset`, *optional*):1106                If provided, will override `self.eval_dataset`. If it is a [`~datasets.Dataset`], columns not accepted1107                by the `model.forward()` method are automatically removed. It must implement `__len__`.1108        """1109        if eval_dataset is None and self.eval_dataset is None:1110            raise ValueError("Trainer: evaluation requires an eval_dataset.")1111        eval_dataset = eval_dataset if eval_dataset is not None else self.eval_dataset1112 1113        if self.precompute_ref_log_probs and not self._precomputed_eval_ref_log_probs:1114            batch_size = self.args.precompute_ref_batch_size or self.args.per_device_eval_batch_size1115            dataloader_params = {1116                "batch_size": batch_size,1117                "collate_fn": self.data_collator,1118                "num_workers": self.args.dataloader_num_workers,1119                "pin_memory": self.args.dataloader_pin_memory,1120                "shuffle": False,1121            }1122 1123            # prepare dataloader1124            data_loader = self.accelerator.prepare(DataLoader(eval_dataset, **dataloader_params))1125 1126            ref_chosen_logps = []1127            ref_rejected_logps = []1128            for padded_batch in tqdm(iterable=data_loader, desc="Eval dataset reference log probs"):1129                ref_chosen_logp, ref_rejected_logp = self.compute_ref_log_probs(padded_batch)1130                ref_chosen_logp, ref_rejected_logp = self.accelerator.gather_for_metrics(1131                    (ref_chosen_logp, ref_rejected_logp)1132                )1133                ref_chosen_logps.append(ref_chosen_logp.cpu())1134                ref_rejected_logps.append(ref_rejected_logp.cpu())1135 1136            all_ref_chosen_logps = torch.cat(ref_chosen_logps).float().numpy()1137            all_ref_rejected_logps = torch.cat(ref_rejected_logps).float().numpy()1138 1139            eval_dataset = eval_dataset.add_column(name="ref_chosen_logps", column=all_ref_chosen_logps)1140            eval_dataset = eval_dataset.add_column(name="ref_rejected_logps", column=all_ref_rejected_logps)1141 1142            # Save calculated ref_chosen_logps and ref_rejected_logps to the eval_dataset for subsequent runs1143            if self.eval_dataset is not None:1144                self.eval_dataset = eval_dataset1145            self._precomputed_eval_ref_log_probs = True1146 1147        return super().get_eval_dataloader(eval_dataset=eval_dataset)1148 1149    @contextmanager1150    def null_ref_context(self):1151        """Context manager for handling null reference model (that is, peft adapter manipulation)."""1152        with (1153            self.accelerator.unwrap_model(self.model).disable_adapter()1154            if self.is_peft_model and not self.ref_adapter_name1155            else nullcontext()1156        ):1157            if self.ref_adapter_name:1158                self.model.set_adapter(self.ref_adapter_name)1159            yield1160            if self.ref_adapter_name:1161                self.model.set_adapter(self.model_adapter_name or "default")1162 1163    def compute_ref_log_probs(self, batch: dict[str, torch.LongTensor]) -> dict:1164        """Computes log probabilities of the reference model for a single padded batch of a DPO specific dataset."""1165        device_type = "xpu" if is_torch_xpu_available() else "cuda"1166        compte_ref_context_manager = amp.autocast(device_type) if self._peft_has_been_casted_to_bf16 else nullcontext()1167        with torch.no_grad(), compte_ref_context_manager:1168            if self.ref_model is None:1169                with self.null_ref_context():1170                    ref_model_output = self.concatenated_forward(self.model, batch)1171            else:1172                ref_model_output = self.concatenated_forward(self.ref_model, batch)1173        return ref_model_output["chosen_logps"], ref_model_output["rejected_logps"]1174 1175    @staticmethod1176    def concatenated_inputs(1177        batch: dict[str, Union[list, torch.LongTensor]], padding_value: int1178    ) -> dict[str, torch.LongTensor]:1179        """1180        Concatenate the `chosen` and `rejected` inputs from the batch into a single tensor for both the prompt1181        and completion sequences.1182 1183        Args:1184            batch (`dict[str, Union[list, torch.LongTensor]]`):1185                A batch of input data. The batch must contain the following keys:1186 1187                - `"prompt_input_ids"`: Tensor of shape `(batch_size, prompt_length)` representing the prompt input IDs.1188                - `"chosen_input_ids"`: Tensor of shape `(batch_size, chosen_length)` representing the chosen completion input IDs.1189                - `"rejected_input_ids"`: Tensor of shape `(batch_size, rejected_length)` representing the rejected completion input IDs.1190                - `"prompt_pixel_values"` (optional): Tensor for pixel values, if available.1191                - `"prompt_pixel_attention_mask"` (optional): Tensor for pixel attention masks, if available.1192 1193            padding_value (`int`):1194                The padding value to use for the concatenated completion sequences (`chosen_input_ids` and1195                `rejected_input_ids`).1196 1197        Returns:1198            `dict[str, torch.LongTensor]`: A dictionary containing:1199 1200                - `"prompt_input_ids"`: Concatenated prompt input IDs of shape `(2 * batch_size, prompt_length)`.

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