Team Ai
Apppublic

Zwounds/Boolean_Search_Query_Model

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

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