Team Ai
Apppublic

Zwounds/Boolean_Search_Query_Model

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
UnslothXPOTrainer.py1011 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.xpo_trainer import (Any, BaseImageProcessor, BasePairwiseJudge, Callable, Dataset, EvalPrediction, F, FeatureExtractionMixin, IterableDataset, OnlineDPOTrainer, OptimizerNames, Optional, PreTrainedModel, PreTrainedTokenizerBase, ProcessorMixin, SIMPLE_CHAT_TEMPLATE, TrainerCallback, Union, XPOConfig, XPOTrainer, empty_cache, generate_model_card, get_comet_experiment_url, get_reward, is_conversational, is_wandb_available, jinja2, maybe_apply_chat_template, nn, os, textwrap, torch, truncate_right, unwrap_model_for_generation)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 UnslothXPOConfig(XPOConfig):44    """45    46    Configuration class for the [`XPOTrainer`].47 48    Subclass of [`OnlineDPOConfig`] we can use all its arguments and add the following:49 50    Parameters:51        alpha (`float` or `list[float]`, *optional*, defaults to `1e-5`):52            Weight of the XPO loss term. If a list of floats is provided then the alpha is selected for each new epoch53            and the last alpha is used for the rest of the epochs.54    55    """56    vllm_sampling_params: Optional[Any] = field(57        default = None,58        metadata = {'help': 'vLLM SamplingParams'},59    )60    unsloth_num_chunks : Optional[int] = field(61        default = -1,62        metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},63    )64    def __init__(65        self,66        output_dir = None,67        overwrite_output_dir = None,68        do_train = False,69        do_eval = False,70        do_predict = False,71        eval_strategy = 'no',72        prediction_loss_only = False,73        per_device_train_batch_size = 4,74        per_device_eval_batch_size = 4,75        per_gpu_train_batch_size = None,76        per_gpu_eval_batch_size = None,77        gradient_accumulation_steps = 2,78        eval_accumulation_steps = 2,79        eval_delay = 0,80        torch_empty_cache_steps = 250,81        learning_rate = 5e-05,82        weight_decay = 0.01,83        adam_beta1 = 0.9,84        adam_beta2 = 0.999,85        adam_epsilon = 1e-08,86        max_grad_norm = 1.0,87        num_train_epochs = 3.0,88        max_steps = -1,89        lr_scheduler_type = 'linear',90        warmup_ratio = 0.1,91        warmup_steps = 0,92        log_level = 'passive',93        log_level_replica = 'warning',94        log_on_each_node = True,95        logging_dir = None,96        logging_strategy = 'steps',97        logging_first_step = False,98        logging_steps = 1,99        logging_nan_inf_filter = False,100        save_strategy = 'steps',101        save_steps = 500,102        save_total_limit = None,103        save_safetensors = True,104        save_on_each_node = False,105        save_only_model = False,106        restore_callback_states_from_checkpoint = False,107        no_cuda = False,108        use_cpu = False,109        use_mps_device = False,110        seed = 3407,111        data_seed = 3407,112        jit_mode_eval = False,113        use_ipex = False,114        bf16 = False,115        fp16 = False,116        fp16_opt_level = 'O1',117        half_precision_backend = 'auto',118        bf16_full_eval = False,119        fp16_full_eval = False,120        tf32 = None,121        local_rank = -1,122        ddp_backend = None,123        tpu_num_cores = None,124        tpu_metrics_debug = False,125        debug = '',126        dataloader_drop_last = False,127        eval_steps = None,128        dataloader_num_workers = 0,129        dataloader_prefetch_factor = None,130        past_index = -1,131        run_name = None,132        disable_tqdm = None,133        remove_unused_columns = True,134        label_names = None,135        load_best_model_at_end = False,136        metric_for_best_model = None,137        greater_is_better = None,138        ignore_data_skip = False,139        fsdp = '',140        fsdp_min_num_params = 0,141        fsdp_config = None,142        tp_size = 0,143        fsdp_transformer_layer_cls_to_wrap = None,144        accelerator_config = None,145        deepspeed = None,146        label_smoothing_factor = 0.0,147        optim = 'adamw_8bit',148        optim_args = None,149        adafactor = False,150        group_by_length = False,151        length_column_name = 'length',152        report_to = None,153        ddp_find_unused_parameters = None,154        ddp_bucket_cap_mb = None,155        ddp_broadcast_buffers = None,156        dataloader_pin_memory = True,157        dataloader_persistent_workers = False,158        skip_memory_metrics = True,159        use_legacy_prediction_loop = False,160        push_to_hub = False,161        resume_from_checkpoint = None,162        hub_model_id = None,163        hub_strategy = 'every_save',164        hub_token = None,165        hub_private_repo = None,166        hub_always_push = False,167        gradient_checkpointing = False,168        gradient_checkpointing_kwargs = None,169        include_inputs_for_metrics = False,170        eval_do_concat_batches = True,171        fp16_backend = 'auto',172        evaluation_strategy = None,173        push_to_hub_model_id = None,174        push_to_hub_organization = None,175        push_to_hub_token = None,176        mp_parameters = '',177        auto_find_batch_size = False,178        full_determinism = False,179        torchdynamo = None,180        ray_scope = 'last',181        ddp_timeout = 1800,182        torch_compile = False,183        torch_compile_backend = None,184        torch_compile_mode = None,185        dispatch_batches = None,186        split_batches = None,187        include_tokens_per_second = False,188        include_num_input_tokens_seen = False,189        neftune_noise_alpha = None,190        optim_target_modules = None,191        batch_eval_metrics = False,192        eval_on_start = False,193        use_liger_kernel = False,194        eval_use_gather_object = False,195        average_tokens_across_devices = False,196        reward_model_path = None,197        judge = None,198        max_new_tokens = 64,199        max_length = 512,200        temperature = 0.9,201        missing_eos_penalty = None,202        loss_type = 'sigmoid',203        dataset_num_proc = None,204        disable_dropout = True,205        use_vllm = False,206        ds3_gather_for_generation = True,207        vllm_sampling_params = None,208        unsloth_num_chunks = -1,209        **kwargs,210    ):211        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!')212        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!')213        if output_dir is None and save_strategy == 'steps' and save_steps == 500:214            output_dir = 'unsloth_training_checkpoints'215            save_strategy = 'no'216        if dataset_num_proc is None:217            from multiprocessing import cpu_count218            dataset_num_proc = cpu_count()219        220        super().__init__(221            output_dir = output_dir,222            overwrite_output_dir = overwrite_output_dir,223            do_train = do_train,224            do_eval = do_eval,225            do_predict = do_predict,226            eval_strategy = eval_strategy,227            prediction_loss_only = prediction_loss_only,228            per_device_train_batch_size = per_device_train_batch_size,229            per_device_eval_batch_size = per_device_eval_batch_size,230            per_gpu_train_batch_size = per_gpu_train_batch_size,231            per_gpu_eval_batch_size = per_gpu_eval_batch_size,232            gradient_accumulation_steps = gradient_accumulation_steps,233            eval_accumulation_steps = eval_accumulation_steps,234            eval_delay = eval_delay,235            torch_empty_cache_steps = torch_empty_cache_steps,236            learning_rate = learning_rate,237            weight_decay = weight_decay,238            adam_beta1 = adam_beta1,239            adam_beta2 = adam_beta2,240            adam_epsilon = adam_epsilon,241            max_grad_norm = max_grad_norm,242            num_train_epochs = num_train_epochs,243            max_steps = max_steps,244            lr_scheduler_type = lr_scheduler_type,245            warmup_ratio = warmup_ratio,246            warmup_steps = warmup_steps,247            log_level = log_level,248            log_level_replica = log_level_replica,249            log_on_each_node = log_on_each_node,250            logging_dir = logging_dir,251            logging_strategy = logging_strategy,252            logging_first_step = logging_first_step,253            logging_steps = logging_steps,254            logging_nan_inf_filter = logging_nan_inf_filter,255            save_strategy = save_strategy,256            save_steps = save_steps,257            save_total_limit = save_total_limit,258            save_safetensors = save_safetensors,259            save_on_each_node = save_on_each_node,260            save_only_model = save_only_model,261            restore_callback_states_from_checkpoint = restore_callback_states_from_checkpoint,262            no_cuda = no_cuda,263            use_cpu = use_cpu,264            use_mps_device = use_mps_device,265            seed = seed,266            data_seed = data_seed,267            jit_mode_eval = jit_mode_eval,268            use_ipex = use_ipex,269            bf16 = bf16,270            fp16 = fp16,271            fp16_opt_level = fp16_opt_level,272            half_precision_backend = half_precision_backend,273            bf16_full_eval = bf16_full_eval,274            fp16_full_eval = fp16_full_eval,275            tf32 = tf32,276            local_rank = local_rank,277            ddp_backend = ddp_backend,278            tpu_num_cores = tpu_num_cores,279            tpu_metrics_debug = tpu_metrics_debug,280            debug = debug,281            dataloader_drop_last = dataloader_drop_last,282            eval_steps = eval_steps,283            dataloader_num_workers = dataloader_num_workers,284            dataloader_prefetch_factor = dataloader_prefetch_factor,285            past_index = past_index,286            run_name = run_name,287            disable_tqdm = disable_tqdm,288            remove_unused_columns = remove_unused_columns,289            label_names = label_names,290            load_best_model_at_end = load_best_model_at_end,291            metric_for_best_model = metric_for_best_model,292            greater_is_better = greater_is_better,293            ignore_data_skip = ignore_data_skip,294            fsdp = fsdp,295            fsdp_min_num_params = fsdp_min_num_params,296            fsdp_config = fsdp_config,297            tp_size = tp_size,298            fsdp_transformer_layer_cls_to_wrap = fsdp_transformer_layer_cls_to_wrap,299            accelerator_config = accelerator_config,300            deepspeed = deepspeed,301            label_smoothing_factor = label_smoothing_factor,302            optim = optim,303            optim_args = optim_args,304            adafactor = adafactor,305            group_by_length = group_by_length,306            length_column_name = length_column_name,307            report_to = report_to,308            ddp_find_unused_parameters = ddp_find_unused_parameters,309            ddp_bucket_cap_mb = ddp_bucket_cap_mb,310            ddp_broadcast_buffers = ddp_broadcast_buffers,311            dataloader_pin_memory = dataloader_pin_memory,312            dataloader_persistent_workers = dataloader_persistent_workers,313            skip_memory_metrics = skip_memory_metrics,314            use_legacy_prediction_loop = use_legacy_prediction_loop,315            push_to_hub = push_to_hub,316            resume_from_checkpoint = resume_from_checkpoint,317            hub_model_id = hub_model_id,318            hub_strategy = hub_strategy,319            hub_token = hub_token,320            hub_private_repo = hub_private_repo,321            hub_always_push = hub_always_push,322            gradient_checkpointing = gradient_checkpointing,323            gradient_checkpointing_kwargs = gradient_checkpointing_kwargs,324            include_inputs_for_metrics = include_inputs_for_metrics,325            eval_do_concat_batches = eval_do_concat_batches,326            fp16_backend = fp16_backend,327            evaluation_strategy = evaluation_strategy,328            push_to_hub_model_id = push_to_hub_model_id,329            push_to_hub_organization = push_to_hub_organization,330            push_to_hub_token = push_to_hub_token,331            mp_parameters = mp_parameters,332            auto_find_batch_size = auto_find_batch_size,333            full_determinism = full_determinism,334            torchdynamo = torchdynamo,335            ray_scope = ray_scope,336            ddp_timeout = ddp_timeout,337            torch_compile = torch_compile,338            torch_compile_backend = torch_compile_backend,339            torch_compile_mode = torch_compile_mode,340            dispatch_batches = dispatch_batches,341            split_batches = split_batches,342            include_tokens_per_second = include_tokens_per_second,343            include_num_input_tokens_seen = include_num_input_tokens_seen,344            neftune_noise_alpha = neftune_noise_alpha,345            optim_target_modules = optim_target_modules,346            batch_eval_metrics = batch_eval_metrics,347            eval_on_start = eval_on_start,348            use_liger_kernel = use_liger_kernel,349            eval_use_gather_object = eval_use_gather_object,350            average_tokens_across_devices = average_tokens_across_devices,351            reward_model_path = reward_model_path,352            judge = judge,353            max_new_tokens = max_new_tokens,354            max_length = max_length,355            temperature = temperature,356            missing_eos_penalty = missing_eos_penalty,357            loss_type = loss_type,358            dataset_num_proc = dataset_num_proc,359            disable_dropout = disable_dropout,360            use_vllm = use_vllm,361            ds3_gather_for_generation = ds3_gather_for_generation,**kwargs)362        self.vllm_sampling_params = vllm_sampling_params363        self.unsloth_num_chunks = unsloth_num_chunks364pass365 366class _UnslothXPOTrainer(OnlineDPOTrainer):367    r""""""368 369    _tag_names = ["trl", "xpo"]370 371    def __init__(372        self,373        model: Union[PreTrainedModel, nn.Module] = None,374        ref_model: Union[PreTrainedModel, nn.Module] = None,375        reward_model: Optional[nn.Module] = None,376        judge: Optional[BasePairwiseJudge] = None,377        args: Optional[XPOConfig] = None,378        data_collator: Optional[Callable] = None,379        train_dataset: Optional[Union[Dataset, IterableDataset]] = None,380        eval_dataset: Optional[Union[Dataset, dict[str, Dataset]]] = None,381        processing_class: Optional[382            Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin]383        ] = None,384        peft_config: Optional[dict] = None,385        compute_metrics: Optional[Callable[[EvalPrediction], dict]] = None,386        callbacks: Optional[list[TrainerCallback]] = None,387        optimizers: tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),388        preprocess_logits_for_metrics: Optional[Callable[[torch.Tensor, torch.Tensor], torch.Tensor]] = None,389    ) -> None:390        super().__init__(391            model=model,392            ref_model=ref_model,393            judge=judge,394            reward_model=reward_model,395            args=args,396            data_collator=data_collator,397            train_dataset=train_dataset,398            eval_dataset=eval_dataset,399            processing_class=processing_class,400            reward_processing_class=processing_class,  # for now, XPOTrainer can't use any reward model401            peft_config=peft_config,402            compute_metrics=compute_metrics,403            callbacks=callbacks,404            optimizers=optimizers,405            preprocess_logits_for_metrics=preprocess_logits_for_metrics,406        )407 408        self._alpha = self.args.alpha409 410        # Overwrite the stats dictionary to include XPO specific statistics411        self.stats = {412            # Remove "non_score_reward", "rlhf_reward", "scores"413            # Add "loss/dpo", "loss/xpo"414            "loss/dpo": [],415            "loss/xpo": [],416            "objective/kl": [],417            "objective/entropy": [],418            "rewards/chosen": [],419            "rewards/rejected": [],420            "rewards/accuracies": [],421            "rewards/margins": [],422            "logps/chosen": [],423            "logps/rejected": [],424            # Replace "contain_eos_token" by "model_contain_eos_token" and "ref_contain_eos_token"425            "val/model_contain_eos_token": [],426            "val/ref_contain_eos_token": [],427            "alpha": [],428            "beta": [],429        }430        if self.reward_model is not None:431            # Replace "scores" by "model_scores" and "ref_scores"432            self.stats["objective/model_scores"] = []433            self.stats["objective/ref_scores"] = []434            self.stats["objective/scores_margin"] = []435 436    @property437    def alpha(self):438        if isinstance(self._alpha, list):439            epoch = self.state.epoch440            return self._alpha[epoch] if epoch < len(self._alpha) else self._alpha[-1]441        else:442            return self._alpha443 444    def _generate_completions(self, prompts, model):445        with unwrap_model_for_generation(model, self.accelerator) as unwrapped_model:446            model_output = unwrapped_model.generate(447                input_ids=prompts["input_ids"],448                attention_mask=prompts["attention_mask"],449                generation_config=self.generation_config,450            )451 452        ref_model = model if self.ref_model is None else self.ref_model453        with torch.no_grad(), unwrap_model_for_generation(ref_model, self.accelerator) as unwrapped_ref_model:454            ref_output = unwrapped_ref_model.generate(455                input_ids=prompts["input_ids"],456                attention_mask=prompts["attention_mask"],457                generation_config=self.generation_config,458            )459 460        return model_output, ref_output461 462    def _process_completions(self, model_output, ref_output, prompts):463        context_length = prompts["input_ids"].shape[1]464 465        # Process model completions466        model_completion_ids = model_output[:, context_length:]467        model_completion_ids, model_completion_mask = truncate_right(468            model_completion_ids, self.processing_class.eos_token_id, self.processing_class.pad_token_id469        )470        model_data = {471            "input_ids": torch.cat((prompts["input_ids"], model_completion_ids), dim=1),472            "attention_mask": torch.cat((prompts["attention_mask"], model_completion_mask), dim=1),473            "raw": prompts["raw"],474        }475 476        # Process reference model completions477        ref_completion_ids = ref_output[:, context_length:]478        ref_completion_ids, ref_completion_mask = truncate_right(479            ref_completion_ids, self.processing_class.eos_token_id, self.processing_class.pad_token_id480        )481        ref_data = {482            "input_ids": torch.cat((prompts["input_ids"], ref_completion_ids), dim=1),483            "attention_mask": torch.cat((prompts["attention_mask"], ref_completion_mask), dim=1),484            "raw": prompts["raw"],485        }486 487        return model_data, ref_data488 489    def _compute_rewards(self, model_data, ref_data, context_length):490        with torch.no_grad():491            _, model_scores, _ = get_reward(492                self.reward_model, model_data["input_ids"], self.processing_class.pad_token_id, context_length493            )494            _, ref_scores, _ = get_reward(495                self.reward_model, ref_data["input_ids"], self.processing_class.pad_token_id, context_length496            )497 498        # Apply EOS penalty if needed499        if self.args.missing_eos_penalty is not None:500            model_contain_eos = torch.any(model_data["input_ids"] == self.processing_class.eos_token_id, dim=-1)501            ref_contain_eos = torch.any(ref_data["input_ids"] == self.processing_class.eos_token_id, dim=-1)502            model_scores[~model_contain_eos] -= self.args.missing_eos_penalty503            ref_scores[~ref_contain_eos] -= self.args.missing_eos_penalty504 505        return model_scores, ref_scores506 507    def _compute_judge(self, model_data, ref_data, context_length):508        prompts = model_data["raw"]509        model_data_completions = self.processing_class.batch_decode(510            model_data["input_ids"][:, context_length:], skip_special_tokens=True511        )512        model_data_completions = [completion.strip() for completion in model_data_completions]513 514        ref_data_completions = self.processing_class.batch_decode(515            ref_data["input_ids"][:, context_length:], skip_special_tokens=True516        )517        ref_data_completions = [completion.strip() for completion in ref_data_completions]518 519        if is_conversational({"prompt": prompts[0]}):520            model_data_completions = [521                [{"role": "assistant", "content": completion}] for completion in model_data_completions522            ]523            environment = jinja2.Environment()524            template = environment.from_string(SIMPLE_CHAT_TEMPLATE)525            prompts = [template.render(messages=message) for message in prompts]526            model_data_completions = [template.render(messages=completion) for completion in model_data_completions]527 528            ref_data_completions = [529                [{"role": "assistant", "content": completion}] for completion in ref_data_completions530            ]531            ref_data_completions = [template.render(messages=completion) for completion in ref_data_completions]532 533        ranks_of_first_completion = self.judge.judge(534            prompts,535            list(zip(model_data_completions, ref_data_completions)),536        )537        # convert ranks to a True/False mask:538        # when rank == 0, it means the first completion is the best539        # when rank == 1, it means the second completion is the best540        return torch.tensor([rank == 0 for rank in ranks_of_first_completion], device=model_data["input_ids"].device)541 542    def _compute_logprobs(self, model, model_data, ref_data, context_length):543        def compute_logprobs_for_data(m, data):544            output = m(data["input_ids"], attention_mask=data["attention_mask"])545            logits = output.logits[:, context_length - 1 : -1]546            token_logprobs = selective_log_softmax(logits, data["input_ids"][:, context_length:])547            return token_logprobs548 549        # Compute logprobs for model completions550        model_logprobs_model_data = compute_logprobs_for_data(model, model_data)551        # Compute logprobs for model on reference completions (for XPO loss)552        model_logprobs_ref_data = compute_logprobs_for_data(model, ref_data)553 554        # Compute logprobs for reference model completions555        with torch.no_grad():556            if self.ref_model is None:557                with model.disable_adapter():558                    ref_logprobs_model_data = compute_logprobs_for_data(model, model_data)559                    ref_logprobs_ref_data = compute_logprobs_for_data(model, ref_data)560            else:561                ref_logprobs_model_data = compute_logprobs_for_data(self.ref_model, model_data)562                ref_logprobs_ref_data = compute_logprobs_for_data(self.ref_model, ref_data)563 564        # Mask padding tokens565        model_padding_mask = model_data["attention_mask"][:, context_length:] == 0566        ref_padding_mask = ref_data["attention_mask"][:, context_length:] == 0567        model_logprobs_model_data = model_logprobs_model_data.masked_fill(model_padding_mask, 0.0)568        model_logprobs_ref_data = model_logprobs_ref_data.masked_fill(ref_padding_mask, 0.0)569        ref_logprobs_ref_data = ref_logprobs_ref_data.masked_fill(ref_padding_mask, 0.0)570        ref_logprobs_model_data = ref_logprobs_model_data.masked_fill(model_padding_mask, 0.0)571 572        return model_logprobs_model_data, model_logprobs_ref_data, ref_logprobs_ref_data, ref_logprobs_model_data573 574    def _compute_losses(575        self,576        model_logprobs_model_data,577        model_logprobs_ref_data,578        ref_logprobs_ref_data,579        ref_logprobs_model_data,580        chosen_mask,581    ):582        # Compute log probs583        model_logprobs_model_data_sum = model_logprobs_model_data.sum(1)584        model_logprobs_ref_data_sum = model_logprobs_ref_data.sum(1)585        ref_logprobs_ref_data_sum = ref_logprobs_ref_data.sum(1)586        ref_logprobs_model_data_sum = ref_logprobs_model_data.sum(1)587 588        chosen_model_logprobs = torch.where(chosen_mask, model_logprobs_model_data_sum, model_logprobs_ref_data_sum)589        chosen_ref_logprobs = torch.where(chosen_mask, ref_logprobs_model_data_sum, ref_logprobs_ref_data_sum)590        chosen_log_ratios = chosen_model_logprobs - chosen_ref_logprobs591 592        rejected_model_logprobs = torch.where(~chosen_mask, model_logprobs_model_data_sum, model_logprobs_ref_data_sum)593        rejected_ref_logprobs = torch.where(~chosen_mask, ref_logprobs_model_data_sum, ref_logprobs_ref_data_sum)594        rejected_log_ratios = rejected_model_logprobs - rejected_ref_logprobs595 596        # Compute logits as the difference between chosen and rejected log ratios597        logits = chosen_log_ratios - rejected_log_ratios598 599        if self.args.loss_type == "sigmoid":600            dpo_losses = -F.logsigmoid(self.beta * logits)601        elif self.args.loss_type == "ipo":602            dpo_losses = (logits - 1 / (2 * self.beta)) ** 2603        else:604            raise NotImplementedError(f"invalid loss type {self.args.loss_type}")605 606        # Compute XPO specific loss607        xpo_losses = self.alpha * model_logprobs_ref_data_sum608 609        # Total loss610        loss = (dpo_losses + xpo_losses).mean()611 612        return loss, dpo_losses, xpo_losses613 614    def _log_statistics(615        self,616        model_data,617        ref_data,618        model_logprobs_model_data,619        model_logprobs_ref_data,620        ref_logprobs_ref_data,621        ref_logprobs_model_data,622        chosen_mask,623        dpo_losses,624        xpo_losses,625        context_length,626        model_scores=None,627        ref_scores=None,628    ):629        # Helper function to gather and compute mean630        def gather_mean(tensor):631            return self.accelerator.gather_for_metrics(tensor).mean().item()632 633        # Log losses634        self.stats["loss/dpo"].append(gather_mean(dpo_losses))635        self.stats["loss/xpo"].append(gather_mean(xpo_losses))636 637        # Log scores638        if self.reward_model is not None:639            self.stats["objective/model_scores"].append(gather_mean(model_scores))640            self.stats["objective/ref_scores"].append(gather_mean(ref_scores))641            self.stats["objective/scores_margin"].append(gather_mean(model_scores - ref_scores))642 643        # Log logprobs644        model_logprobs_model_data_sum = model_logprobs_model_data.sum(1)645        model_logprobs_ref_data_sum = model_logprobs_ref_data.sum(1)646        ref_logprobs_ref_data_sum = ref_logprobs_ref_data.sum(1)647        ref_logprobs_model_data_sum = ref_logprobs_model_data.sum(1)648 649        chosen_model_logprobs = torch.where(chosen_mask, model_logprobs_model_data_sum, model_logprobs_ref_data_sum)650        chosen_ref_logprobs = torch.where(chosen_mask, ref_logprobs_model_data_sum, ref_logprobs_ref_data_sum)651        chosen_log_ratios = chosen_model_logprobs - chosen_ref_logprobs652 653        rejected_model_logprobs = torch.where(~chosen_mask, model_logprobs_model_data_sum, model_logprobs_ref_data_sum)654        rejected_ref_logprobs = torch.where(~chosen_mask, ref_logprobs_model_data_sum, ref_logprobs_ref_data_sum)655        rejected_log_ratios = rejected_model_logprobs - rejected_ref_logprobs656 657        self.stats["logps/chosen"].append(gather_mean(chosen_model_logprobs.mean() + chosen_ref_logprobs.mean()))658        self.stats["logps/rejected"].append(gather_mean(rejected_model_logprobs.mean() + rejected_ref_logprobs.mean()))659 660        # Log rewards661        # Compute various statistics662        chosen_rewards = chosen_log_ratios * self.beta663        rejected_rewards = rejected_log_ratios * self.beta664        self.stats["rewards/chosen"].append(gather_mean(chosen_rewards.mean()))665        self.stats["rewards/rejected"].append(gather_mean(rejected_rewards.mean()))666 667        # Calculate KL divergence for model and ref data668        kl_model_data = model_logprobs_model_data - ref_logprobs_model_data669        kl_ref_data = model_logprobs_ref_data - ref_logprobs_ref_data670        mean_kl = (kl_model_data.sum(1) + kl_ref_data.sum(1)).mean() / 2671        self.stats["objective/kl"].append(gather_mean(mean_kl))672 673        # Calculate entropy for model and ref data674        entropy_model_data = -model_logprobs_model_data.sum(1)675        entropy_ref_data = -model_logprobs_ref_data.sum(1)676        mean_entropy = (entropy_model_data.mean() + entropy_ref_data.mean()) / 2677        self.stats["objective/entropy"].append(gather_mean(mean_entropy))678 679        # Calculate margins680        margin = chosen_rewards - rejected_rewards681        self.stats["rewards/margins"].append(gather_mean(margin.mean()))682 683        # Calculate accuracy684        accuracy = (margin > 0).float()685        self.stats["rewards/accuracies"].append(gather_mean(accuracy.mean()))686 687        # Log EOS token statistics688        model_eos = (model_data["input_ids"][:, context_length:] == self.processing_class.eos_token_id).any(dim=1)689        ref_eos = (ref_data["input_ids"][:, context_length:] == self.processing_class.eos_token_id).any(dim=1)690        self.stats["val/model_contain_eos_token"].append(gather_mean(model_eos.float()))691        self.stats["val/ref_contain_eos_token"].append(gather_mean(ref_eos.float()))692 693        # Log alpha and beta694        self.stats["alpha"].append(self.alpha)695        self.stats["beta"].append(self.beta)696 697    def training_step(698        self, model: nn.Module, inputs: dict[str, Union[torch.Tensor, Any]], num_items_in_batch: Optional[int] = None699    ) -> torch.Tensor:700        model.train()701 702        # Apply chat template and tokenize the input703        batch_size = len(next(iter(inputs.values())))704        prompts = inputs["prompt"]705        inputs = [{k: v[i] for k, v in inputs.items()} for i in range(batch_size)]706        inputs = [maybe_apply_chat_template(x, self.processing_class) for x in inputs]707        inputs = [self.tokenize_row(x, self.model.config.is_encoder_decoder, self.processing_class) for x in inputs]708        inputs = self.data_collator(inputs)709 710        # need the prompt_ only711        inputs = self._prepare_inputs(inputs)712        context_length = inputs["prompt_input_ids"].shape[1]713        prompts = {714            "input_ids": inputs["prompt_input_ids"],715            "attention_mask": inputs["prompt_attention_mask"],716            "raw": prompts,717        }718        del inputs719 720        # Sample completions from both the model and the reference model721        model_output, ref_output = self._generate_completions(prompts, model)722 723        # Process model completions724        model_data, ref_data = self._process_completions(model_output, ref_output, prompts)725 726        # Compute rewards727        if self.reward_model is not None:728            model_scores, ref_scores = self._compute_rewards(model_data, ref_data, context_length)729            chosen_mask = model_scores >= ref_scores730        else:731            model_scores, ref_scores = None, None732            chosen_mask = self._compute_judge(model_data, ref_data, context_length)733 734        # Compute logprobs735        model_logprobs_model_data, model_logprobs_ref_data, ref_logprobs_ref_data, ref_logprobs_model_data = (736            self._compute_logprobs(model, model_data, ref_data, context_length)737        )738 739        # Compute loss740        loss, dpo_losses, xpo_losses = self._compute_losses(741            model_logprobs_model_data,742            model_logprobs_ref_data,743            ref_logprobs_ref_data,744            ref_logprobs_model_data,745            chosen_mask,746        )747 748        # Log everything749        self._log_statistics(750            model_data,751            ref_data,752            model_logprobs_model_data.detach(),753            model_logprobs_ref_data.detach(),754            ref_logprobs_ref_data,755            ref_logprobs_model_data,756            chosen_mask,757            dpo_losses.detach(),758            xpo_losses.detach(),759            context_length,760            model_scores,761            ref_scores,762        )763 764        if (765            self.args.torch_empty_cache_steps is not None766            and self.state.global_step % self.args.torch_empty_cache_steps == 0767        ):768            empty_cache()769 770        kwargs = {}771        # For LOMO optimizers you need to explicitly use the learning rate772        if self.args.optim in [OptimizerNames.LOMO, OptimizerNames.ADALOMO]:773            kwargs["learning_rate"] = self._get_learning_rate()774 775        if self.args.n_gpu > 1:776            loss = loss.mean()  # mean() to average on multi-gpu parallel training777 778        if self.use_apex:779            with amp.scale_loss(loss, self.optimizer) as scaled_loss:780                scaled_loss.backward()781        else:782            self.accelerator.backward(loss, **kwargs)783 784        return loss.detach() / self.args.gradient_accumulation_steps785 786    def create_model_card(787        self,788        model_name: Optional[str] = None,789        dataset_name: Optional[str] = None,790        tags: Union[str, list[str], None] = None,791    ):792        """793        Creates a draft of a model card using the information available to the `Trainer`.794 795        Args:796            model_name (`str` or `None`, *optional*, defaults to `None`):797                Name of the model.798            dataset_name (`str` or `None`, *optional*, defaults to `None`):799                Name of the dataset used for training.800            tags (`str`, `list[str]` or `None`, *optional*, defaults to `None`):801                Tags to be associated with the model card.802        """803        if not self.is_world_process_zero():804            return805 806        if hasattr(self.model.config, "_name_or_path") and not os.path.isdir(self.model.config._name_or_path):807            base_model = self.model.config._name_or_path808        else:809            base_model = None810 811        tags = tags or []812        if isinstance(tags, str):813            tags = [tags]814 815        if hasattr(self.model.config, "unsloth_version"):816            tags.append("unsloth")817 818        citation = textwrap.dedent("""\819        @article{jung2024binary,820            title        = {{Exploratory Preference Optimization: Harnessing Implicit Q*-Approximation for Sample-Efficient RLHF}},821            author       = {Tengyang Xie and Dylan J. Foster and Akshay Krishnamurthy and Corby Rosset and Ahmed Awadallah and Alexander Rakhlin},822            year         = 2024,823            eprint       = {arXiv:2405.21046}824        }""")825 826        model_card = generate_model_card(827            base_model=base_model,828            model_name=model_name,829            hub_model_id=self.hub_model_id,830            dataset_name=dataset_name,831            tags=tags,832            wandb_url=wandb.run.get_url() if is_wandb_available() and wandb.run is not None else None,833            comet_url=get_comet_experiment_url(),834            trainer_name="XPO",835            trainer_citation=citation,836            paper_title="Exploratory Preference Optimization: Harnessing Implicit Q*-Approximation for Sample-Efficient RLHF",837            paper_id="2405.21046",838        )839 840        model_card.save(os.path.join(self.args.output_dir, "README.md"))841class UnslothXPOTrainer(_UnslothXPOTrainer):842    """843    844    Initialize XPOTrainer as a subclass of [`OnlineDPOConfig`].845 846    Args:847        model (`transformers.PreTrainedModel`):848            The model to train, preferably an `AutoModelForCausalLM`.849        ref_model (`PreTrainedModelWrapper`):850            Hugging Face transformer model with a casual language modelling head. Used for implicit reward computation and loss. If no851            reference model is provided, the trainer will create a reference model with the same architecture as the model to be optimized.852        reward_model (`transformers.PreTrainedModel`):853            The reward model to score completions with, preferably an `AutoModelForSequenceClassification`.854        judge (`BasePairwiseJudge`):855            The judge to use for pairwise comparison of model completions.856        args (`XPOConfig`):857            The XPO config arguments to use for training.858        data_collator (`transformers.DataCollator`):859            The data collator to use for training. If None is specified, the default data collator (`DPODataCollatorWithPadding`) will be used860            which will pad the sequences to the maximum length of the sequences in the batch, given a dataset of paired sequences.861        train_dataset (`datasets.Dataset`):862            The dataset to use for training.863        eval_dataset (`datasets.Dataset`):864            The dataset to use for evaluation.865        processing_class (`PreTrainedTokenizerBase` or `BaseImageProcessor` or `FeatureExtractionMixin` or `ProcessorMixin`, *optional*):866            Processing class used to process the data. If provided, will be used to automatically process the inputs867            for the model, and it will be saved along the model to make it easier to rerun an interrupted training or868            reuse the fine-tuned model.869        peft_config (`dict`):870            The peft config to use for training.871        compute_metrics (`Callable[[EvalPrediction], dict]`, *optional*):872            The function to use to compute the metrics. Must take a `EvalPrediction` and return873            a dictionary string to metric values.874        callbacks (`list[transformers.TrainerCallback]`):875            The callbacks to use for training.876        optimizers (`tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR]`):877            The optimizer and scheduler to use for training.878        preprocess_logits_for_metrics (`Callable[[torch.Tensor, torch.Tensor], torch.Tensor]`):879            The function to use to preprocess the logits before computing the metrics.880    881    """882    def __init__(883        self,884        model = None,885        ref_model = None,886        reward_model = None,887        judge = None,888        args = None,889        data_collator = None,890        train_dataset = None,891        eval_dataset = None,892        processing_class = None,893        peft_config = None,894        compute_metrics = None,895        callbacks = None,896        preprocess_logits_for_metrics = None,897        **kwargs898    ):899        if args is None: args = UnslothXPOConfig()900        use_bf16 = getattr(args, 'bf16', False)901        use_fp16 = getattr(args, 'fp16', False)902        force_float32 = False903        if os.environ.get('UNSLOTH_FORCE_FLOAT32', '0') == '1':904            print('Unsloth: Switching to float32 training since model cannot work with float16')905            force_float32 = True906        mixed_precision_dtype = os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32')907        dtype = getattr(model.config, 'torch_dtype', None)908        if dtype is None: dtype = model.get_input_embeddings().dtype909        from unsloth_zoo.utils import _get_dtype910        dtype = _get_dtype(dtype)911        float16 = dtype == torch.float16912        if not force_float32 and (float16 and use_bf16): raise TypeError('Unsloth: Model is in float16 precision but you want to use bfloat16 precision. Set fp16 to `True` and bf16 to `False`')913        if not force_float32 and (not float16 and use_fp16): raise TypeError('Unsloth: Model is in bfloat16 precision but you want to use float16 precision. Set fp16 to `False` and bf16 to `True`')914        if force_float32:915            args.fp16 = False916            args.bf16 = False917            os.environ['ACCELERATE_MIXED_PRECISION'] = 'no'918        elif (not use_bf16 and not use_fp16) and mixed_precision_dtype == 'float32':919            args.fp16 = float16920            args.bf16 = not float16921            os.environ['ACCELERATE_MIXED_PRECISION'] = 'fp16' if float16 else 'bf16'922        if getattr(args, 'eval_dataset', None) is not None and getattr(args, 'eval_strategy', 'no') == 'no':923            args.eval_strategy = 'steps'924            if getattr(args, 'eval_steps', None) is None: args.eval_steps = 0.1925        ga_steps = getattr(args, 'gradient_accumulation_steps', None)926        if ga_steps is not None and ga_steps > 1:927            from transformers import __version__ as transformers_version928            if Version(transformers_version) <= Version('4.45.2'):929                print('**** Unsloth: Please use our fixed gradient_accumulation_steps by updating transformers, TRL and Unsloth!\n'930                      '`pip install --upgrade --no-cache-dir --force-reinstall --no-deps unsloth transformers trl unsloth_zoo`')931        if getattr(args, 'eval_strategy', 'no') != 'no':932            eval_bsz = getattr(args, 'per_device_eval_batch_size', 8)933            if eval_bsz == 8 and args.per_device_train_batch_size < eval_bsz: args.per_device_eval_batch_size = args.per_device_train_batch_size934            if getattr(args, 'eval_accumulation_steps', None) is None and ga_steps is not None: args.eval_accumulation_steps = ga_steps935        fp16_full_eval = getattr(args, 'fp16_full_eval', False)936        bf16_full_eval = getattr(args, 'bf16_full_eval', False)937        if args.fp16 and bf16_full_eval: args.bf16_full_eval = False; args.fp16_full_eval = True938        if args.bf16 and fp16_full_eval: args.bf16_full_eval = True; args.fp16_full_eval = False939        if force_float32:940            args.bf16_full_eval = False941            args.fp16_full_eval = False942        elif os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32') == 'bfloat16':943            args.bf16_full_eval = True944            args.fp16_full_eval = False945        elif not bf16_full_eval and not fp16_full_eval:946            args.bf16_full_eval = args.bf16947            args.fp16_full_eval = args.fp16948        _output_logits = False949        if locals().get('compute_metrics', None) is not None: _output_logits = True950        if locals().get('preprocess_logits_for_metrics', None) is not None: _output_logits = True951        if _output_logits:952            os.environ['UNSLOTH_RETURN_LOGITS'] = '1'953        if 'max_seq_length' not in locals() and not hasattr(args, 'max_seq_length'):954            pass955        else:956            model_max_seq_length = getattr(model, 'max_seq_length', None)957            args_max_seq_length  = getattr(args,  'max_seq_length', None)958            if args_max_seq_length is None and model_max_seq_length is not None:959                max_seq_length = model.max_seq_length960                if hasattr(args, 'max_seq_length'): args.max_seq_length = max_seq_length961        if model is not None and hasattr(model, 'for_training'):962            model.for_training()963        if 'tokenizer' in locals() and hasattr(tokenizer, 'padding_side'): tokenizer.padding_side = 'right'964        if 'processing_class' in locals():965            if hasattr(processing_class, 'padding_side'): processing_class.padding_side = 'right'966            if hasattr(processing_class, 'tokenizer') and hasattr(processing_class.tokenizer, 'padding_side'): processing_class.tokenizer.padding_side = 'right'967        __tokenizer = processing_class if 'processing_class' in locals() else tokenizer968        from unsloth_zoo.vision_utils import UnslothVisionDataCollator969        if not isinstance(data_collator, UnslothVisionDataCollator):970            if isinstance(data_collator, DataCollatorForSeq2Seq) and 'labels' not in train_dataset.column_names:971                data_collator = DataCollatorForLanguageModeling(__tokenizer, mlm = False)972            elif isinstance(data_collator, DataCollatorForLanguageModeling) and 'labels' in train_dataset.column_names:973                data_collator = DataCollatorForSeq2Seq(__tokenizer)974        else:975            if hasattr(args, 'remove_unused_columns'): args.remove_unused_columns = False976            if hasattr(args, 'dataset_text_field'): args.dataset_text_field = ''977            if hasattr(args, 'dataset_kwargs'): args.dataset_kwargs = {'skip_prepare_dataset': True}978        if not isinstance(data_collator, UnslothVisionDataCollator):979            if not hasattr(__tokenizer, 'pad') and hasattr(__tokenizer, 'tokenizer'):980                if isinstance(data_collator, DataCollatorForSeq2Seq):981                    data_collator = DataCollatorForSeq2Seq(__tokenizer.tokenizer)982                else:983                    data_collator = DataCollatorForLanguageModeling(__tokenizer.tokenizer, mlm = False)984        other_metrics = []985        986        from unsloth_zoo.logging_utils import PatchRLStatistics987        PatchRLStatistics('xpo_trainer', other_metrics)988        989        super().__init__(990            model = model,991            ref_model = ref_model,992            reward_model = reward_model,993            judge = judge,994            args = args,995            data_collator = data_collator,996            train_dataset = train_dataset,997            eval_dataset = eval_dataset,998            processing_class = processing_class,999            peft_config = peft_config,1000            compute_metrics = compute_metrics,1001            callbacks = callbacks,1002            preprocess_logits_for_metrics = preprocess_logits_for_metrics,**kwargs)1003        if hasattr(self, 'neftune_hook_handle'):1004            self.neftune_hook_handle.remove()1005            if hasattr(self, 'neftune_hook_handle'): del self.neftune_hook_handle1006        if getattr(args, 'neftune_noise_alpha', None) is not None:1007            model.get_input_embeddings().neftune_noise_alpha = self.neftune_noise_alpha1008        pass1009        1010pass1011