Team Ai
Apppublic

Zwounds/Boolean_Search_Query_Model

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
UnslothGKDTrainer.py864 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.gkd_trainer import (Any, AutoModelForCausalLM, BaseImageProcessor, Callable, DataCollator, DataCollatorForChatML, Dataset, EvalPrediction, F, FeatureExtractionMixin, GKDConfig, GKDTrainer, GenerationConfig, Optional, PeftConfig, PreTrainedModel, PreTrainedModelWrapper, PreTrainedTokenizerBase, ProcessorMixin, SFTTrainer, TrainerCallback, Union, deepcopy, disable_dropout_in_model, empty_cache, generate_model_card, get_comet_experiment_url, is_wandb_available, nn, os, random, textwrap, torch, 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 UnslothGKDConfig(GKDConfig):44    """45    46    Configuration class for [`GKDTrainer`].47 48    Args:49        temperature (`float`, *optional*, defaults to `0.9`):50            Temperature for sampling. The higher the temperature, the more random the completions.51        lmbda (`float`, *optional*, defaults to `0.5`):52            Lambda parameter that controls the student data fraction (i.e., the proportion of on-policy53            student-generated outputs).54        beta (`float`, *optional*, defaults to `0.5`):55            Interpolation coefficient between `0.0` and `1.0` of the Generalized Jensen-Shannon Divergence loss. When56            beta is `0.0`, the loss is the KL divergence. When beta is `1.0`, the loss is the Inverse KL Divergence.57        max_new_tokens (`int`, *optional*, defaults to `128`):58            Maximum number of tokens to generate per completion.59        teacher_model_name_or_path (`str` or `None`, *optional*, defaults to `None`):60            Model name or path of the teacher model. If `None`, the teacher model will be the same as the model61            being trained.62        teacher_model_init_kwargs (`dict[str, Any]]` or `None`, *optional*, defaults to `None`):63            Keyword arguments to pass to `AutoModelForCausalLM.from_pretrained` when instantiating the teacher model64            from a string.65        disable_dropout (`bool`, *optional*, defaults to `True`):66            Whether to disable dropout in the model.67        seq_kd (`bool`, *optional*, defaults to `False`):68            Seq_kd parameter that controls whether to perform Sequence-Level KD (can be viewed as supervised FT69            on teacher-generated output).70    71    """72    vllm_sampling_params: Optional[Any] = field(73        default = None,74        metadata = {'help': 'vLLM SamplingParams'},75    )76    unsloth_num_chunks : Optional[int] = field(77        default = -1,78        metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},79    )80    def __init__(81        self,82        output_dir = None,83        overwrite_output_dir = None,84        do_train = False,85        do_eval = False,86        do_predict = False,87        eval_strategy = 'no',88        prediction_loss_only = False,89        per_device_train_batch_size = 4,90        per_device_eval_batch_size = 4,91        per_gpu_train_batch_size = None,92        per_gpu_eval_batch_size = None,93        gradient_accumulation_steps = 2,94        eval_accumulation_steps = 2,95        eval_delay = 0,96        torch_empty_cache_steps = 250,97        learning_rate = 5e-05,98        weight_decay = 0.01,99        adam_beta1 = 0.9,100        adam_beta2 = 0.999,101        adam_epsilon = 1e-08,102        max_grad_norm = 1.0,103        num_train_epochs = 3.0,104        max_steps = -1,105        lr_scheduler_type = 'linear',106        warmup_ratio = 0.1,107        warmup_steps = 0,108        log_level = 'passive',109        log_level_replica = 'warning',110        log_on_each_node = True,111        logging_dir = None,112        logging_strategy = 'steps',113        logging_first_step = False,114        logging_steps = 1,115        logging_nan_inf_filter = False,116        save_strategy = 'steps',117        save_steps = 500,118        save_total_limit = None,119        save_safetensors = True,120        save_on_each_node = False,121        save_only_model = False,122        restore_callback_states_from_checkpoint = False,123        no_cuda = False,124        use_cpu = False,125        use_mps_device = False,126        seed = 3407,127        data_seed = 3407,128        jit_mode_eval = False,129        use_ipex = False,130        bf16 = False,131        fp16 = False,132        fp16_opt_level = 'O1',133        half_precision_backend = 'auto',134        bf16_full_eval = False,135        fp16_full_eval = False,136        tf32 = None,137        local_rank = -1,138        ddp_backend = None,139        tpu_num_cores = None,140        tpu_metrics_debug = False,141        debug = '',142        dataloader_drop_last = False,143        eval_steps = None,144        dataloader_num_workers = 0,145        dataloader_prefetch_factor = None,146        past_index = -1,147        run_name = None,148        disable_tqdm = None,149        remove_unused_columns = True,150        label_names = None,151        load_best_model_at_end = False,152        metric_for_best_model = None,153        greater_is_better = None,154        ignore_data_skip = False,155        fsdp = '',156        fsdp_min_num_params = 0,157        fsdp_config = None,158        tp_size = 0,159        fsdp_transformer_layer_cls_to_wrap = None,160        accelerator_config = None,161        deepspeed = None,162        label_smoothing_factor = 0.0,163        optim = 'adamw_8bit',164        optim_args = None,165        adafactor = False,166        group_by_length = False,167        length_column_name = 'length',168        report_to = None,169        ddp_find_unused_parameters = None,170        ddp_bucket_cap_mb = None,171        ddp_broadcast_buffers = None,172        dataloader_pin_memory = True,173        dataloader_persistent_workers = False,174        skip_memory_metrics = True,175        use_legacy_prediction_loop = False,176        push_to_hub = False,177        resume_from_checkpoint = None,178        hub_model_id = None,179        hub_strategy = 'every_save',180        hub_token = None,181        hub_private_repo = None,182        hub_always_push = False,183        gradient_checkpointing = False,184        gradient_checkpointing_kwargs = None,185        include_inputs_for_metrics = False,186        eval_do_concat_batches = True,187        fp16_backend = 'auto',188        evaluation_strategy = None,189        push_to_hub_model_id = None,190        push_to_hub_organization = None,191        push_to_hub_token = None,192        mp_parameters = '',193        auto_find_batch_size = False,194        full_determinism = False,195        torchdynamo = None,196        ray_scope = 'last',197        ddp_timeout = 1800,198        torch_compile = False,199        torch_compile_backend = None,200        torch_compile_mode = None,201        dispatch_batches = None,202        split_batches = None,203        include_tokens_per_second = False,204        include_num_input_tokens_seen = False,205        neftune_noise_alpha = None,206        optim_target_modules = None,207        batch_eval_metrics = False,208        eval_on_start = False,209        use_liger_kernel = False,210        eval_use_gather_object = False,211        average_tokens_across_devices = False,212        model_init_kwargs = None,213        use_liger = False,214        dataset_text_field = 'text',215        dataset_kwargs = None,216        dataset_num_proc = None,217        max_seq_length = None,218        packing = False,219        eval_packing = None,220        dataset_batch_size = None,221        num_of_sequences = None,222        chars_per_token = None,223        temperature = 0.9,224        lmbda = 0.5,225        beta = 0.5,226        max_new_tokens = 128,227        teacher_model_name_or_path = None,228        teacher_model_init_kwargs = None,229        disable_dropout = True,230        seq_kd = False,231        vllm_sampling_params = None,232        unsloth_num_chunks = -1,233        **kwargs,234    ):235        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!')236        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!')237        if output_dir is None and save_strategy == 'steps' and save_steps == 500:238            output_dir = 'unsloth_training_checkpoints'239            save_strategy = 'no'240        if dataset_num_proc is None:241            from multiprocessing import cpu_count242            dataset_num_proc = cpu_count()243        244        super().__init__(245            output_dir = output_dir,246            overwrite_output_dir = overwrite_output_dir,247            do_train = do_train,248            do_eval = do_eval,249            do_predict = do_predict,250            eval_strategy = eval_strategy,251            prediction_loss_only = prediction_loss_only,252            per_device_train_batch_size = per_device_train_batch_size,253            per_device_eval_batch_size = per_device_eval_batch_size,254            per_gpu_train_batch_size = per_gpu_train_batch_size,255            per_gpu_eval_batch_size = per_gpu_eval_batch_size,256            gradient_accumulation_steps = gradient_accumulation_steps,257            eval_accumulation_steps = eval_accumulation_steps,258            eval_delay = eval_delay,259            torch_empty_cache_steps = torch_empty_cache_steps,260            learning_rate = learning_rate,261            weight_decay = weight_decay,262            adam_beta1 = adam_beta1,263            adam_beta2 = adam_beta2,264            adam_epsilon = adam_epsilon,265            max_grad_norm = max_grad_norm,266            num_train_epochs = num_train_epochs,267            max_steps = max_steps,268            lr_scheduler_type = lr_scheduler_type,269            warmup_ratio = warmup_ratio,270            warmup_steps = warmup_steps,271            log_level = log_level,272            log_level_replica = log_level_replica,273            log_on_each_node = log_on_each_node,274            logging_dir = logging_dir,275            logging_strategy = logging_strategy,276            logging_first_step = logging_first_step,277            logging_steps = logging_steps,278            logging_nan_inf_filter = logging_nan_inf_filter,279            save_strategy = save_strategy,280            save_steps = save_steps,281            save_total_limit = save_total_limit,282            save_safetensors = save_safetensors,283            save_on_each_node = save_on_each_node,284            save_only_model = save_only_model,285            restore_callback_states_from_checkpoint = restore_callback_states_from_checkpoint,286            no_cuda = no_cuda,287            use_cpu = use_cpu,288            use_mps_device = use_mps_device,289            seed = seed,290            data_seed = data_seed,291            jit_mode_eval = jit_mode_eval,292            use_ipex = use_ipex,293            bf16 = bf16,294            fp16 = fp16,295            fp16_opt_level = fp16_opt_level,296            half_precision_backend = half_precision_backend,297            bf16_full_eval = bf16_full_eval,298            fp16_full_eval = fp16_full_eval,299            tf32 = tf32,300            local_rank = local_rank,301            ddp_backend = ddp_backend,302            tpu_num_cores = tpu_num_cores,303            tpu_metrics_debug = tpu_metrics_debug,304            debug = debug,305            dataloader_drop_last = dataloader_drop_last,306            eval_steps = eval_steps,307            dataloader_num_workers = dataloader_num_workers,308            dataloader_prefetch_factor = dataloader_prefetch_factor,309            past_index = past_index,310            run_name = run_name,311            disable_tqdm = disable_tqdm,312            remove_unused_columns = remove_unused_columns,313            label_names = label_names,314            load_best_model_at_end = load_best_model_at_end,315            metric_for_best_model = metric_for_best_model,316            greater_is_better = greater_is_better,317            ignore_data_skip = ignore_data_skip,318            fsdp = fsdp,319            fsdp_min_num_params = fsdp_min_num_params,320            fsdp_config = fsdp_config,321            tp_size = tp_size,322            fsdp_transformer_layer_cls_to_wrap = fsdp_transformer_layer_cls_to_wrap,323            accelerator_config = accelerator_config,324            deepspeed = deepspeed,325            label_smoothing_factor = label_smoothing_factor,326            optim = optim,327            optim_args = optim_args,328            adafactor = adafactor,329            group_by_length = group_by_length,330            length_column_name = length_column_name,331            report_to = report_to,332            ddp_find_unused_parameters = ddp_find_unused_parameters,333            ddp_bucket_cap_mb = ddp_bucket_cap_mb,334            ddp_broadcast_buffers = ddp_broadcast_buffers,335            dataloader_pin_memory = dataloader_pin_memory,336            dataloader_persistent_workers = dataloader_persistent_workers,337            skip_memory_metrics = skip_memory_metrics,338            use_legacy_prediction_loop = use_legacy_prediction_loop,339            push_to_hub = push_to_hub,340            resume_from_checkpoint = resume_from_checkpoint,341            hub_model_id = hub_model_id,342            hub_strategy = hub_strategy,343            hub_token = hub_token,344            hub_private_repo = hub_private_repo,345            hub_always_push = hub_always_push,346            gradient_checkpointing = gradient_checkpointing,347            gradient_checkpointing_kwargs = gradient_checkpointing_kwargs,348            include_inputs_for_metrics = include_inputs_for_metrics,349            eval_do_concat_batches = eval_do_concat_batches,350            fp16_backend = fp16_backend,351            evaluation_strategy = evaluation_strategy,352            push_to_hub_model_id = push_to_hub_model_id,353            push_to_hub_organization = push_to_hub_organization,354            push_to_hub_token = push_to_hub_token,355            mp_parameters = mp_parameters,356            auto_find_batch_size = auto_find_batch_size,357            full_determinism = full_determinism,358            torchdynamo = torchdynamo,359            ray_scope = ray_scope,360            ddp_timeout = ddp_timeout,361            torch_compile = torch_compile,362            torch_compile_backend = torch_compile_backend,363            torch_compile_mode = torch_compile_mode,364            dispatch_batches = dispatch_batches,365            split_batches = split_batches,366            include_tokens_per_second = include_tokens_per_second,367            include_num_input_tokens_seen = include_num_input_tokens_seen,368            neftune_noise_alpha = neftune_noise_alpha,369            optim_target_modules = optim_target_modules,370            batch_eval_metrics = batch_eval_metrics,371            eval_on_start = eval_on_start,372            use_liger_kernel = use_liger_kernel,373            eval_use_gather_object = eval_use_gather_object,374            average_tokens_across_devices = average_tokens_across_devices,375            model_init_kwargs = model_init_kwargs,376            use_liger = use_liger,377            dataset_text_field = dataset_text_field,378            dataset_kwargs = dataset_kwargs,379            dataset_num_proc = dataset_num_proc,380            max_seq_length = max_seq_length,381            packing = packing,382            eval_packing = eval_packing,383            dataset_batch_size = dataset_batch_size,384            num_of_sequences = num_of_sequences,385            chars_per_token = chars_per_token,386            temperature = temperature,387            lmbda = lmbda,388            beta = beta,389            max_new_tokens = max_new_tokens,390            teacher_model_name_or_path = teacher_model_name_or_path,391            teacher_model_init_kwargs = teacher_model_init_kwargs,392            disable_dropout = disable_dropout,393            seq_kd = seq_kd,**kwargs)394        self.vllm_sampling_params = vllm_sampling_params395        self.unsloth_num_chunks = unsloth_num_chunks396pass397 398class _UnslothGKDTrainer(SFTTrainer):399    _tag_names = ["trl", "gkd"]400 401    def __init__(402        self,403        model: Optional[Union[PreTrainedModel, nn.Module, str]] = None,404        teacher_model: Union[PreTrainedModel, nn.Module, str] = None,405        args: Optional[GKDConfig] = None,406        data_collator: Optional[DataCollator] = None,  # type: ignore407        train_dataset: Optional[Dataset] = None,408        eval_dataset: Optional[Union[Dataset, dict[str, Dataset]]] = None,409        processing_class: Optional[410            Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin]411        ] = None,412        compute_metrics: Optional[Callable[[EvalPrediction], dict]] = None,413        callbacks: Optional[list[TrainerCallback]] = None,414        optimizers: tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),415        preprocess_logits_for_metrics: Optional[Callable[[torch.Tensor, torch.Tensor], torch.Tensor]] = None,416        peft_config: Optional["PeftConfig"] = None,417        formatting_func: Optional[Callable] = None,418    ):419        # add remove_unused_columns=False to the dataclass args420        args.remove_unused_columns = False421        data_collator = DataCollatorForChatML(tokenizer=processing_class, max_length=args.max_seq_length)422 423        super().__init__(424            model,425            args=args,426            data_collator=data_collator,427            train_dataset=train_dataset,428            eval_dataset=eval_dataset,429            processing_class=processing_class,430            compute_metrics=compute_metrics,431            callbacks=callbacks,432            optimizers=optimizers,433            preprocess_logits_for_metrics=preprocess_logits_for_metrics,434            peft_config=peft_config,435            formatting_func=formatting_func,436        )437 438        if args.teacher_model_init_kwargs is None:439            teacher_model_init_kwargs = {}440        elif not isinstance(teacher_model, str):441            raise ValueError(442                "You passed teacher_model_init_kwargs to the GKDConfig, but your teacher_model is already instantiated."443            )444        else:445            teacher_model_init_kwargs = args.teacher_model_init_kwargs446            teacher_model_init_kwargs["torch_dtype"] = (447                teacher_model_init_kwargs["torch_dtype"]448                if teacher_model_init_kwargs["torch_dtype"] in ["auto", None]449                else getattr(torch, teacher_model_init_kwargs["torch_dtype"])450            )451 452        if isinstance(teacher_model, str):453            if args.use_liger:454                teacher_model = AutoLigerKernelForCausalLM.from_pretrained(teacher_model, **teacher_model_init_kwargs)455            else:456                teacher_model = AutoModelForCausalLM.from_pretrained(teacher_model, **teacher_model_init_kwargs)457 458        # Disable dropout in the model459        if args.disable_dropout:460            disable_dropout_in_model(self.model)461 462        if self.is_deepspeed_enabled:463            self.teacher_model = self._prepare_deepspeed(teacher_model)464        else:465            self.teacher_model = self.accelerator.prepare_model(teacher_model, evaluation_mode=True)466 467        self.lmbda = args.lmbda468        self.beta = args.beta469        self.temperature = args.temperature470        self.seq_kd = args.seq_kd471 472        self.generation_config = GenerationConfig(473            max_new_tokens=args.max_new_tokens,474            temperature=args.temperature,475            do_sample=True,476            top_k=0,477            use_cache=False if args.gradient_checkpointing else True,478            pad_token_id=self.processing_class.pad_token_id,479        )480        # Set custom EOS tokens if they are specified by the model's generation481        # config. This is important for models with the Llama 3 chat template,482        # which use special tokens <|eot_id|> and <|eom_id|> to mark the end of483        # turns or messages.484        if (485            hasattr(self.model.generation_config, "eos_token_id")486            and self.model.generation_config.eos_token_id is not None487        ):488            self.generation_config.eos_token_id = self.model.generation_config.eos_token_id489 490    def _prepare_dataset(self, dataset, *args):491        # SFTTrainer._prepare_dataset() applies the chat template and rename the messages column to text. However, we492        # need to keep the messages column as it is. We use the following workaround to keep the messages column.493        dataset = dataset.add_column("_messages", dataset["messages"])494        dataset = super()._prepare_dataset(dataset, *args)495        dataset = dataset.rename_column("_messages", "messages")496        return dataset497 498    @staticmethod499    def generalized_jsd_loss(500        student_logits, teacher_logits, labels=None, beta=0.5, temperature=1.0, reduction="batchmean"501    ):502        """503        Compute the generalized Jensen-Shannon Divergence loss for knowledge distillation using F.kl_div. See Eq. (1)504        of https://huggingface.co/papers/2306.13649 for the definition.505 506        Args:507            student_logits: Tensor of shape (batch_size, sequence_length, vocab_size)508            teacher_logits: Tensor of shape (batch_size, sequence_length, vocab_size)509            labels: Tensor of shape (batch_size, sequence_length) with -100 for padding tokens to ignore when computing loss510            beta: Interpolation coefficient between 0 and 1 (default: 0.5)511            temperature: Softmax temperature (default: 1.0)512            reduction: Specifies the reduction to apply to the output (default: 'batchmean')513 514        Returns:515            loss: Scalar tensor with the generalized JSD loss516        """517 518        # Apply temperature scaling519        student_logits = student_logits / temperature520        teacher_logits = teacher_logits / temperature521 522        # Compute log probabilities for student and probabilities for teacher523        student_log_probs = F.log_softmax(student_logits, dim=-1)524        teacher_log_probs = F.log_softmax(teacher_logits, dim=-1)525 526        # Compute the log of the mixture distribution527        # log(a + b) = log(exp(log(a)) + exp(log(b))) -> for mixture528        beta = torch.tensor(beta, dtype=student_log_probs.dtype)529        mixture_log_probs = torch.logsumexp(530            torch.stack([student_log_probs + torch.log(beta), teacher_log_probs + torch.log(1 - beta)]),531            dim=0,532        )533 534        # Compute KL divergences using F.kl_div535        # PyTorch differs from the standard mathematical definition, so the order of the probability distributions is swapped compared to that defined in the paper.536        kl_teacher = F.kl_div(mixture_log_probs, teacher_log_probs, reduction="none", log_target=True)537        kl_student = F.kl_div(mixture_log_probs, student_log_probs, reduction="none", log_target=True)538 539        # Compute the Generalized Jensen-Shannon Divergence540        jsd = beta * kl_teacher + (1 - beta) * kl_student541 542        # Masking543        if labels is not None:544            mask = labels != -100545            jsd = jsd[mask]546 547        # Apply reduction548        if reduction == "batchmean":549            return jsd.sum() / mask.sum() if labels is not None else jsd.sum() / (jsd.size(0) * jsd.size(1))550        elif reduction == "sum":551            return jsd.sum()552        elif reduction == "mean":553            return jsd.mean()554        else:555            return jsd556 557    def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None):558        # compute student output559        outputs_student = model(560            input_ids=inputs["input_ids"],561            attention_mask=inputs["attention_mask"],562        )563 564        # compute teacher output in eval mode565        self.teacher_model.eval()566        with torch.no_grad():567            outputs_teacher = self.teacher_model(568                input_ids=inputs["input_ids"],569                attention_mask=inputs["attention_mask"],570            )571 572        # slice the logits for the generated tokens using the inputs["prompts"] lengths573        prompt_lengths = inputs["prompts"].shape[1]574        shifted_student_logits = outputs_student.logits[:, prompt_lengths - 1 : -1, :]575        shifted_teacher_logits = outputs_teacher.logits[:, prompt_lengths - 1 : -1, :]576        shifted_labels = inputs["labels"][:, prompt_lengths:]577 578        # compute loss579        loss = self.generalized_jsd_loss(580            student_logits=shifted_student_logits,581            teacher_logits=shifted_teacher_logits,582            labels=shifted_labels,583            beta=self.beta,584        )585 586        # empty cache587        empty_cache()588 589        # Return loss590        return (loss, outputs_student) if return_outputs else loss591 592    @staticmethod593    def generate_on_policy_outputs(model, inputs, generation_config, pad_token_id=None):594        # Generate output with respect to the prompt only595        generated_outputs = model.generate(596            input_ids=inputs["prompts"],597            attention_mask=inputs.get("prompt_attention_mask", None),598            generation_config=generation_config,599            return_dict_in_generate=True,600        )601 602        # Get the generated token IDs603        generated_tokens = generated_outputs.sequences604        # Calculate new attention mask605        new_attention_mask = torch.ones_like(generated_tokens)606        new_labels = generated_tokens.clone()607 608        # If there's pad_token_id, set attention mask to 0 for padding tokens609        if pad_token_id is not None:610            new_labels[new_labels == pad_token_id] = -100611            new_attention_mask[generated_tokens == pad_token_id] = 0612 613        return generated_tokens, new_attention_mask, new_labels614 615    def training_step(616        self, model: nn.Module, inputs: dict[str, Union[torch.Tensor, Any]], num_items_in_batch: Optional[int] = None617    ) -> torch.Tensor:618        """619        Perform a training step for the Generalized Knowledge Distillation (GKD) model.620 621        This method implements the on-policy learning approach described in the GKD paper.622        With probability `self.lmbda`, it generates new responses using the student model,623        which are then used for training instead of the original inputs.624        """625        if self.seq_kd:626            with unwrap_model_for_generation(self.teacher_model, self.accelerator) as unwrapped_model:627                new_input_ids, new_attention_mask, new_labels = self.generate_on_policy_outputs(628                    unwrapped_model, inputs, self.generation_config, self.processing_class.pad_token_id629                )630            inputs["input_ids"] = new_input_ids631            inputs["attention_mask"] = new_attention_mask632            inputs["labels"] = new_labels633        if random.random() <= self.lmbda:634            with unwrap_model_for_generation(model, self.accelerator) as unwrapped_model:635                new_input_ids, new_attention_mask, new_labels = self.generate_on_policy_outputs(636                    unwrapped_model, inputs, self.generation_config, self.processing_class.pad_token_id637                )638            inputs["input_ids"] = new_input_ids639            inputs["attention_mask"] = new_attention_mask640            inputs["labels"] = new_labels641 642        loss = super().training_step(model, inputs, num_items_in_batch)643        return loss644 645    def _prepare_deepspeed(self, model: PreTrainedModelWrapper):646        # Adapted from accelerate: https://github.com/huggingface/accelerate/blob/739b135f8367becb67ffaada12fe76e3aa60fefd/src/accelerate/accelerator.py#L1473647        deepspeed_plugin = self.accelerator.state.deepspeed_plugin648        config_kwargs = deepcopy(deepspeed_plugin.deepspeed_config)649 650        if model is not None:651            if hasattr(model, "config"):652                hidden_size = (653                    max(model.config.hidden_sizes)654                    if getattr(model.config, "hidden_sizes", None)655                    else getattr(model.config, "hidden_size", None)656                )657                if hidden_size is not None and config_kwargs["zero_optimization"]["stage"] == 3:658                    # Note that `stage3_prefetch_bucket_size` can produce DeepSpeed messages like: `Invalidate trace cache @ step 0: expected module 1, but got module 0`659                    # This is expected and is not an error, see: https://github.com/microsoft/DeepSpeed/discussions/4081660                    config_kwargs.update(661                        {662                            "zero_optimization.reduce_bucket_size": hidden_size * hidden_size,663                            "zero_optimization.stage3_param_persistence_threshold": 10 * hidden_size,664                            "zero_optimization.stage3_prefetch_bucket_size": 0.9 * hidden_size * hidden_size,665                        }666                    )667 668        # If ZeRO-3 is used, we shard both the active and reference model.669        # Otherwise, we assume the reference model fits in memory and is initialized on each device with ZeRO disabled (stage 0)670        if config_kwargs["zero_optimization"]["stage"] != 3:671            config_kwargs["zero_optimization"]["stage"] = 0672        model, *_ = deepspeed.initialize(model=model, config=config_kwargs)673        model.eval()674        return model675 676    def create_model_card(677        self,678        model_name: Optional[str] = None,679        dataset_name: Optional[str] = None,680        tags: Union[str, list[str], None] = None,681    ):682        """683        Creates a draft of a model card using the information available to the `Trainer`.684 685        Args:686            model_name (`str` or `None`, *optional*, defaults to `None`):687                Name of the model.688            dataset_name (`str` or `None`, *optional*, defaults to `None`):689                Name of the dataset used for training.690            tags (`str`, `list[str]` or `None`, *optional*, defaults to `None`):691                Tags to be associated with the model card.692        """693        if not self.is_world_process_zero():694            return695 696        if hasattr(self.model.config, "_name_or_path") and not os.path.isdir(self.model.config._name_or_path):697            base_model = self.model.config._name_or_path698        else:699            base_model = None700 701        tags = tags or []702        if isinstance(tags, str):703            tags = [tags]704 705        if hasattr(self.model.config, "unsloth_version"):706            tags.append("unsloth")707 708        citation = textwrap.dedent("""\709        @inproceedings{agarwal2024on-policy,710            title        = {{On-Policy Distillation of Language Models: Learning from Self-Generated Mistakes}},711            author       = {Rishabh Agarwal and Nino Vieillard and Yongchao Zhou and Piotr Stanczyk and Sabela Ramos Garea and Matthieu Geist and Olivier Bachem},712            year         = 2024,713            booktitle    = {The Twelfth International Conference on Learning Representations, {ICLR} 2024, Vienna, Austria, May 7-11, 2024},714            publisher    = {OpenReview.net},715            url          = {https://openreview.net/forum?id=3zKtaqxLhW},716        }""")717 718        model_card = generate_model_card(719            base_model=base_model,720            model_name=model_name,721            hub_model_id=self.hub_model_id,722            dataset_name=dataset_name,723            tags=tags,724            wandb_url=wandb.run.get_url() if is_wandb_available() and wandb.run is not None else None,725            comet_url=get_comet_experiment_url(),726            trainer_name="GKD",727            trainer_citation=citation,728            paper_title="On-Policy Distillation of Language Models: Learning from Self-Generated Mistakes",729            paper_id="2306.13649",730        )731 732        model_card.save(os.path.join(self.args.output_dir, "README.md"))733class UnslothGKDTrainer(_UnslothGKDTrainer):734    """735    736    """737    def __init__(738        self,739        model = None,740        teacher_model = None,741        args = None,742        data_collator = None,743        train_dataset = None,744        eval_dataset = None,745        processing_class = None,746        compute_metrics = None,747        callbacks = None,748        preprocess_logits_for_metrics = None,749        peft_config = None,750        formatting_func = None,751        **kwargs752    ):753        if args is None: args = UnslothGKDConfig()754        use_bf16 = getattr(args, 'bf16', False)755        use_fp16 = getattr(args, 'fp16', False)756        force_float32 = False757        if os.environ.get('UNSLOTH_FORCE_FLOAT32', '0') == '1':758            print('Unsloth: Switching to float32 training since model cannot work with float16')759            force_float32 = True760        mixed_precision_dtype = os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32')761        dtype = getattr(model.config, 'torch_dtype', None)762        if dtype is None: dtype = model.get_input_embeddings().dtype763        from unsloth_zoo.utils import _get_dtype764        dtype = _get_dtype(dtype)765        float16 = dtype == torch.float16766        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`')767        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`')768        if force_float32:769            args.fp16 = False770            args.bf16 = False771            os.environ['ACCELERATE_MIXED_PRECISION'] = 'no'772        elif (not use_bf16 and not use_fp16) and mixed_precision_dtype == 'float32':773            args.fp16 = float16774            args.bf16 = not float16775            os.environ['ACCELERATE_MIXED_PRECISION'] = 'fp16' if float16 else 'bf16'776        if getattr(args, 'eval_dataset', None) is not None and getattr(args, 'eval_strategy', 'no') == 'no':777            args.eval_strategy = 'steps'778            if getattr(args, 'eval_steps', None) is None: args.eval_steps = 0.1779        ga_steps = getattr(args, 'gradient_accumulation_steps', None)780        if ga_steps is not None and ga_steps > 1:781            from transformers import __version__ as transformers_version782            if Version(transformers_version) <= Version('4.45.2'):783                print('**** Unsloth: Please use our fixed gradient_accumulation_steps by updating transformers, TRL and Unsloth!\n'784                      '`pip install --upgrade --no-cache-dir --force-reinstall --no-deps unsloth transformers trl unsloth_zoo`')785        if getattr(args, 'eval_strategy', 'no') != 'no':786            eval_bsz = getattr(args, 'per_device_eval_batch_size', 8)787            if eval_bsz == 8 and args.per_device_train_batch_size < eval_bsz: args.per_device_eval_batch_size = args.per_device_train_batch_size788            if getattr(args, 'eval_accumulation_steps', None) is None and ga_steps is not None: args.eval_accumulation_steps = ga_steps789        fp16_full_eval = getattr(args, 'fp16_full_eval', False)790        bf16_full_eval = getattr(args, 'bf16_full_eval', False)791        if args.fp16 and bf16_full_eval: args.bf16_full_eval = False; args.fp16_full_eval = True792        if args.bf16 and fp16_full_eval: args.bf16_full_eval = True; args.fp16_full_eval = False793        if force_float32:794            args.bf16_full_eval = False795            args.fp16_full_eval = False796        elif os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32') == 'bfloat16':797            args.bf16_full_eval = True798            args.fp16_full_eval = False799        elif not bf16_full_eval and not fp16_full_eval:800            args.bf16_full_eval = args.bf16801            args.fp16_full_eval = args.fp16802        _output_logits = False803        if locals().get('compute_metrics', None) is not None: _output_logits = True804        if locals().get('preprocess_logits_for_metrics', None) is not None: _output_logits = True805        if _output_logits:806            os.environ['UNSLOTH_RETURN_LOGITS'] = '1'807        if 'max_seq_length' not in locals() and not hasattr(args, 'max_seq_length'):808            pass809        else:810            model_max_seq_length = getattr(model, 'max_seq_length', None)811            args_max_seq_length  = getattr(args,  'max_seq_length', None)812            if args_max_seq_length is None and model_max_seq_length is not None:813                max_seq_length = model.max_seq_length814                if hasattr(args, 'max_seq_length'): args.max_seq_length = max_seq_length815        if model is not None and hasattr(model, 'for_training'):816            model.for_training()817        if 'tokenizer' in locals() and hasattr(tokenizer, 'padding_side'): tokenizer.padding_side = 'right'818        if 'processing_class' in locals():819            if hasattr(processing_class, 'padding_side'): processing_class.padding_side = 'right'820            if hasattr(processing_class, 'tokenizer') and hasattr(processing_class.tokenizer, 'padding_side'): processing_class.tokenizer.padding_side = 'right'821        __tokenizer = processing_class if 'processing_class' in locals() else tokenizer822        from unsloth_zoo.vision_utils import UnslothVisionDataCollator823        if not isinstance(data_collator, UnslothVisionDataCollator):824            if isinstance(data_collator, DataCollatorForSeq2Seq) and 'labels' not in train_dataset.column_names:825                data_collator = DataCollatorForLanguageModeling(__tokenizer, mlm = False)826            elif isinstance(data_collator, DataCollatorForLanguageModeling) and 'labels' in train_dataset.column_names:827                data_collator = DataCollatorForSeq2Seq(__tokenizer)828        else:829            if hasattr(args, 'remove_unused_columns'): args.remove_unused_columns = False830            if hasattr(args, 'dataset_text_field'): args.dataset_text_field = ''831            if hasattr(args, 'dataset_kwargs'): args.dataset_kwargs = {'skip_prepare_dataset': True}832        if not isinstance(data_collator, UnslothVisionDataCollator):833            if not hasattr(__tokenizer, 'pad') and hasattr(__tokenizer, 'tokenizer'):834                if isinstance(data_collator, DataCollatorForSeq2Seq):835                    data_collator = DataCollatorForSeq2Seq(__tokenizer.tokenizer)836                else:837                    data_collator = DataCollatorForLanguageModeling(__tokenizer.tokenizer, mlm = False)838        other_metrics = []839        840        from unsloth_zoo.logging_utils import PatchRLStatistics841        PatchRLStatistics('gkd_trainer', other_metrics)842        843        super().__init__(844            model = model,845            teacher_model = teacher_model,846            args = args,847            data_collator = data_collator,848            train_dataset = train_dataset,849            eval_dataset = eval_dataset,850            processing_class = processing_class,851            compute_metrics = compute_metrics,852            callbacks = callbacks,853            preprocess_logits_for_metrics = preprocess_logits_for_metrics,854            peft_config = peft_config,855            formatting_func = formatting_func,**kwargs)856        if hasattr(self, 'neftune_hook_handle'):857            self.neftune_hook_handle.remove()858            if hasattr(self, 'neftune_hook_handle'): del self.neftune_hook_handle859        if getattr(args, 'neftune_noise_alpha', None) is not None:860            model.get_input_embeddings().neftune_noise_alpha = self.neftune_noise_alpha861        pass862        863pass864