Team Ai
Apppublic

Zwounds/Boolean_Search_Query_Model

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
UnslothNashMDTrainer.py956 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.nash_md_trainer import (Any, BaseImageProcessor, BasePairwiseJudge, Callable, Dataset, EvalPrediction, F, FeatureExtractionMixin, GeometricMixtureWrapper, IterableDataset, NashMDConfig, NashMDTrainer, OnlineDPOTrainer, OptimizerNames, Optional, PreTrainedModel, PreTrainedTokenizerBase, ProcessorMixin, SIMPLE_CHAT_TEMPLATE, TrainerCallback, Union, 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 UnslothNashMDConfig(NashMDConfig):44    """45    46    Configuration class for the [`NashMDTrainer`].47 48    Subclass of [`OnlineDPOConfig`] we can use all its arguments and add the following:49 50    Parameters:51        mixture_coef (`float` or `list[float]`, *optional*, defaults to `0.5`):52            Logit mixture coefficient for the model and reference model. If a list of floats is provided then the53            mixture coefficient is selected for each new epoch and the last coefficient is used for the rest of the54            epochs.55    56    """57    vllm_sampling_params: Optional[Any] = field(58        default = None,59        metadata = {'help': 'vLLM SamplingParams'},60    )61    unsloth_num_chunks : Optional[int] = field(62        default = -1,63        metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},64    )65    def __init__(66        self,67        output_dir = None,68        overwrite_output_dir = None,69        do_train = False,70        do_eval = False,71        do_predict = False,72        eval_strategy = 'no',73        prediction_loss_only = False,74        per_device_train_batch_size = 4,75        per_device_eval_batch_size = 4,76        per_gpu_train_batch_size = None,77        per_gpu_eval_batch_size = None,78        gradient_accumulation_steps = 2,79        eval_accumulation_steps = 2,80        eval_delay = 0,81        torch_empty_cache_steps = 250,82        learning_rate = 5e-05,83        weight_decay = 0.01,84        adam_beta1 = 0.9,85        adam_beta2 = 0.999,86        adam_epsilon = 1e-08,87        max_grad_norm = 1.0,88        num_train_epochs = 3.0,89        max_steps = -1,90        lr_scheduler_type = 'linear',91        warmup_ratio = 0.1,92        warmup_steps = 0,93        log_level = 'passive',94        log_level_replica = 'warning',95        log_on_each_node = True,96        logging_dir = None,97        logging_strategy = 'steps',98        logging_first_step = False,99        logging_steps = 1,100        logging_nan_inf_filter = False,101        save_strategy = 'steps',102        save_steps = 500,103        save_total_limit = None,104        save_safetensors = True,105        save_on_each_node = False,106        save_only_model = False,107        restore_callback_states_from_checkpoint = False,108        no_cuda = False,109        use_cpu = False,110        use_mps_device = False,111        seed = 3407,112        data_seed = 3407,113        jit_mode_eval = False,114        use_ipex = False,115        bf16 = False,116        fp16 = False,117        fp16_opt_level = 'O1',118        half_precision_backend = 'auto',119        bf16_full_eval = False,120        fp16_full_eval = False,121        tf32 = None,122        local_rank = -1,123        ddp_backend = None,124        tpu_num_cores = None,125        tpu_metrics_debug = False,126        debug = '',127        dataloader_drop_last = False,128        eval_steps = None,129        dataloader_num_workers = 0,130        dataloader_prefetch_factor = None,131        past_index = -1,132        run_name = None,133        disable_tqdm = None,134        remove_unused_columns = True,135        label_names = None,136        load_best_model_at_end = False,137        metric_for_best_model = None,138        greater_is_better = None,139        ignore_data_skip = False,140        fsdp = '',141        fsdp_min_num_params = 0,142        fsdp_config = None,143        tp_size = 0,144        fsdp_transformer_layer_cls_to_wrap = None,145        accelerator_config = None,146        deepspeed = None,147        label_smoothing_factor = 0.0,148        optim = 'adamw_8bit',149        optim_args = None,150        adafactor = False,151        group_by_length = False,152        length_column_name = 'length',153        report_to = None,154        ddp_find_unused_parameters = None,155        ddp_bucket_cap_mb = None,156        ddp_broadcast_buffers = None,157        dataloader_pin_memory = True,158        dataloader_persistent_workers = False,159        skip_memory_metrics = True,160        use_legacy_prediction_loop = False,161        push_to_hub = False,162        resume_from_checkpoint = None,163        hub_model_id = None,164        hub_strategy = 'every_save',165        hub_token = None,166        hub_private_repo = None,167        hub_always_push = False,168        gradient_checkpointing = False,169        gradient_checkpointing_kwargs = None,170        include_inputs_for_metrics = False,171        eval_do_concat_batches = True,172        fp16_backend = 'auto',173        evaluation_strategy = None,174        push_to_hub_model_id = None,175        push_to_hub_organization = None,176        push_to_hub_token = None,177        mp_parameters = '',178        auto_find_batch_size = False,179        full_determinism = False,180        torchdynamo = None,181        ray_scope = 'last',182        ddp_timeout = 1800,183        torch_compile = False,184        torch_compile_backend = None,185        torch_compile_mode = None,186        dispatch_batches = None,187        split_batches = None,188        include_tokens_per_second = False,189        include_num_input_tokens_seen = False,190        neftune_noise_alpha = None,191        optim_target_modules = None,192        batch_eval_metrics = False,193        eval_on_start = False,194        use_liger_kernel = False,195        eval_use_gather_object = False,196        average_tokens_across_devices = False,197        reward_model_path = None,198        judge = None,199        max_new_tokens = 64,200        max_length = 512,201        temperature = 0.9,202        missing_eos_penalty = None,203        loss_type = 'sigmoid',204        dataset_num_proc = None,205        disable_dropout = True,206        use_vllm = False,207        ds3_gather_for_generation = True,208        vllm_sampling_params = None,209        unsloth_num_chunks = -1,210        **kwargs,211    ):212        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!')213        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!')214        if output_dir is None and save_strategy == 'steps' and save_steps == 500:215            output_dir = 'unsloth_training_checkpoints'216            save_strategy = 'no'217        if dataset_num_proc is None:218            from multiprocessing import cpu_count219            dataset_num_proc = cpu_count()220        221        super().__init__(222            output_dir = output_dir,223            overwrite_output_dir = overwrite_output_dir,224            do_train = do_train,225            do_eval = do_eval,226            do_predict = do_predict,227            eval_strategy = eval_strategy,228            prediction_loss_only = prediction_loss_only,229            per_device_train_batch_size = per_device_train_batch_size,230            per_device_eval_batch_size = per_device_eval_batch_size,231            per_gpu_train_batch_size = per_gpu_train_batch_size,232            per_gpu_eval_batch_size = per_gpu_eval_batch_size,233            gradient_accumulation_steps = gradient_accumulation_steps,234            eval_accumulation_steps = eval_accumulation_steps,235            eval_delay = eval_delay,236            torch_empty_cache_steps = torch_empty_cache_steps,237            learning_rate = learning_rate,238            weight_decay = weight_decay,239            adam_beta1 = adam_beta1,240            adam_beta2 = adam_beta2,241            adam_epsilon = adam_epsilon,242            max_grad_norm = max_grad_norm,243            num_train_epochs = num_train_epochs,244            max_steps = max_steps,245            lr_scheduler_type = lr_scheduler_type,246            warmup_ratio = warmup_ratio,247            warmup_steps = warmup_steps,248            log_level = log_level,249            log_level_replica = log_level_replica,250            log_on_each_node = log_on_each_node,251            logging_dir = logging_dir,252            logging_strategy = logging_strategy,253            logging_first_step = logging_first_step,254            logging_steps = logging_steps,255            logging_nan_inf_filter = logging_nan_inf_filter,256            save_strategy = save_strategy,257            save_steps = save_steps,258            save_total_limit = save_total_limit,259            save_safetensors = save_safetensors,260            save_on_each_node = save_on_each_node,261            save_only_model = save_only_model,262            restore_callback_states_from_checkpoint = restore_callback_states_from_checkpoint,263            no_cuda = no_cuda,264            use_cpu = use_cpu,265            use_mps_device = use_mps_device,266            seed = seed,267            data_seed = data_seed,268            jit_mode_eval = jit_mode_eval,269            use_ipex = use_ipex,270            bf16 = bf16,271            fp16 = fp16,272            fp16_opt_level = fp16_opt_level,273            half_precision_backend = half_precision_backend,274            bf16_full_eval = bf16_full_eval,275            fp16_full_eval = fp16_full_eval,276            tf32 = tf32,277            local_rank = local_rank,278            ddp_backend = ddp_backend,279            tpu_num_cores = tpu_num_cores,280            tpu_metrics_debug = tpu_metrics_debug,281            debug = debug,282            dataloader_drop_last = dataloader_drop_last,283            eval_steps = eval_steps,284            dataloader_num_workers = dataloader_num_workers,285            dataloader_prefetch_factor = dataloader_prefetch_factor,286            past_index = past_index,287            run_name = run_name,288            disable_tqdm = disable_tqdm,289            remove_unused_columns = remove_unused_columns,290            label_names = label_names,291            load_best_model_at_end = load_best_model_at_end,292            metric_for_best_model = metric_for_best_model,293            greater_is_better = greater_is_better,294            ignore_data_skip = ignore_data_skip,295            fsdp = fsdp,296            fsdp_min_num_params = fsdp_min_num_params,297            fsdp_config = fsdp_config,298            tp_size = tp_size,299            fsdp_transformer_layer_cls_to_wrap = fsdp_transformer_layer_cls_to_wrap,300            accelerator_config = accelerator_config,301            deepspeed = deepspeed,302            label_smoothing_factor = label_smoothing_factor,303            optim = optim,304            optim_args = optim_args,305            adafactor = adafactor,306            group_by_length = group_by_length,307            length_column_name = length_column_name,308            report_to = report_to,309            ddp_find_unused_parameters = ddp_find_unused_parameters,310            ddp_bucket_cap_mb = ddp_bucket_cap_mb,311            ddp_broadcast_buffers = ddp_broadcast_buffers,312            dataloader_pin_memory = dataloader_pin_memory,313            dataloader_persistent_workers = dataloader_persistent_workers,314            skip_memory_metrics = skip_memory_metrics,315            use_legacy_prediction_loop = use_legacy_prediction_loop,316            push_to_hub = push_to_hub,317            resume_from_checkpoint = resume_from_checkpoint,318            hub_model_id = hub_model_id,319            hub_strategy = hub_strategy,320            hub_token = hub_token,321            hub_private_repo = hub_private_repo,322            hub_always_push = hub_always_push,323            gradient_checkpointing = gradient_checkpointing,324            gradient_checkpointing_kwargs = gradient_checkpointing_kwargs,325            include_inputs_for_metrics = include_inputs_for_metrics,326            eval_do_concat_batches = eval_do_concat_batches,327            fp16_backend = fp16_backend,328            evaluation_strategy = evaluation_strategy,329            push_to_hub_model_id = push_to_hub_model_id,330            push_to_hub_organization = push_to_hub_organization,331            push_to_hub_token = push_to_hub_token,332            mp_parameters = mp_parameters,333            auto_find_batch_size = auto_find_batch_size,334            full_determinism = full_determinism,335            torchdynamo = torchdynamo,336            ray_scope = ray_scope,337            ddp_timeout = ddp_timeout,338            torch_compile = torch_compile,339            torch_compile_backend = torch_compile_backend,340            torch_compile_mode = torch_compile_mode,341            dispatch_batches = dispatch_batches,342            split_batches = split_batches,343            include_tokens_per_second = include_tokens_per_second,344            include_num_input_tokens_seen = include_num_input_tokens_seen,345            neftune_noise_alpha = neftune_noise_alpha,346            optim_target_modules = optim_target_modules,347            batch_eval_metrics = batch_eval_metrics,348            eval_on_start = eval_on_start,349            use_liger_kernel = use_liger_kernel,350            eval_use_gather_object = eval_use_gather_object,351            average_tokens_across_devices = average_tokens_across_devices,352            reward_model_path = reward_model_path,353            judge = judge,354            max_new_tokens = max_new_tokens,355            max_length = max_length,356            temperature = temperature,357            missing_eos_penalty = missing_eos_penalty,358            loss_type = loss_type,359            dataset_num_proc = dataset_num_proc,360            disable_dropout = disable_dropout,361            use_vllm = use_vllm,362            ds3_gather_for_generation = ds3_gather_for_generation,**kwargs)363        self.vllm_sampling_params = vllm_sampling_params364        self.unsloth_num_chunks = unsloth_num_chunks365pass366 367class _UnslothNashMDTrainer(OnlineDPOTrainer):368    r""""""369 370    _tag_names = ["trl", "nash-md"]371 372    def __init__(373        self,374        model: Union[PreTrainedModel, nn.Module] = None,375        ref_model: Union[PreTrainedModel, nn.Module] = None,376        reward_model: Union[PreTrainedModel, nn.Module, None] = None,377        judge: Optional[BasePairwiseJudge] = None,378        args: Optional[NashMDConfig] = None,379        data_collator: Optional[Callable] = None,380        train_dataset: Optional[Union[Dataset, IterableDataset]] = None,381        eval_dataset: Optional[Union[Dataset, dict[str, Dataset]]] = None,382        processing_class: Optional[383            Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin]384        ] = None,385        peft_config: Optional[dict] = None,386        compute_metrics: Optional[Callable[[EvalPrediction], dict]] = None,387        callbacks: Optional[list[TrainerCallback]] = None,388        optimizers: tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),389        preprocess_logits_for_metrics: Optional[Callable[[torch.Tensor, torch.Tensor], torch.Tensor]] = None,390    ) -> None:391        super().__init__(392            model=model,393            ref_model=ref_model,394            reward_model=reward_model,395            judge=judge,396            args=args,397            data_collator=data_collator,398            train_dataset=train_dataset,399            eval_dataset=eval_dataset,400            processing_class=processing_class,401            reward_processing_class=processing_class,  # for now, NashMDTrainer can't use any reward model402            peft_config=peft_config,403            compute_metrics=compute_metrics,404            callbacks=callbacks,405            optimizers=optimizers,406            preprocess_logits_for_metrics=preprocess_logits_for_metrics,407        )408 409        self._mixture_coef = self.args.mixture_coef410 411        # Overwrite the stats dictionary to include NashMD specific statistics412        self.stats = {413            # Remove "non_score_reward", "rlhf_reward", "scores_margin"414            # Add "mixture_coef"415            "loss/kl": [],416            "objective/entropy": [],417            "loss/score": [],418            "rewards/probabilities": [],419            "rewards/accuracies": [],420            "rewards/margins": [],421            "logps/chosen": [],422            "logps/rejected": [],423            "val/model_contain_eos_token": [],424            "val/ref_contain_eos_token": [],425            "beta": [],426            "mixture_coef": [],427        }428        if self.reward_model is not None:429            self.stats["rewards/chosen"] = []430            self.stats["rewards/rejected"] = []431 432    @property433    def mixture_coef(self):434        if isinstance(self._mixture_coef, list):435            epoch = self.state.epoch436            return self._mixture_coef[epoch] if epoch < len(self._mixture_coef) else self._mixture_coef[-1]437        else:438            return self._mixture_coef439 440    def _generate_completions(self, model, prompts):441        with unwrap_model_for_generation(model, self.accelerator) as unwrapped_model:442            model_output = unwrapped_model.generate(443                input_ids=prompts["input_ids"],444                attention_mask=prompts["attention_mask"],445                generation_config=self.generation_config,446            )447 448            ref_model = model if self.ref_model is None else self.ref_model449            with torch.no_grad(), unwrap_model_for_generation(ref_model, self.accelerator) as unwrapped_ref_model:450                mixture_model = GeometricMixtureWrapper(451                    model=unwrapped_model,452                    ref_model=unwrapped_ref_model,453                    generation_config=self.generation_config,454                    mixture_coef=self.mixture_coef,455                    device=self.accelerator.device,456                )457 458                mixture_output = mixture_model.generate(459                    input_ids=prompts["input_ids"],460                    attention_mask=prompts["attention_mask"],461                    generation_config=self.generation_config,462                )463 464        return model_output, mixture_output465 466    def _process_completions(self, model_output, mixture_output, prompts):467        context_length = prompts["input_ids"].shape[1]468 469        # Process model completions470        model_completion_ids = model_output[:, context_length:]471        model_completion_ids, model_completion_mask = truncate_right(472            model_completion_ids, self.processing_class.eos_token_id, self.processing_class.pad_token_id473        )474        model_data = {475            "input_ids": torch.cat((prompts["input_ids"], model_completion_ids), dim=1),476            "attention_mask": torch.cat((prompts["attention_mask"], model_completion_mask), dim=1),477            "raw": prompts["raw"],478        }479 480        # Process reference model completions481        mixture_completion_ids = mixture_output[:, context_length:]482        mixture_completion_ids, mixture_completion_mask = truncate_right(483            mixture_completion_ids, self.processing_class.eos_token_id, self.processing_class.pad_token_id484        )485        mixture_data = {486            "input_ids": torch.cat((prompts["input_ids"], mixture_completion_ids), dim=1),487            "attention_mask": torch.cat((prompts["attention_mask"], mixture_completion_mask), dim=1),488            "raw": prompts["raw"],489        }490 491        return model_data, mixture_data492 493    def _compute_rewards(self, model_data, mixture_data, context_length):494        with torch.no_grad():495            _, model_scores, _ = get_reward(496                self.reward_model, model_data["input_ids"], self.processing_class.pad_token_id, context_length497            )498            _, mixture_scores, _ = get_reward(499                self.reward_model, mixture_data["input_ids"], self.processing_class.pad_token_id, context_length500            )501 502        # Apply EOS penalty if needed503        if self.args.missing_eos_penalty is not None:504            model_contain_eos = torch.any(model_data["input_ids"] == self.processing_class.eos_token_id, dim=-1)505            mixture_contain_eos = torch.any(mixture_data["input_ids"] == self.processing_class.eos_token_id, dim=-1)506            model_scores[~model_contain_eos] -= self.args.missing_eos_penalty507            mixture_scores[~mixture_contain_eos] -= self.args.missing_eos_penalty508 509        return model_scores, mixture_scores510 511    def _compute_judge(self, model_data, mixture_data, context_length):512        prompts = model_data["raw"]513        model_data_completions = self.processing_class.batch_decode(514            model_data["input_ids"][:, context_length:], skip_special_tokens=True515        )516        model_data_completions = [completion.strip() for completion in model_data_completions]517 518        mixture_data_completions = self.processing_class.batch_decode(519            mixture_data["input_ids"][:, context_length:], skip_special_tokens=True520        )521        mixture_data_completions = [completion.strip() for completion in mixture_data_completions]522        if is_conversational({"prompt": prompts[0]}):523            model_data_completions = [524                [{"role": "assistant", "content": completion}] for completion in model_data_completions525            ]526            environment = jinja2.Environment()527            template = environment.from_string(SIMPLE_CHAT_TEMPLATE)528            prompts = [template.render(messages=message) for message in prompts]529            model_data_completions = [template.render(messages=completion) for completion in model_data_completions]530 531            mixture_data_completions = [532                [{"role": "assistant", "content": completion}] for completion in mixture_data_completions533            ]534            mixture_data_completions = [535                template.render(messages=completion) for completion in mixture_data_completions536            ]537 538        probability = self.judge.judge(539            prompts,540            list(zip(model_data_completions, mixture_data_completions)),541            return_scores=True,542        )543        return torch.tensor(probability, device=model_data["input_ids"].device)544 545    def _compute_logprobs(self, model, model_data, context_length):546        def compute_logprobs_for_data(m, data):547            output = m(data["input_ids"], attention_mask=data["attention_mask"])548            logits = output.logits[:, context_length - 1 : -1]549            token_logprobs = selective_log_softmax(logits, data["input_ids"][:, context_length:])550            return token_logprobs551 552        # Compute logprobs for model completions under the model553        model_logprobs_model_data = compute_logprobs_for_data(model, model_data)554 555        # Compute logprobs of model completions under the reference model556        with torch.no_grad():557            if self.ref_model is None:558                with model.disable_adapter():559                    ref_logprobs_model_data = compute_logprobs_for_data(model, model_data)560            else:561                ref_logprobs_model_data = compute_logprobs_for_data(self.ref_model, model_data)562 563        # Mask padding tokens564        model_padding_mask = model_data["attention_mask"][:, context_length:] == 0565        model_logprobs_model_data = model_logprobs_model_data.masked_fill(model_padding_mask, 0.0)566        ref_logprobs_model_data = ref_logprobs_model_data.masked_fill(model_padding_mask, 0.0)567 568        return (model_logprobs_model_data, ref_logprobs_model_data)569 570    def _compute_losses(571        self,572        model_logprobs_model_data,573        ref_logprobs_model_data,574        probability,575    ):576        # reinforce score where 0.5 is a control variate577        score = (probability - 0.5) * model_logprobs_model_data.sum(1)578 579        # kl divergence via reinforce580        with torch.no_grad():581            log_ratio = model_logprobs_model_data - ref_logprobs_model_data582            kl_div_log = log_ratio.sum(1)583        kl_div_loss = (log_ratio * model_logprobs_model_data).sum(1)584 585        # final loss586        loss = self.beta * kl_div_loss - score587 588        return loss.mean(), score, kl_div_log589 590    def _log_statistics(591        self,592        model_data,593        mixture_data,594        model_logprobs_model_data,595        ref_logprobs_model_data,596        probability,597        score,598        kl_div,599        context_length,600        model_scores=None,601        mixture_scores=None,602    ):603        # Helper function to gather and compute mean604        def gather_mean(tensor):605            return self.accelerator.gather_for_metrics(tensor).mean().item()606 607        # Log score608        self.stats["loss/score"].append(gather_mean(score))609        # Log KL divergence610        self.stats["loss/kl"].append(gather_mean(kl_div))611 612        # Log logprobs613        model_logprobs_model_data_sum = model_logprobs_model_data.sum(1)614        ref_logprobs_model_data_sum = ref_logprobs_model_data.sum(1)615 616        self.stats["logps/chosen"].append(gather_mean(model_logprobs_model_data_sum))617        self.stats["logps/rejected"].append(gather_mean(ref_logprobs_model_data_sum))618 619        # Log rewards620        if self.reward_model is not None:621            self.stats["rewards/chosen"].append(gather_mean(model_scores))622            self.stats["rewards/rejected"].append(gather_mean(mixture_scores))623 624        # Log probabilities625        self.stats["rewards/probabilities"].append(gather_mean(probability))626 627        # Calculate entropy for model data628        entropy_model_data = -model_logprobs_model_data.sum(1)629        self.stats["objective/entropy"].append(gather_mean(entropy_model_data))630 631        # Calculate margins632        margin = model_logprobs_model_data_sum - ref_logprobs_model_data_sum633        self.stats["rewards/margins"].append(gather_mean(margin))634 635        # Calculate accuracy636        accuracy = (margin > 0).float()637        self.stats["rewards/accuracies"].append(gather_mean(accuracy))638 639        # Log EOS token statistics640        model_eos = (model_data["input_ids"][:, context_length:] == self.processing_class.eos_token_id).any(dim=1)641        mixture_eos = (mixture_data["input_ids"][:, context_length:] == self.processing_class.eos_token_id).any(dim=1)642        self.stats["val/model_contain_eos_token"].append(gather_mean(model_eos.float()))643        self.stats["val/ref_contain_eos_token"].append(gather_mean(mixture_eos.float()))644 645        # Log beta and mixture coef646        self.stats["beta"].append(self.beta)647        self.stats["mixture_coef"].append(self.mixture_coef)648 649    def training_step(650        self, model: nn.Module, inputs: dict[str, Union[torch.Tensor, Any]], num_items_in_batch: Optional[int] = None651    ) -> torch.Tensor:652        model.train()653 654        # Apply chat template and tokenize the input655        batch_size = len(next(iter(inputs.values())))656        prompts = inputs["prompt"]657        inputs = [{k: v[i] for k, v in inputs.items()} for i in range(batch_size)]658        inputs = [maybe_apply_chat_template(x, self.processing_class) for x in inputs]659        inputs = [self.tokenize_row(x, self.model.config.is_encoder_decoder, self.processing_class) for x in inputs]660        inputs = self.data_collator(inputs)661 662        # need the prompt_ only663        inputs = self._prepare_inputs(inputs)664        context_length = inputs["prompt_input_ids"].shape[1]665        prompts = {666            "input_ids": inputs["prompt_input_ids"],667            "attention_mask": inputs["prompt_attention_mask"],668            "raw": prompts,669        }670        del inputs671 672        # Sample completions from both the model and the reference model673        model_output, mixture_output = self._generate_completions(model, prompts)674 675        # Process model completions676        model_data, mixture_data = self._process_completions(model_output, mixture_output, prompts)677 678        # Compute rewards679        if self.reward_model is not None:680            model_scores, mixture_scores = self._compute_rewards(model_data, mixture_data, context_length)681            # probability of the model data vs the mixture data682            probability = F.sigmoid(model_scores - mixture_scores)683        else:684            model_scores, mixture_scores = None, None685            probability = self._compute_judge(model_data, mixture_data, context_length)686 687        # Compute logprobs688        model_logprobs_model_data, ref_logprobs_model_data = self._compute_logprobs(model, model_data, context_length)689 690        # Compute loss691        loss, score, kl_div = self._compute_losses(model_logprobs_model_data, ref_logprobs_model_data, probability)692 693        # Log everything694        self._log_statistics(695            model_data,696            mixture_data,697            model_logprobs_model_data.detach(),698            ref_logprobs_model_data,699            probability,700            score.detach(),701            kl_div.detach(),702            context_length,703            model_scores,704            mixture_scores,705        )706 707        if (708            self.args.torch_empty_cache_steps is not None709            and self.state.global_step % self.args.torch_empty_cache_steps == 0710        ):711            empty_cache()712 713        kwargs = {}714        # For LOMO optimizers you need to explicitly use the learning rate715        if self.args.optim in [OptimizerNames.LOMO, OptimizerNames.ADALOMO]:716            kwargs["learning_rate"] = self._get_learning_rate()717 718        if self.args.n_gpu > 1:719            loss = loss.mean()  # mean() to average on multi-gpu parallel training720 721        if self.use_apex:722            with amp.scale_loss(loss, self.optimizer) as scaled_loss:723                scaled_loss.backward()724        else:725            self.accelerator.backward(loss, **kwargs)726 727        return loss.detach() / self.args.gradient_accumulation_steps728 729    def create_model_card(730        self,731        model_name: Optional[str] = None,732        dataset_name: Optional[str] = None,733        tags: Union[str, list[str], None] = None,734    ):735        """736        Creates a draft of a model card using the information available to the `Trainer`.737 738        Args:739            model_name (`str` or `None`, *optional*, defaults to `None`):740                Name of the model.741            dataset_name (`str` or `None`, *optional*, defaults to `None`):742                Name of the dataset used for training.743            tags (`str`, `list[str]` or `None`, *optional*, defaults to `None`):744                Tags to be associated with the model card.745        """746        if not self.is_world_process_zero():747            return748 749        if hasattr(self.model.config, "_name_or_path") and not os.path.isdir(self.model.config._name_or_path):750            base_model = self.model.config._name_or_path751        else:752            base_model = None753 754        tags = tags or []755        if isinstance(tags, str):756            tags = [tags]757 758        if hasattr(self.model.config, "unsloth_version"):759            tags.append("unsloth")760 761        citation = textwrap.dedent("""\762        @inproceedings{munos2024nash,763            title        = {{Nash Learning from Human Feedback}},764            author       = {R{\'{e}}mi Munos and Michal Valko and Daniele Calandriello and Mohammad Gheshlaghi Azar and Mark Rowland and Zhaohan Daniel Guo and Yunhao Tang and Matthieu Geist and Thomas Mesnard and C{\\^{o}}me Fiegel and Andrea Michi and Marco Selvi and Sertan Girgin and Nikola Momchev and Olivier Bachem and Daniel J. Mankowitz and Doina Precup and Bilal Piot},765            year         = 2024,766            booktitle    = {Forty-first International Conference on Machine Learning, {ICML} 2024, Vienna, Austria, July 21-27, 2024},767            publisher    = {OpenReview.net},768            url          = {https://openreview.net/forum?id=Y5AmNYiyCQ}769        }""")770 771        model_card = generate_model_card(772            base_model=base_model,773            model_name=model_name,774            hub_model_id=self.hub_model_id,775            dataset_name=dataset_name,776            tags=tags,777            wandb_url=wandb.run.get_url() if is_wandb_available() and wandb.run is not None else None,778            comet_url=get_comet_experiment_url(),779            trainer_name="Nash-MD",780            trainer_citation=citation,781            paper_title="Nash Learning from Human Feedback",782            paper_id="2312.00886",783        )784 785        model_card.save(os.path.join(self.args.output_dir, "README.md"))786class UnslothNashMDTrainer(_UnslothNashMDTrainer):787    """788    789    Initialize NashMDTrainer as a subclass of [`OnlineDPOConfig`].790 791    Args:792        model (`transformers.PreTrainedModel`):793            The model to train, preferably an `AutoModelForCausalLM`.794        ref_model (`PreTrainedModelWrapper`):795            Hugging Face transformer model with a casual language modelling head. Used for implicit reward computation and loss. If no796            reference model is provided, the trainer will create a reference model with the same architecture as the model to be optimized.797        reward_model (`transformers.PreTrainedModel`):798            The reward model to score completions with, preferably an `AutoModelForSequenceClassification`.799        judge (`BasePairwiseJudge`):800            The judge to use for pairwise comparison of model completions.801        args (`NashMDConfig`):802            The NashMD config arguments to use for training.803        data_collator (`transformers.DataCollator`):804            The data collator to use for training. If None is specified, the default data collator (`DPODataCollatorWithPadding`) will be used805            which will pad the sequences to the maximum length of the sequences in the batch, given a dataset of paired sequences.806        train_dataset (`datasets.Dataset`):807            The dataset to use for training.808        eval_dataset (`datasets.Dataset`):809            The dataset to use for evaluation.810        processing_class (`PreTrainedTokenizerBase` or `BaseImageProcessor` or `FeatureExtractionMixin` or `ProcessorMixin`, *optional*):811            Processing class used to process the data. If provided, will be used to automatically process the inputs812            for the model, and it will be saved along the model to make it easier to rerun an interrupted training or813            reuse the fine-tuned model.814        peft_config (`dict`):815            The peft config to use for training.816        compute_metrics (`Callable[[EvalPrediction], dict]`, *optional*):817            The function to use to compute the metrics. Must take a `EvalPrediction` and return818            a dictionary string to metric values.819        callbacks (`list[transformers.TrainerCallback]`):820            The callbacks to use for training.821        optimizers (`tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR]`):822            The optimizer and scheduler to use for training.823        preprocess_logits_for_metrics (`Callable[[torch.Tensor, torch.Tensor], torch.Tensor]`):824            The function to use to preprocess the logits before computing the metrics.825    826    """827    def __init__(828        self,829        model = None,830        ref_model = None,831        reward_model = None,832        judge = None,833        args = None,834        data_collator = None,835        train_dataset = None,836        eval_dataset = None,837        processing_class = None,838        peft_config = None,839        compute_metrics = None,840        callbacks = None,841        preprocess_logits_for_metrics = None,842        **kwargs843    ):844        if args is None: args = UnslothNashMDConfig()845        use_bf16 = getattr(args, 'bf16', False)846        use_fp16 = getattr(args, 'fp16', False)847        force_float32 = False848        if os.environ.get('UNSLOTH_FORCE_FLOAT32', '0') == '1':849            print('Unsloth: Switching to float32 training since model cannot work with float16')850            force_float32 = True851        mixed_precision_dtype = os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32')852        dtype = getattr(model.config, 'torch_dtype', None)853        if dtype is None: dtype = model.get_input_embeddings().dtype854        from unsloth_zoo.utils import _get_dtype855        dtype = _get_dtype(dtype)856        float16 = dtype == torch.float16857        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`')858        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`')859        if force_float32:860            args.fp16 = False861            args.bf16 = False862            os.environ['ACCELERATE_MIXED_PRECISION'] = 'no'863        elif (not use_bf16 and not use_fp16) and mixed_precision_dtype == 'float32':864            args.fp16 = float16865            args.bf16 = not float16866            os.environ['ACCELERATE_MIXED_PRECISION'] = 'fp16' if float16 else 'bf16'867        if getattr(args, 'eval_dataset', None) is not None and getattr(args, 'eval_strategy', 'no') == 'no':868            args.eval_strategy = 'steps'869            if getattr(args, 'eval_steps', None) is None: args.eval_steps = 0.1870        ga_steps = getattr(args, 'gradient_accumulation_steps', None)871        if ga_steps is not None and ga_steps > 1:872            from transformers import __version__ as transformers_version873            if Version(transformers_version) <= Version('4.45.2'):874                print('**** Unsloth: Please use our fixed gradient_accumulation_steps by updating transformers, TRL and Unsloth!\n'875                      '`pip install --upgrade --no-cache-dir --force-reinstall --no-deps unsloth transformers trl unsloth_zoo`')876        if getattr(args, 'eval_strategy', 'no') != 'no':877            eval_bsz = getattr(args, 'per_device_eval_batch_size', 8)878            if eval_bsz == 8 and args.per_device_train_batch_size < eval_bsz: args.per_device_eval_batch_size = args.per_device_train_batch_size879            if getattr(args, 'eval_accumulation_steps', None) is None and ga_steps is not None: args.eval_accumulation_steps = ga_steps880        fp16_full_eval = getattr(args, 'fp16_full_eval', False)881        bf16_full_eval = getattr(args, 'bf16_full_eval', False)882        if args.fp16 and bf16_full_eval: args.bf16_full_eval = False; args.fp16_full_eval = True883        if args.bf16 and fp16_full_eval: args.bf16_full_eval = True; args.fp16_full_eval = False884        if force_float32:885            args.bf16_full_eval = False886            args.fp16_full_eval = False887        elif os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32') == 'bfloat16':888            args.bf16_full_eval = True889            args.fp16_full_eval = False890        elif not bf16_full_eval and not fp16_full_eval:891            args.bf16_full_eval = args.bf16892            args.fp16_full_eval = args.fp16893        _output_logits = False894        if locals().get('compute_metrics', None) is not None: _output_logits = True895        if locals().get('preprocess_logits_for_metrics', None) is not None: _output_logits = True896        if _output_logits:897            os.environ['UNSLOTH_RETURN_LOGITS'] = '1'898        if 'max_seq_length' not in locals() and not hasattr(args, 'max_seq_length'):899            pass900        else:901            model_max_seq_length = getattr(model, 'max_seq_length', None)902            args_max_seq_length  = getattr(args,  'max_seq_length', None)903            if args_max_seq_length is None and model_max_seq_length is not None:904                max_seq_length = model.max_seq_length905                if hasattr(args, 'max_seq_length'): args.max_seq_length = max_seq_length906        if model is not None and hasattr(model, 'for_training'):907            model.for_training()908        if 'tokenizer' in locals() and hasattr(tokenizer, 'padding_side'): tokenizer.padding_side = 'right'909        if 'processing_class' in locals():910            if hasattr(processing_class, 'padding_side'): processing_class.padding_side = 'right'911            if hasattr(processing_class, 'tokenizer') and hasattr(processing_class.tokenizer, 'padding_side'): processing_class.tokenizer.padding_side = 'right'912        __tokenizer = processing_class if 'processing_class' in locals() else tokenizer913        from unsloth_zoo.vision_utils import UnslothVisionDataCollator914        if not isinstance(data_collator, UnslothVisionDataCollator):915            if isinstance(data_collator, DataCollatorForSeq2Seq) and 'labels' not in train_dataset.column_names:916                data_collator = DataCollatorForLanguageModeling(__tokenizer, mlm = False)917            elif isinstance(data_collator, DataCollatorForLanguageModeling) and 'labels' in train_dataset.column_names:918                data_collator = DataCollatorForSeq2Seq(__tokenizer)919        else:920            if hasattr(args, 'remove_unused_columns'): args.remove_unused_columns = False921            if hasattr(args, 'dataset_text_field'): args.dataset_text_field = ''922            if hasattr(args, 'dataset_kwargs'): args.dataset_kwargs = {'skip_prepare_dataset': True}923        if not isinstance(data_collator, UnslothVisionDataCollator):924            if not hasattr(__tokenizer, 'pad') and hasattr(__tokenizer, 'tokenizer'):925                if isinstance(data_collator, DataCollatorForSeq2Seq):926                    data_collator = DataCollatorForSeq2Seq(__tokenizer.tokenizer)927                else:928                    data_collator = DataCollatorForLanguageModeling(__tokenizer.tokenizer, mlm = False)929        other_metrics = []930        931        from unsloth_zoo.logging_utils import PatchRLStatistics932        PatchRLStatistics('nash_md_trainer', other_metrics)933        934        super().__init__(935            model = model,936            ref_model = ref_model,937            reward_model = reward_model,938            judge = judge,939            args = args,940            data_collator = data_collator,941            train_dataset = train_dataset,942            eval_dataset = eval_dataset,943            processing_class = processing_class,944            peft_config = peft_config,945            compute_metrics = compute_metrics,946            callbacks = callbacks,947            preprocess_logits_for_metrics = preprocess_logits_for_metrics,**kwargs)948        if hasattr(self, 'neftune_hook_handle'):949            self.neftune_hook_handle.remove()950            if hasattr(self, 'neftune_hook_handle'): del self.neftune_hook_handle951        if getattr(args, 'neftune_noise_alpha', None) is not None:952            model.get_input_embeddings().neftune_noise_alpha = self.neftune_noise_alpha953        pass954        955pass956