Team Ai
Apppublic

Zwounds/Boolean_Search_Query_Model

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

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