Team Ai
Apppublic

Zwounds/Boolean_Search_Query_Model

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
UnslothORPOTrainer.py1544 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.orpo_trainer import (Any, AutoModelForCausalLM, BaseImageProcessor, Callable, DPODataCollatorWithPadding, DataCollator, DataLoader, Dataset, EvalLoopOutput, F, FeatureExtractionMixin, Literal, ORPOConfig, ORPOTrainer, Optional, PartialState, PeftModel, PreTrainedModel, PreTrainedModelWrapper, PreTrainedTokenizerBase, ProcessorMixin, Trainer, TrainerCallback, Union, add_bos_token_if_needed, add_eos_token_if_needed, amp, deepcopy, defaultdict, disable_dropout_in_model, generate_model_card, get_comet_experiment_url, inspect, is_comet_available, is_peft_available, is_torch_fx_proxy, is_torch_xla_available, 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 UnslothORPOConfig(ORPOConfig):44    """45    46    Configuration class for the [`ORPOTrainer`].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 relative ratio loss weight in the ORPO loss. In the [paper](https://huggingface.co/papers/2403.07691),66            it is denoted by λ. In the [code](https://github.com/xfactlab/orpo), it is denoted by `alpha`.67        disable_dropout (`bool`, *optional*, defaults to `True`):68            Whether to disable dropout in the model.69        label_pad_token_id (`int`, *optional*, defaults to `-100`):70            Label pad token id. This argument is required if you want to use the default data collator.71        padding_value (`int` or `None`, *optional*, defaults to `None`):72            Padding value to use. If `None`, the padding value of the tokenizer is used.73        truncation_mode (`str`, *optional*, defaults to `"keep_end"`):74            Truncation mode to use when the prompt is too long. Possible values are `"keep_end"` or `"keep_start"`.75            This argument is required if you want to use the default data collator.76        generate_during_eval (`bool`, *optional*, defaults to `False`):77            If `True`, generates and logs completions from the model to W&B or Comet during evaluation.78        is_encoder_decoder (`bool` or `None`, *optional*, defaults to `None`):79            When using the `model_init` argument (callable) to instantiate the model instead of the `model` argument,80            you need to specify if the model returned by the callable is an encoder-decoder model.81        model_init_kwargs (`dict[str, Any]` or `None`, *optional*, defaults to `None`):82            Keyword arguments to pass to `AutoModelForCausalLM.from_pretrained` when instantiating the model from a83            string.84        dataset_num_proc (`int` or `None`, *optional*, defaults to `None`):85            Number of processes to use for processing the dataset.86    87    """88    vllm_sampling_params: Optional[Any] = field(89        default = None,90        metadata = {'help': 'vLLM SamplingParams'},91    )92    unsloth_num_chunks : Optional[int] = field(93        default = -1,94        metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},95    )96    def __init__(97        self,98        output_dir = None,99        overwrite_output_dir = None,100        do_train = False,101        do_eval = False,102        do_predict = False,103        eval_strategy = 'no',104        prediction_loss_only = False,105        per_device_train_batch_size = 4,106        per_device_eval_batch_size = 4,107        per_gpu_train_batch_size = None,108        per_gpu_eval_batch_size = None,109        gradient_accumulation_steps = 2,110        eval_accumulation_steps = 2,111        eval_delay = 0,112        torch_empty_cache_steps = 250,113        learning_rate = 5e-05,114        weight_decay = 0.01,115        adam_beta1 = 0.9,116        adam_beta2 = 0.999,117        adam_epsilon = 1e-08,118        max_grad_norm = 1.0,119        num_train_epochs = 3.0,120        max_steps = -1,121        lr_scheduler_type = 'linear',122        warmup_ratio = 0.1,123        warmup_steps = 0,124        log_level = 'passive',125        log_level_replica = 'warning',126        log_on_each_node = True,127        logging_dir = None,128        logging_strategy = 'steps',129        logging_first_step = False,130        logging_steps = 1,131        logging_nan_inf_filter = False,132        save_strategy = 'steps',133        save_steps = 500,134        save_total_limit = None,135        save_safetensors = True,136        save_on_each_node = False,137        save_only_model = False,138        restore_callback_states_from_checkpoint = False,139        no_cuda = False,140        use_cpu = False,141        use_mps_device = False,142        seed = 3407,143        data_seed = 3407,144        jit_mode_eval = False,145        use_ipex = False,146        bf16 = False,147        fp16 = False,148        fp16_opt_level = 'O1',149        half_precision_backend = 'auto',150        bf16_full_eval = False,151        fp16_full_eval = False,152        tf32 = None,153        local_rank = -1,154        ddp_backend = None,155        tpu_num_cores = None,156        tpu_metrics_debug = False,157        debug = '',158        dataloader_drop_last = False,159        eval_steps = None,160        dataloader_num_workers = 0,161        dataloader_prefetch_factor = None,162        past_index = -1,163        run_name = None,164        disable_tqdm = None,165        remove_unused_columns = True,166        label_names = None,167        load_best_model_at_end = False,168        metric_for_best_model = None,169        greater_is_better = None,170        ignore_data_skip = False,171        fsdp = '',172        fsdp_min_num_params = 0,173        fsdp_config = None,174        tp_size = 0,175        fsdp_transformer_layer_cls_to_wrap = None,176        accelerator_config = None,177        deepspeed = None,178        label_smoothing_factor = 0.0,179        optim = 'adamw_8bit',180        optim_args = None,181        adafactor = False,182        group_by_length = False,183        length_column_name = 'length',184        report_to = None,185        ddp_find_unused_parameters = None,186        ddp_bucket_cap_mb = None,187        ddp_broadcast_buffers = None,188        dataloader_pin_memory = True,189        dataloader_persistent_workers = False,190        skip_memory_metrics = True,191        use_legacy_prediction_loop = False,192        push_to_hub = False,193        resume_from_checkpoint = None,194        hub_model_id = None,195        hub_strategy = 'every_save',196        hub_token = None,197        hub_private_repo = None,198        hub_always_push = False,199        gradient_checkpointing = False,200        gradient_checkpointing_kwargs = None,201        include_inputs_for_metrics = False,202        eval_do_concat_batches = True,203        fp16_backend = 'auto',204        evaluation_strategy = None,205        push_to_hub_model_id = None,206        push_to_hub_organization = None,207        push_to_hub_token = None,208        mp_parameters = '',209        auto_find_batch_size = False,210        full_determinism = False,211        torchdynamo = None,212        ray_scope = 'last',213        ddp_timeout = 1800,214        torch_compile = False,215        torch_compile_backend = None,216        torch_compile_mode = None,217        dispatch_batches = None,218        split_batches = None,219        include_tokens_per_second = False,220        include_num_input_tokens_seen = False,221        neftune_noise_alpha = None,222        optim_target_modules = None,223        batch_eval_metrics = False,224        eval_on_start = False,225        use_liger_kernel = False,226        eval_use_gather_object = False,227        average_tokens_across_devices = False,228        max_length = 1024,229        max_prompt_length = 512,230        max_completion_length = None,231        beta = 0.1,232        disable_dropout = True,233        label_pad_token_id = -100,234        padding_value = None,235        truncation_mode = 'keep_end',236        generate_during_eval = False,237        is_encoder_decoder = None,238        model_init_kwargs = None,239        dataset_num_proc = None,240        vllm_sampling_params = None,241        unsloth_num_chunks = -1,242        **kwargs,243    ):244        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!')245        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!')246        if output_dir is None and save_strategy == 'steps' and save_steps == 500:247            output_dir = 'unsloth_training_checkpoints'248            save_strategy = 'no'249        if dataset_num_proc is None:250            from multiprocessing import cpu_count251            dataset_num_proc = cpu_count()252        253        super().__init__(254            output_dir = output_dir,255            overwrite_output_dir = overwrite_output_dir,256            do_train = do_train,257            do_eval = do_eval,258            do_predict = do_predict,259            eval_strategy = eval_strategy,260            prediction_loss_only = prediction_loss_only,261            per_device_train_batch_size = per_device_train_batch_size,262            per_device_eval_batch_size = per_device_eval_batch_size,263            per_gpu_train_batch_size = per_gpu_train_batch_size,264            per_gpu_eval_batch_size = per_gpu_eval_batch_size,265            gradient_accumulation_steps = gradient_accumulation_steps,266            eval_accumulation_steps = eval_accumulation_steps,267            eval_delay = eval_delay,268            torch_empty_cache_steps = torch_empty_cache_steps,269            learning_rate = learning_rate,270            weight_decay = weight_decay,271            adam_beta1 = adam_beta1,272            adam_beta2 = adam_beta2,273            adam_epsilon = adam_epsilon,274            max_grad_norm = max_grad_norm,275            num_train_epochs = num_train_epochs,276            max_steps = max_steps,277            lr_scheduler_type = lr_scheduler_type,278            warmup_ratio = warmup_ratio,279            warmup_steps = warmup_steps,280            log_level = log_level,281            log_level_replica = log_level_replica,282            log_on_each_node = log_on_each_node,283            logging_dir = logging_dir,284            logging_strategy = logging_strategy,285            logging_first_step = logging_first_step,286            logging_steps = logging_steps,287            logging_nan_inf_filter = logging_nan_inf_filter,288            save_strategy = save_strategy,289            save_steps = save_steps,290            save_total_limit = save_total_limit,291            save_safetensors = save_safetensors,292            save_on_each_node = save_on_each_node,293            save_only_model = save_only_model,294            restore_callback_states_from_checkpoint = restore_callback_states_from_checkpoint,295            no_cuda = no_cuda,296            use_cpu = use_cpu,297            use_mps_device = use_mps_device,298            seed = seed,299            data_seed = data_seed,300            jit_mode_eval = jit_mode_eval,301            use_ipex = use_ipex,302            bf16 = bf16,303            fp16 = fp16,304            fp16_opt_level = fp16_opt_level,305            half_precision_backend = half_precision_backend,306            bf16_full_eval = bf16_full_eval,307            fp16_full_eval = fp16_full_eval,308            tf32 = tf32,309            local_rank = local_rank,310            ddp_backend = ddp_backend,311            tpu_num_cores = tpu_num_cores,312            tpu_metrics_debug = tpu_metrics_debug,313            debug = debug,314            dataloader_drop_last = dataloader_drop_last,315            eval_steps = eval_steps,316            dataloader_num_workers = dataloader_num_workers,317            dataloader_prefetch_factor = dataloader_prefetch_factor,318            past_index = past_index,319            run_name = run_name,320            disable_tqdm = disable_tqdm,321            remove_unused_columns = remove_unused_columns,322            label_names = label_names,323            load_best_model_at_end = load_best_model_at_end,324            metric_for_best_model = metric_for_best_model,325            greater_is_better = greater_is_better,326            ignore_data_skip = ignore_data_skip,327            fsdp = fsdp,328            fsdp_min_num_params = fsdp_min_num_params,329            fsdp_config = fsdp_config,330            tp_size = tp_size,331            fsdp_transformer_layer_cls_to_wrap = fsdp_transformer_layer_cls_to_wrap,332            accelerator_config = accelerator_config,333            deepspeed = deepspeed,334            label_smoothing_factor = label_smoothing_factor,335            optim = optim,336            optim_args = optim_args,337            adafactor = adafactor,338            group_by_length = group_by_length,339            length_column_name = length_column_name,340            report_to = report_to,341            ddp_find_unused_parameters = ddp_find_unused_parameters,342            ddp_bucket_cap_mb = ddp_bucket_cap_mb,343            ddp_broadcast_buffers = ddp_broadcast_buffers,344            dataloader_pin_memory = dataloader_pin_memory,345            dataloader_persistent_workers = dataloader_persistent_workers,346            skip_memory_metrics = skip_memory_metrics,347            use_legacy_prediction_loop = use_legacy_prediction_loop,348            push_to_hub = push_to_hub,349            resume_from_checkpoint = resume_from_checkpoint,350            hub_model_id = hub_model_id,351            hub_strategy = hub_strategy,352            hub_token = hub_token,353            hub_private_repo = hub_private_repo,354            hub_always_push = hub_always_push,355            gradient_checkpointing = gradient_checkpointing,356            gradient_checkpointing_kwargs = gradient_checkpointing_kwargs,357            include_inputs_for_metrics = include_inputs_for_metrics,358            eval_do_concat_batches = eval_do_concat_batches,359            fp16_backend = fp16_backend,360            evaluation_strategy = evaluation_strategy,361            push_to_hub_model_id = push_to_hub_model_id,362            push_to_hub_organization = push_to_hub_organization,363            push_to_hub_token = push_to_hub_token,364            mp_parameters = mp_parameters,365            auto_find_batch_size = auto_find_batch_size,366            full_determinism = full_determinism,367            torchdynamo = torchdynamo,368            ray_scope = ray_scope,369            ddp_timeout = ddp_timeout,370            torch_compile = torch_compile,371            torch_compile_backend = torch_compile_backend,372            torch_compile_mode = torch_compile_mode,373            dispatch_batches = dispatch_batches,374            split_batches = split_batches,375            include_tokens_per_second = include_tokens_per_second,376            include_num_input_tokens_seen = include_num_input_tokens_seen,377            neftune_noise_alpha = neftune_noise_alpha,378            optim_target_modules = optim_target_modules,379            batch_eval_metrics = batch_eval_metrics,380            eval_on_start = eval_on_start,381            use_liger_kernel = use_liger_kernel,382            eval_use_gather_object = eval_use_gather_object,383            average_tokens_across_devices = average_tokens_across_devices,384            max_length = max_length,385            max_prompt_length = max_prompt_length,386            max_completion_length = max_completion_length,387            beta = beta,388            disable_dropout = disable_dropout,389            label_pad_token_id = label_pad_token_id,390            padding_value = padding_value,391            truncation_mode = truncation_mode,392            generate_during_eval = generate_during_eval,393            is_encoder_decoder = is_encoder_decoder,394            model_init_kwargs = model_init_kwargs,395            dataset_num_proc = dataset_num_proc,**kwargs)396        self.vllm_sampling_params = vllm_sampling_params397        self.unsloth_num_chunks = unsloth_num_chunks398pass399 400class _UnslothORPOTrainer(Trainer):401    r""""""402 403    _tag_names = ["trl", "orpo"]404 405    def __init__(406        self,407        model: Optional[Union[PreTrainedModel, nn.Module, str]] = None,408        args: Optional[ORPOConfig] = None,409        data_collator: Optional[DataCollator] = None,410        train_dataset: Optional[Dataset] = None,411        eval_dataset: Optional[Union[Dataset, dict[str, Dataset]]] = None,412        processing_class: Optional[413            Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin]414        ] = None,415        model_init: Optional[Callable[[], PreTrainedModel]] = None,416        callbacks: Optional[list[TrainerCallback]] = None,417        optimizers: tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),418        preprocess_logits_for_metrics: Optional[Callable[[torch.Tensor, torch.Tensor], torch.Tensor]] = None,419        peft_config: Optional[dict] = None,420        compute_metrics: Optional[Callable[[EvalLoopOutput], dict]] = None,421    ):422        if args.model_init_kwargs is None:423            model_init_kwargs = {}424        elif not isinstance(model, str):425            raise ValueError("You passed model_kwargs to the ORPOTrainer. But your model is already instantiated.")426        else:427            model_init_kwargs = args.model_init_kwargs428            torch_dtype = model_init_kwargs.get("torch_dtype")429            if torch_dtype is not None:430                # Convert to `torch.dtype` if an str is passed431                if isinstance(torch_dtype, str) and torch_dtype != "auto":432                    torch_dtype = getattr(torch, torch_dtype)433                if torch_dtype != "auto" and not isinstance(torch_dtype, torch.dtype):434                    raise ValueError(435                        f"Invalid `torch_dtype` passed to the ORPOConfig. Expected a string with either `torch.dtype` or 'auto', but got {torch_dtype}."436                    )437                model_init_kwargs["torch_dtype"] = torch_dtype438 439        if isinstance(model, str):440            model = AutoModelForCausalLM.from_pretrained(model, **model_init_kwargs)441 442        # Initialize this variable to False. This helps tracking the case when `peft_module_casting_to_bf16`443        # has been called in order to properly call autocast if needed.444        self._peft_has_been_casted_to_bf16 = False445 446        if not is_peft_available() and peft_config is not None:447            raise ValueError(448                "PEFT is not installed and you passed a `peft_config` in the trainer's kwargs, please install it to use the PEFT models"449            )450        elif is_peft_available() and peft_config is not None:451            # if model is a peft model and we have a peft_config, we merge and unload it first452            if isinstance(model, PeftModel):453                model = model.merge_and_unload()454 455            if getattr(model, "is_loaded_in_8bit", False) or getattr(model, "is_loaded_in_4bit", False):456                _support_gc_kwargs = hasattr(457                    args, "gradient_checkpointing_kwargs"458                ) and "gradient_checkpointing_kwargs" in list(459                    inspect.signature(prepare_model_for_kbit_training).parameters460                )461 462                prepare_model_kwargs = {"use_gradient_checkpointing": args.gradient_checkpointing}463 464                if _support_gc_kwargs:465                    prepare_model_kwargs["gradient_checkpointing_kwargs"] = args.gradient_checkpointing_kwargs466 467                model = prepare_model_for_kbit_training(model, **prepare_model_kwargs)468            elif getattr(args, "gradient_checkpointing", False):469                # For backward compatibility with older versions of transformers470                if hasattr(model, "enable_input_require_grads"):471                    model.enable_input_require_grads()472                else:473 474                    def make_inputs_require_grad(module, input, output):475                        output.requires_grad_(True)476 477                    model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)478 479            # get peft model with the given config480            model = model481            if args.bf16 and getattr(model, "is_loaded_in_4bit", False):482                peft_module_casting_to_bf16(model)483                # If args.bf16 we need to explicitly call `generate` with torch amp autocast context manager484                self._peft_has_been_casted_to_bf16 = True485 486        # For models that use gradient_checkpointing, we need to attach a hook that enables input487        # to explicitly have `requires_grad=True`, otherwise training will either silently488        # fail or completely fail.489        elif getattr(args, "gradient_checkpointing", False):490            # For backward compatibility with older versions of transformers491            if hasattr(model, "enable_input_require_grads"):492                model.enable_input_require_grads()493            else:494 495                def make_inputs_require_grad(module, input, output):496                    output.requires_grad_(True)497 498                model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)499 500        if args.generate_during_eval and not (is_wandb_available() or is_comet_available()):501            raise ValueError(502                "`generate_during_eval=True` requires Weights and Biases or Comet to be installed."503                " Please install `wandb` or `comet-ml` to resolve."504            )505 506        if model is not None:507            self.is_encoder_decoder = model.config.is_encoder_decoder508        elif args.is_encoder_decoder is None:509            raise ValueError("When no model is provided, you need to pass the parameter is_encoder_decoder.")510        else:511            self.is_encoder_decoder = args.is_encoder_decoder512 513        if self.is_encoder_decoder:514            self.decoder_start_token_id = model.config.decoder_start_token_id515            self.pad_token_id = model.config.pad_token_id516 517        if processing_class is None:518            raise ValueError("processing_class must be specified to tokenize a ORPO dataset.")519        if args.max_length is None:520            warnings.warn(521                "`max_length` is not set in the ORPOConfig's init"522                " it will default to `512` by default, but you should do it yourself in the future.",523                UserWarning,524            )525            max_length = 512526        else:527            max_length = args.max_length528        if args.max_prompt_length is None:529            warnings.warn(530                "`max_prompt_length` is not set in the ORPOConfig's init"531                " it will default to `128` by default, but you should do it yourself in the future.",532                UserWarning,533            )534            max_prompt_length = 128535        else:536            max_prompt_length = args.max_prompt_length537 538        if args.max_completion_length is None and self.is_encoder_decoder:539            warnings.warn(540                "When using an encoder decoder architecture, you should set `max_completion_length` in the ORPOConfig's init"541                " it will default to `128` by default, but you should do it yourself in the future.",542                UserWarning,543            )544            self.max_completion_length = 128545        else:546            self.max_completion_length = args.max_completion_length547 548        if data_collator is None:549            data_collator = DPODataCollatorWithPadding(550                pad_token_id=processing_class.pad_token_id,551                label_pad_token_id=args.label_pad_token_id,552                is_encoder_decoder=self.is_encoder_decoder,553            )554 555            if args.remove_unused_columns:556                args.remove_unused_columns = False557                # warn users558                warnings.warn(559                    "When using DPODataCollatorWithPadding, you should set `remove_unused_columns=False` in your TrainingArguments"560                    " we have set it for you, but you should do it yourself in the future.",561                    UserWarning,562                )563 564            self.use_dpo_data_collator = True565        else:566            self.use_dpo_data_collator = False567 568        # Disable dropout in the model and reference model569        if args.disable_dropout:570            disable_dropout_in_model(model)571 572        self.max_length = max_length573        self.generate_during_eval = args.generate_during_eval574        self.label_pad_token_id = args.label_pad_token_id575        self.padding_value = args.padding_value if args.padding_value is not None else processing_class.pad_token_id576        self.max_prompt_length = max_prompt_length577        self.truncation_mode = args.truncation_mode578        self.processing_class = processing_class579 580        self.beta = args.beta581        self.aux_loss_enabled = getattr(model.config, "output_router_logits", False)582        self.aux_loss_coef = getattr(model.config, "router_aux_loss_coef", 0.0)583        if self.aux_loss_enabled and self.aux_loss_coef == 0.0:584            warnings.warn(585                "You set `output_router_logits` to `True` in the model config, but `router_aux_loss_coef` is set to "586                "`0.0`, meaning the auxiliary loss will not be used. Either set `router_aux_loss_coef` to a value "587                "greater than `0.0`, or set `output_router_logits` to `False` if you don't want to use the auxiliary "588                "loss.",589                UserWarning,590            )591 592        self._stored_metrics = defaultdict(lambda: defaultdict(list))593 594        # The trainer estimates the number of FLOPs (floating-point operations) using the number of elements in the595        # input tensor associated with the key "input_ids". However, in ORPO, the sampled data does not include the596        # "input_ids" key. Instead, the available keys are "prompt_input_ids", "chosen_input_ids", and597        # "rejected_input_ids". As a result, the trainer issues the warning: "Could not estimate the number of tokens598        # of the input, floating-point operations will not be computed." To suppress this warning, we set the599        # "estimate_tokens" key in the model's "warnings_issued" dictionary to True. This acts as a flag to indicate600        # that the warning has already been issued.601        model.warnings_issued["estimate_tokens"] = True602 603        # Compute that only on the main process for faster data processing.604        # see: https://github.com/huggingface/trl/pull/1255605        with PartialState().local_main_process_first():606            # Extract the prompt if needed, and apply the chat template if needed607            train_dataset = train_dataset.map(maybe_extract_prompt, num_proc=args.dataset_num_proc)608            train_dataset = train_dataset.map(609                maybe_apply_chat_template, fn_kwargs={"tokenizer": processing_class}, num_proc=args.dataset_num_proc610            )611            train_dataset = train_dataset.map(self.tokenize_row, num_proc=args.dataset_num_proc)612            if eval_dataset is not None:613                eval_dataset = eval_dataset.map(maybe_extract_prompt, num_proc=args.dataset_num_proc)614                eval_dataset = eval_dataset.map(615                    maybe_apply_chat_template,616                    fn_kwargs={"tokenizer": processing_class},617                    num_proc=args.dataset_num_proc,618                )619                eval_dataset = eval_dataset.map(self.tokenize_row, num_proc=args.dataset_num_proc)620 621        super().__init__(622            model=model,623            args=args,624            data_collator=data_collator,625            train_dataset=train_dataset,626            eval_dataset=eval_dataset,627            processing_class=processing_class,628            model_init=model_init,629            compute_metrics=compute_metrics,630            callbacks=callbacks,631            optimizers=optimizers,632            preprocess_logits_for_metrics=preprocess_logits_for_metrics,633        )634 635        # Add tags for models that have been loaded with the correct transformers version636        if hasattr(self.model, "add_model_tags"):637            self.model.add_model_tags(self._tag_names)638 639        if not hasattr(self, "accelerator"):640            raise AttributeError(641                "Your `Trainer` does not have an `accelerator` object. Consider upgrading `transformers`."642            )643 644    def _prepare_deepspeed(self, model: PreTrainedModelWrapper):645        # Adapted from accelerate: https://github.com/huggingface/accelerate/blob/739b135f8367becb67ffaada12fe76e3aa60fefd/src/accelerate/accelerator.py#L1473646        deepspeed_plugin = self.accelerator.state.deepspeed_plugin647        config_kwargs = deepcopy(deepspeed_plugin.deepspeed_config)648 649        if model is not None:650            if hasattr(model, "config"):651                hidden_size = (652                    max(model.config.hidden_sizes)653                    if getattr(model.config, "hidden_sizes", None)654                    else getattr(model.config, "hidden_size", None)655                )656                if hidden_size is not None and config_kwargs["zero_optimization"]["stage"] == 3:657                    # Note that `stage3_prefetch_bucket_size` can produce DeepSpeed messages like: `Invalidate trace cache @ step 0: expected module 1, but got module 0`658                    # This is expected and is not an error, see: https://github.com/microsoft/DeepSpeed/discussions/4081659                    config_kwargs.update(660                        {661                            "zero_optimization.reduce_bucket_size": hidden_size * hidden_size,662                            "zero_optimization.stage3_param_persistence_threshold": 10 * hidden_size,663                            "zero_optimization.stage3_prefetch_bucket_size": 0.9 * hidden_size * hidden_size,664                        }665                    )666 667        # If ZeRO-3 is used, we shard both the active and reference model.668        # Otherwise, we assume the reference model fits in memory and is initialized on each device with ZeRO disabled (stage 0)669        if config_kwargs["zero_optimization"]["stage"] != 3:670            config_kwargs["zero_optimization"]["stage"] = 0671        model, *_ = deepspeed.initialize(model=model, config=config_kwargs)672        model.eval()673        return model674 675    def build_tokenized_answer(self, prompt, answer):676        """677        Llama tokenizer does satisfy `enc(a + b) = enc(a) + enc(b)`.678        It does ensure `enc(a + b) = enc(a) + enc(a + b)[len(enc(a)):]`.679        Reference:680            https://github.com/EleutherAI/lm-evaluation-harness/pull/531#issuecomment-1595586257681        """682 683        full_tokenized = self.processing_class(prompt + answer, add_special_tokens=False)684        prompt_input_ids = self.processing_class(prompt, add_special_tokens=False)["input_ids"]685 686        answer_input_ids = full_tokenized["input_ids"][len(prompt_input_ids) :]687        answer_attention_mask = full_tokenized["attention_mask"][len(prompt_input_ids) :]688 689        # Concat tokens to form `enc(a) + enc(a + b)[len(enc(a)):]`690        full_concat_input_ids = np.concatenate([prompt_input_ids, answer_input_ids])691 692        # Prepare input tokens for token by token comparison693        full_input_ids = np.array(full_tokenized["input_ids"])694 695        if len(full_input_ids) != len(full_concat_input_ids):696            raise ValueError("Prompt input ids and answer input ids should have the same length.")697 698        # On some tokenizers, like Llama-2 tokenizer, there are occasions where tokens699        # can be merged together when tokenizing prompt+answer. This could result700        # on the last token from the prompt being different when tokenized on its own701        # vs when done as prompt+answer.702        response_token_ids_start_idx = len(prompt_input_ids)703 704        # If tokenized prompt is different than both prompt+answer, then it means the705        # last token has changed due to merging.706        if prompt_input_ids != full_tokenized["input_ids"][:response_token_ids_start_idx]:707            response_token_ids_start_idx -= 1708 709        prompt_input_ids = full_tokenized["input_ids"][:response_token_ids_start_idx]710        prompt_attention_mask = full_tokenized["attention_mask"][:response_token_ids_start_idx]711 712        if len(prompt_input_ids) != len(prompt_attention_mask):713            raise ValueError("Prompt input ids and attention mask should have the same length.")714 715        answer_input_ids = full_tokenized["input_ids"][response_token_ids_start_idx:]716        answer_attention_mask = full_tokenized["attention_mask"][response_token_ids_start_idx:]717 718        return dict(719            prompt_input_ids=prompt_input_ids,720            prompt_attention_mask=prompt_attention_mask,721            input_ids=answer_input_ids,722            attention_mask=answer_attention_mask,723        )724 725    def tokenize_row(self, feature, model: Optional[Union[PreTrainedModel, nn.Module]] = None) -> dict:726        """Tokenize a single row from a ORPO specific dataset.727 728        At this stage, we don't convert to PyTorch tensors yet; we just handle the truncation729        in case the prompt + chosen or prompt + rejected responses is/are too long. First730        we truncate the prompt; if we're still too long, we truncate the chosen/rejected.731 732        We also create the labels for the chosen/rejected responses, which are of length equal to733        the sum of the length of the prompt and the chosen/rejected response, with734        label_pad_token_id  for the prompt tokens.735        """736        batch = {}737        prompt = feature["prompt"]738        chosen = feature["chosen"]739        rejected = feature["rejected"]740 741        if not self.is_encoder_decoder:742            # Check issues below for more details743            #  1. https://github.com/huggingface/trl/issues/907744            #  2. https://github.com/EleutherAI/lm-evaluation-harness/pull/531#issuecomment-1595586257745            #  3. https://github.com/LianjiaTech/BELLE/issues/337746 747            if not isinstance(prompt, str):748                raise ValueError(f"prompt should be an str but got {type(prompt)}")749            prompt_tokens = self.processing_class(prompt, add_special_tokens=False)750            prompt_tokens = {f"prompt_{k}": v for k, v in prompt_tokens.items()}751 752            if not isinstance(chosen, str):753                raise ValueError(f"chosen should be an str but got {type(chosen)}")754            chosen_tokens = self.build_tokenized_answer(prompt, chosen)755 756            if not isinstance(rejected, str):757                raise ValueError(f"rejected should be an str but got {type(rejected)}")758            rejected_tokens = self.build_tokenized_answer(prompt, rejected)759 760            # Last prompt token might get merged by tokenizer and761            # it should not be included for generation if that happens762            prompt_len_input_ids = len(prompt_tokens["prompt_input_ids"])763 764            chosen_prompt_len_input_ids = len(chosen_tokens["prompt_input_ids"])765            rejected_prompt_len_input_ids = len(rejected_tokens["prompt_input_ids"])766            prompt_len_input_ids = min(chosen_prompt_len_input_ids, rejected_prompt_len_input_ids)767 768            for k, v in prompt_tokens.items():769                prompt_tokens[k] = v[:prompt_len_input_ids]770 771            # Make sure prompts only have one different token at most an772            # and length only differs by 1 at most773            num_diff_tokens = sum(774                [a != b for a, b in zip(chosen_tokens["prompt_input_ids"], rejected_tokens["prompt_input_ids"])]775            )776            num_diff_len = abs(chosen_prompt_len_input_ids - rejected_prompt_len_input_ids)777            if num_diff_tokens > 1 or num_diff_len > 1:778                raise ValueError(779                    "Chosen and rejected prompt_input_ids might only differ on the "780                    "last token due to tokenizer merge ops."781                )782 783            # add BOS token to head of prompt. Avoid adding if it's already there784            prompt_tokens, chosen_tokens, rejected_tokens = add_bos_token_if_needed(785                self.processing_class.bos_token_id,786                prompt_len_input_ids,787                prompt_tokens,788                chosen_prompt_len_input_ids,789                chosen_tokens,790                rejected_prompt_len_input_ids,791                rejected_tokens,792            )793 794            # add EOS token to end of answer. Avoid adding if it's already there795            chosen_tokens, rejected_tokens = add_eos_token_if_needed(796                self.processing_class.eos_token_id, chosen_tokens, rejected_tokens797            )798 799            longer_response_length = max(len(chosen_tokens["input_ids"]), len(rejected_tokens["input_ids"]))800 801            # if combined sequence is too long, truncate the prompt802            for answer_tokens in [chosen_tokens, rejected_tokens, prompt_tokens]:803                if len(answer_tokens["prompt_input_ids"]) + longer_response_length > self.max_length:804                    if self.truncation_mode == "keep_start":805                        for k in ["prompt_input_ids", "prompt_attention_mask"]:806                            answer_tokens[k] = answer_tokens[k][: self.max_prompt_length]807                    elif self.truncation_mode == "keep_end":808                        for k in ["prompt_input_ids", "prompt_attention_mask"]:809                            answer_tokens[k] = answer_tokens[k][-self.max_prompt_length :]810                    else:811                        raise ValueError(f"Unknown truncation mode: {self.truncation_mode}")812 813            # if that's still too long, truncate the response814            for answer_tokens in [chosen_tokens, rejected_tokens]:815                if len(answer_tokens["prompt_input_ids"]) + longer_response_length > self.max_length:816                    for k in ["input_ids", "attention_mask"]:817                        answer_tokens[k] = answer_tokens[k][: self.max_length - self.max_prompt_length]818 819            # Create labels820            chosen_sequence_tokens = {821                k: chosen_tokens[f"prompt_{k}"] + chosen_tokens[k] for k in ["input_ids", "attention_mask"]822            }823            rejected_sequence_tokens = {824                k: rejected_tokens[f"prompt_{k}"] + rejected_tokens[k] for k in ["input_ids", "attention_mask"]825            }826            chosen_sequence_tokens["labels"] = chosen_sequence_tokens["input_ids"][:]827            chosen_sequence_tokens["labels"][: len(chosen_tokens["prompt_input_ids"])] = [828                self.label_pad_token_id829            ] * len(chosen_tokens["prompt_input_ids"])830            rejected_sequence_tokens["labels"] = rejected_sequence_tokens["input_ids"][:]831            rejected_sequence_tokens["labels"][: len(rejected_tokens["prompt_input_ids"])] = [832                self.label_pad_token_id833            ] * len(rejected_tokens["prompt_input_ids"])834 835            for k, toks in {836                "chosen_": chosen_sequence_tokens,837                "rejected_": rejected_sequence_tokens,838                "": prompt_tokens,839            }.items():840                for type_key, tokens in toks.items():841                    if type_key == "token_type_ids":842                        continue843                    batch[f"{k}{type_key}"] = tokens844 845        else:846            chosen_tokens = self.processing_class(847                chosen, truncation=True, max_length=self.max_completion_length, add_special_tokens=True848            )849            rejected_tokens = self.processing_class(850                rejected, truncation=True, max_length=self.max_completion_length, add_special_tokens=True851            )852            prompt_tokens = self.processing_class(853                prompt, truncation=True, max_length=self.max_prompt_length, add_special_tokens=True854            )855 856            batch["chosen_labels"] = chosen_tokens["input_ids"]857            batch["rejected_labels"] = rejected_tokens["input_ids"]858            batch["prompt_input_ids"] = prompt_tokens["input_ids"]859            batch["prompt_attention_mask"] = prompt_tokens["attention_mask"]860 861            if model is not None and hasattr(model, "prepare_decoder_input_ids_from_labels"):862                batch["rejected_decoder_input_ids"] = model.prepare_decoder_input_ids_from_labels(863                    labels=torch.tensor(batch["rejected_labels"])864                )865                batch["chosen_decoder_input_ids"] = model.prepare_decoder_input_ids_from_labels(866                    labels=torch.tensor(batch["chosen_labels"])867                )868 869        if is_torch_xla_available():870            # Pad the sequences to global max_length to avoid TorchXLA recompilation871            for k in batch:872                if "labels" in k or self.is_encoder_decoder:873                    pad_value = self.label_pad_token_id874                elif k.endswith("_input_ids"):875                    pad_value = self.padding_value876                elif k.endswith("_attention_mask"):877                    pad_value = 0878                batch[k] = batch[k] + [pad_value] * (self.max_length - len(batch[k]))879        return batch880 881    @staticmethod882    def concatenated_inputs(883        batch: dict[str, Union[list, torch.LongTensor]],884        is_encoder_decoder: bool = False,885        label_pad_token_id: int = -100,886        padding_value: int = 0,887        device: Optional[torch.device] = None,888    ) -> dict[str, torch.LongTensor]:889        """Concatenate the chosen and rejected inputs into a single tensor.890 891        Args:892            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).893            is_encoder_decoder: Whether the model is an encoder-decoder model.894            label_pad_token_id: The label pad token id.895            padding_value: The padding value to use for the concatenated inputs_ids.896            device: The device for the concatenated inputs.897 898        Returns:899            A dictionary containing the concatenated inputs under the key 'concatenated_input_ids'.900        """901        concatenated_batch = {}902 903        if is_encoder_decoder:904            max_length = max(batch["chosen_labels"].shape[1], batch["rejected_labels"].shape[1])905        else:906            max_length = max(batch["chosen_input_ids"].shape[1], batch["rejected_input_ids"].shape[1])907 908        for k in batch:909            if k.startswith("chosen") and isinstance(batch[k], torch.Tensor):910                if "labels" in k or is_encoder_decoder:911                    pad_value = label_pad_token_id912                elif k.endswith("_input_ids"):913                    pad_value = padding_value914                elif k.endswith("_attention_mask"):915                    pad_value = 0916                concatenated_key = k.replace("chosen", "concatenated")917                concatenated_batch[concatenated_key] = pad_to_length(batch[k], max_length, pad_value=pad_value)918        for k in batch:919            if k.startswith("rejected") and isinstance(batch[k], torch.Tensor):920                if "labels" in k or is_encoder_decoder:921                    pad_value = label_pad_token_id922                elif k.endswith("_input_ids"):923                    pad_value = padding_value924                elif k.endswith("_attention_mask"):925                    pad_value = 0926                concatenated_key = k.replace("rejected", "concatenated")927                concatenated_batch[concatenated_key] = torch.cat(928                    (929                        concatenated_batch[concatenated_key],930                        pad_to_length(batch[k], max_length, pad_value=pad_value),931                    ),932                    dim=0,933                ).to(device=device)934 935        if is_encoder_decoder:936            concatenated_batch["concatenated_input_ids"] = batch["prompt_input_ids"].repeat(2, 1).to(device=device)937            concatenated_batch["concatenated_attention_mask"] = (938                batch["prompt_attention_mask"].repeat(2, 1).to(device=device)939            )940 941        return concatenated_batch942 943    def odds_ratio_loss(944        self,945        policy_chosen_logps: torch.FloatTensor,946        policy_rejected_logps: torch.FloatTensor,947    ) -> tuple[torch.FloatTensor, torch.FloatTensor, torch.FloatTensor, torch.FloatTensor, torch.FloatTensor]:948        """Compute ORPO's odds ratio (OR) loss for a batch of policy and reference model log probabilities.949 950        Args:951            policy_chosen_logps: Log probabilities of the policy model for the chosen responses. Shape: (batch_size,)952            policy_rejected_logps: Log probabilities of the policy model for the rejected responses. Shape: (batch_size,)953 954        Returns:955            A tuple of three tensors: (losses, chosen_rewards, rejected_rewards).956            The losses tensor contains the ORPO loss for each example in the batch.957            The chosen_rewards and rejected_rewards tensors contain the rewards for the chosen and rejected responses, respectively.958            The log odds ratio of the chosen responses over the rejected responses ratio for logging purposes.959            The `log(sigmoid(log_odds_chosen))` for logging purposes.960        """961 962        # Derived from Eqs. (4) and (7) from https://huggingface.co/papers/2403.07691 by using log identities and exp(log(P(y|x)) = P(y|x)963        log_odds = (policy_chosen_logps - policy_rejected_logps) - (964            torch.log1p(-torch.exp(policy_chosen_logps)) - torch.log1p(-torch.exp(policy_rejected_logps))965        )966        ratio = F.logsigmoid(log_odds)967        losses = self.beta * ratio968 969        chosen_rewards = self.beta * (policy_chosen_logps.to(self.accelerator.device)).detach()970        rejected_rewards = self.beta * (policy_rejected_logps.to(self.accelerator.device)).detach()971 972        return losses, chosen_rewards, rejected_rewards, torch.mean(ratio), torch.mean(log_odds)973 974    @staticmethod975    def get_batch_logps(976        logits: torch.FloatTensor,977        labels: torch.LongTensor,978        average_log_prob: bool = False,979        label_pad_token_id: int = -100,980        is_encoder_decoder: bool = False,981    ) -> torch.FloatTensor:982        """Compute the log probabilities of the given labels under the given logits.983 984        Args:985            logits: Logits of the model (unnormalized). Shape: (batch_size, sequence_length, vocab_size)986            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)987            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.988            label_pad_token_id: The label pad token id.989            is_encoder_decoder: Whether the model is an encoder-decoder model.990 991        Returns:992            A tensor of shape (batch_size,) containing the average/sum log probabilities of the given labels under the given logits.993        """994        if logits.shape[:-1] != labels.shape:995            raise ValueError("Logits (batch and sequence length dim) and labels must have the same shape.")996 997        if not is_encoder_decoder:998            labels = labels[:, 1:].clone()999            logits = logits[:, :-1, :]1000        loss_mask = labels != label_pad_token_id1001 1002        # dummy token; we'll ignore the losses on these tokens later1003        labels = torch.where(labels == label_pad_token_id, 0, labels)1004 1005        per_token_logps = selective_log_softmax(logits, labels)1006 1007        if average_log_prob:1008            return (per_token_logps * loss_mask).sum(-1) / loss_mask.sum(-1)1009        else:1010            return (per_token_logps * loss_mask).sum(-1)1011 1012    def concatenated_forward(1013        self, model: nn.Module, batch: dict[str, Union[list, torch.LongTensor]]1014    ) -> tuple[torch.FloatTensor, torch.FloatTensor, torch.FloatTensor, torch.FloatTensor]:1015        """Run the given model on the given batch of inputs, concatenating the chosen and rejected inputs together.1016 1017        We do this to avoid doing two forward passes, because it's faster for FSDP.1018        """1019        concatenated_batch = self.concatenated_inputs(1020            batch,1021            is_encoder_decoder=self.is_encoder_decoder,1022            label_pad_token_id=self.label_pad_token_id,1023            padding_value=self.padding_value,1024            device=self.accelerator.device,1025        )1026        len_chosen = batch["chosen_labels"].shape[0]1027 1028        model_kwargs = (1029            {1030                "decoder_input_ids": self._shift_right(concatenated_batch["concatenated_labels"]),1031            }1032            if self.is_encoder_decoder1033            else {}1034        )1035 1036        if self.aux_loss_enabled:1037            model_kwargs["output_router_logits"] = True1038 1039        outputs = model(1040            concatenated_batch["concatenated_input_ids"],1041            attention_mask=concatenated_batch["concatenated_attention_mask"],1042            use_cache=False,1043            **model_kwargs,1044        )1045        all_logits = outputs.logits1046 1047        def cross_entropy_loss(logits, labels):1048            if not self.is_encoder_decoder:1049                # Shift so that tokens < n predict n1050                logits = logits[..., :-1, :].contiguous()1051                labels = labels[..., 1:].contiguous()1052            # Flatten the tokens1053            loss_fct = nn.CrossEntropyLoss()1054            logits = logits.view(-1, logits.shape[-1])1055            labels = labels.view(-1)1056            # Enable model parallelism1057            labels = labels.to(logits.device)1058            loss = loss_fct(logits, labels)1059            return loss1060 1061        if self.is_encoder_decoder:1062            labels = concatenated_batch["concatenated_labels"].clone()1063        else:1064            labels = concatenated_batch["concatenated_input_ids"].clone()1065            attention_mask = concatenated_batch["concatenated_attention_mask"]1066            labels = torch.where(attention_mask == 1, labels, self.label_pad_token_id)1067        # orpo chosen nll loss is computed over the full prompt and response1068        chosen_nll_loss = cross_entropy_loss(all_logits[:len_chosen], labels[:len_chosen])1069 1070        all_logps = self.get_batch_logps(1071            all_logits,1072            concatenated_batch["concatenated_labels"],1073            average_log_prob=True,1074            is_encoder_decoder=self.is_encoder_decoder,1075            label_pad_token_id=self.label_pad_token_id,1076        )1077 1078        chosen_logps = all_logps[:len_chosen]1079        rejected_logps = all_logps[len_chosen:]1080 1081        if not self.is_encoder_decoder:1082            chosen_logits = all_logits[:len_chosen, :-1, :]1083            rejected_logits = all_logits[len_chosen:, :-1, :]1084        else:1085            chosen_logits = all_logits[:len_chosen]1086            rejected_logits = all_logits[len_chosen:]1087 1088        if self.aux_loss_enabled:1089            return (chosen_logps, rejected_logps, chosen_logits, rejected_logits, chosen_nll_loss, outputs.aux_loss)1090 1091        return (chosen_logps, rejected_logps, chosen_logits, rejected_logits, chosen_nll_loss)1092 1093    def get_batch_loss_metrics(1094        self,1095        model,1096        batch: dict[str, Union[list, torch.LongTensor]],1097        train_eval: Literal["train", "eval"] = "train",1098    ):1099        """Compute the ORPO loss and other metrics for the given batch of inputs for train or test."""1100        metrics = {}1101 1102        forward_output = self.concatenated_forward(model, batch)1103        (1104            policy_chosen_logps,1105            policy_rejected_logps,1106            policy_chosen_logits,1107            policy_rejected_logits,1108            policy_nll_loss,1109        ) = forward_output[:5]1110        if self.aux_loss_enabled:1111            aux_loss = forward_output[5]1112 1113        losses, chosen_rewards, rejected_rewards, log_odds_ratio, log_odds_chosen = self.odds_ratio_loss(1114            policy_chosen_logps, policy_rejected_logps1115        )1116        # full ORPO loss1117        loss = policy_nll_loss - losses.mean()1118 1119        reward_accuracies = (chosen_rewards > rejected_rewards).float()1120 1121        prefix = "eval_" if train_eval == "eval" else ""1122        metrics[f"{prefix}rewards/chosen"] = self.accelerator.gather_for_metrics(chosen_rewards).mean()1123        metrics[f"{prefix}rewards/rejected"] = self.accelerator.gather_for_metrics(rejected_rewards).mean()1124        metrics[f"{prefix}rewards/accuracies"] = self.accelerator.gather_for_metrics(reward_accuracies).mean()1125        metrics[f"{prefix}rewards/margins"] = self.accelerator.gather_for_metrics(1126            chosen_rewards - rejected_rewards1127        ).mean()1128        metrics[f"{prefix}logps/rejected"] = self.accelerator.gather_for_metrics(policy_rejected_logps).detach().mean()1129        metrics[f"{prefix}logps/chosen"] = self.accelerator.gather_for_metrics(policy_chosen_logps).detach().mean()1130        metrics[f"{prefix}logits/rejected"] = (1131            self.accelerator.gather_for_metrics(policy_rejected_logits).detach().mean()1132        )1133        metrics[f"{prefix}logits/chosen"] = self.accelerator.gather_for_metrics(policy_chosen_logits).detach().mean()1134        metrics[f"{prefix}nll_loss"] = self.accelerator.gather_for_metrics(policy_nll_loss).detach().mean()1135        metrics[f"{prefix}log_odds_ratio"] = self.accelerator.gather_for_metrics(log_odds_ratio).mean()1136        metrics[f"{prefix}log_odds_chosen"] = self.accelerator.gather_for_metrics(log_odds_chosen).mean()1137        if is_torch_xla_available():1138            xm.mark_step()  # needed because .item() calls1139        for k, v in metrics.items():1140            metrics[k] = v.item()1141        if self.aux_loss_enabled:1142            loss += self.aux_loss_coef * aux_loss1143 1144        return loss, metrics1145 1146    def compute_loss(1147        self,1148        model: Union[PreTrainedModel, nn.Module],1149        inputs: dict[str, Union[torch.Tensor, Any]],1150        return_outputs=False,1151        num_items_in_batch=None,1152    ) -> Union[torch.Tensor, tuple[torch.Tensor, dict[str, torch.Tensor]]]:1153        compute_loss_context_manager = amp.autocast("cuda") if self._peft_has_been_casted_to_bf16 else nullcontext()1154 1155        with compute_loss_context_manager:1156            loss, metrics = self.get_batch_loss_metrics(model, inputs, train_eval="train")1157 1158        # Make sure to move the loss to the device the original accumulating loss is at back in the `Trainer` class:1159        loss = loss.to(self.args.device)1160 1161        # force log the metrics1162        self.store_metrics(metrics, train_eval="train")1163 1164        if return_outputs:1165            return (loss, metrics)1166        return loss1167 1168    def generate_from_model(self, model, batch: dict[str, torch.LongTensor]) -> str:1169        """Generate samples from the model and reference model for the given batch of inputs."""1170 1171        # If one uses `generate_during_eval` with peft + bf16, we need to explicitly call generate with1172        # the torch cuda amp context manager as some hidden states are silently casted to full precision.1173        generate_context_manager = amp.autocast("cuda") if self._peft_has_been_casted_to_bf16 else nullcontext()1174 1175        with generate_context_manager:1176            policy_output = model.generate(1177                input_ids=batch["prompt_input_ids"],1178                attention_mask=batch["prompt_attention_mask"],1179                max_length=self.max_length,1180                do_sample=True,1181                pad_token_id=self.processing_class.pad_token_id,1182            )1183 1184        policy_output = pad_to_length(policy_output, self.max_length, self.processing_class.pad_token_id)1185        policy_output_decoded = self.processing_class.batch_decode(policy_output, skip_special_tokens=True)1186 1187        return policy_output_decoded1188 1189    def prediction_step(1190        self,1191        model: Union[PreTrainedModel, nn.Module],1192        inputs: dict[str, Union[torch.Tensor, Any]],1193        prediction_loss_only: bool,1194        ignore_keys: Optional[list[str]] = None,1195    ):1196        if not self.use_dpo_data_collator:1197            warnings.warn(1198                "prediction_step is only implemented for DPODataCollatorWithPadding, and you passed a datacollator that is different than "1199                "DPODataCollatorWithPadding - you might see unexpected behavior. Alternatively, you can implement your own prediction_step method if you are using a custom data collator"1200            )

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