Team Ai
Apppublic

Zwounds/Boolean_Search_Query_Model

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
UnslothSFTTrainer.py1028 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.sft_trainer import (Any, AutoModelForCausalLM, AutoTokenizer, BaseImageProcessor, Callable, ConstantLengthDataset, DataCollator, DataCollatorForLanguageModeling, Dataset, EvalPrediction, FeatureExtractionMixin, IterableDataset, Optional, PeftConfig, PeftModel, PreTrainedModel, PreTrainedTokenizerBase, ProcessorMixin, SFTConfig, SFTTrainer, Trainer, TrainerCallback, TrainingArguments, Type, Union, dataclasses, defaultdict, deprecate_kwarg, generate_model_card, get_comet_experiment_url, get_peft_model, is_liger_kernel_available, is_peft_available, is_wandb_available, nn, os, pack_examples, peft, peft_module_casting_to_bf16, prepare_model_for_kbit_training, torch, transformers, version, warnings, Callable, ConstantLengthDataset, DataCollator, DataCollatorForLanguageModeling, Dataset, IterableDataset, Optional, Union, os, pack_examples, transformers, os)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 UnslothSFTConfig(SFTConfig):44    """45    46    Configuration class for the [`SFTTrainer`].47 48    Only the parameters specific to SFT training are listed here. For details on other parameters, refer to the49    [`~transformers.TrainingArguments`] documentation.50 51    Using [`~transformers.HfArgumentParser`] we can turn this class into52    [argparse](https://docs.python.org/3/library/argparse#module-argparse) arguments that can be specified on the53    command line.54 55    Parameters:56        > Parameters that control the model57 58        model_init_kwargs (`dict[str, Any]` or `None`, *optional*, defaults to `None`):59            Keyword arguments for [`~transformers.AutoModelForCausalLM.from_pretrained`], used when the `model`60            argument of the [`SFTTrainer`] is provided as a string.61        use_liger (`bool`, *optional*, defaults to `False`):62            Monkey patch the model with Liger kernels to increase throughput and reduce memory usage.63 64        > Parameters that control the data preprocessing65 66        dataset_text_field (`str`, *optional*, defaults to `"text"`):67            Name of the column that contains text data in the dataset.68        dataset_kwargs (`dict[str, Any]` or `None`, *optional*, defaults to `None`):69            Dictionary of optional keyword arguments for the dataset preparation. The only supported key is70            `skip_prepare_dataset`.71        dataset_num_proc (`int` or `None`, *optional*, defaults to `None`):72            Number of processes to use for processing the dataset.73        max_seq_length (`int` or `None`, *optional*, defaults to `1024`):74            Maximum length of the tokenized sequence. Sequences longer than `max_seq_length` are truncated from the75            right.76            If `None`, no truncation is applied. When packing is enabled, this value sets the sequence length.77        packing (`bool`, *optional*, defaults to `False`):78            Whether to pack multiple sequences into a fixed-length format. Uses `max_seq_length` to define sequence79            length.80        eval_packing (`bool` or `None`, *optional*, defaults to `None`):81            Whether to pack the eval dataset. If `None`, uses the same value as `packing`.82 83        > Parameters that control the training84 85        learning_rate (`float`, *optional*, defaults to `2e-5`):86            Initial learning rate for [`AdamW`] optimizer. The default value replaces that of87            [`~transformers.TrainingArguments`].88    89    """90    vllm_sampling_params: Optional[Any] = field(91        default = None,92        metadata = {'help': 'vLLM SamplingParams'},93    )94    unsloth_num_chunks : Optional[int] = field(95        default = -1,96        metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},97    )98    def __init__(99        self,100        output_dir = None,101        overwrite_output_dir = None,102        do_train = False,103        do_eval = False,104        do_predict = False,105        eval_strategy = 'no',106        prediction_loss_only = False,107        per_device_train_batch_size = 4,108        per_device_eval_batch_size = 4,109        per_gpu_train_batch_size = None,110        per_gpu_eval_batch_size = None,111        gradient_accumulation_steps = 2,112        eval_accumulation_steps = 2,113        eval_delay = 0,114        torch_empty_cache_steps = 250,115        learning_rate = 5e-05,116        weight_decay = 0.01,117        adam_beta1 = 0.9,118        adam_beta2 = 0.999,119        adam_epsilon = 1e-08,120        max_grad_norm = 1.0,121        num_train_epochs = 3.0,122        max_steps = -1,123        lr_scheduler_type = 'linear',124        warmup_ratio = 0.1,125        warmup_steps = 0,126        log_level = 'passive',127        log_level_replica = 'warning',128        log_on_each_node = True,129        logging_dir = None,130        logging_strategy = 'steps',131        logging_first_step = False,132        logging_steps = 1,133        logging_nan_inf_filter = False,134        save_strategy = 'steps',135        save_steps = 500,136        save_total_limit = None,137        save_safetensors = True,138        save_on_each_node = False,139        save_only_model = False,140        restore_callback_states_from_checkpoint = False,141        no_cuda = False,142        use_cpu = False,143        use_mps_device = False,144        seed = 3407,145        data_seed = 3407,146        jit_mode_eval = False,147        use_ipex = False,148        bf16 = False,149        fp16 = False,150        fp16_opt_level = 'O1',151        half_precision_backend = 'auto',152        bf16_full_eval = False,153        fp16_full_eval = False,154        tf32 = None,155        local_rank = -1,156        ddp_backend = None,157        tpu_num_cores = None,158        tpu_metrics_debug = False,159        debug = '',160        dataloader_drop_last = False,161        eval_steps = None,162        dataloader_num_workers = 0,163        dataloader_prefetch_factor = None,164        past_index = -1,165        run_name = None,166        disable_tqdm = None,167        remove_unused_columns = True,168        label_names = None,169        load_best_model_at_end = False,170        metric_for_best_model = None,171        greater_is_better = None,172        ignore_data_skip = False,173        fsdp = '',174        fsdp_min_num_params = 0,175        fsdp_config = None,176        tp_size = 0,177        fsdp_transformer_layer_cls_to_wrap = None,178        accelerator_config = None,179        deepspeed = None,180        label_smoothing_factor = 0.0,181        optim = 'adamw_8bit',182        optim_args = None,183        adafactor = False,184        group_by_length = False,185        length_column_name = 'length',186        report_to = None,187        ddp_find_unused_parameters = None,188        ddp_bucket_cap_mb = None,189        ddp_broadcast_buffers = None,190        dataloader_pin_memory = True,191        dataloader_persistent_workers = False,192        skip_memory_metrics = True,193        use_legacy_prediction_loop = False,194        push_to_hub = False,195        resume_from_checkpoint = None,196        hub_model_id = None,197        hub_strategy = 'every_save',198        hub_token = None,199        hub_private_repo = None,200        hub_always_push = False,201        gradient_checkpointing = False,202        gradient_checkpointing_kwargs = None,203        include_inputs_for_metrics = False,204        eval_do_concat_batches = True,205        fp16_backend = 'auto',206        evaluation_strategy = None,207        push_to_hub_model_id = None,208        push_to_hub_organization = None,209        push_to_hub_token = None,210        mp_parameters = '',211        auto_find_batch_size = False,212        full_determinism = False,213        torchdynamo = None,214        ray_scope = 'last',215        ddp_timeout = 1800,216        torch_compile = False,217        torch_compile_backend = None,218        torch_compile_mode = None,219        dispatch_batches = None,220        split_batches = None,221        include_tokens_per_second = False,222        include_num_input_tokens_seen = False,223        neftune_noise_alpha = None,224        optim_target_modules = None,225        batch_eval_metrics = False,226        eval_on_start = False,227        use_liger_kernel = False,228        eval_use_gather_object = False,229        average_tokens_across_devices = False,230        model_init_kwargs = None,231        use_liger = False,232        dataset_text_field = 'text',233        dataset_kwargs = None,234        dataset_num_proc = None,235        max_seq_length = None,236        packing = False,237        eval_packing = None,238        dataset_batch_size = None,239        num_of_sequences = None,240        chars_per_token = None,241        vllm_sampling_params = None,242        unsloth_num_chunks = -1,243        **kwargs,244    ):245        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!')246        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!')247        if output_dir is None and save_strategy == 'steps' and save_steps == 500:248            output_dir = 'unsloth_training_checkpoints'249            save_strategy = 'no'250        if dataset_num_proc is None:251            from multiprocessing import cpu_count252            dataset_num_proc = cpu_count()253        254        super().__init__(255            output_dir = output_dir,256            overwrite_output_dir = overwrite_output_dir,257            do_train = do_train,258            do_eval = do_eval,259            do_predict = do_predict,260            eval_strategy = eval_strategy,261            prediction_loss_only = prediction_loss_only,262            per_device_train_batch_size = per_device_train_batch_size,263            per_device_eval_batch_size = per_device_eval_batch_size,264            per_gpu_train_batch_size = per_gpu_train_batch_size,265            per_gpu_eval_batch_size = per_gpu_eval_batch_size,266            gradient_accumulation_steps = gradient_accumulation_steps,267            eval_accumulation_steps = eval_accumulation_steps,268            eval_delay = eval_delay,269            torch_empty_cache_steps = torch_empty_cache_steps,270            learning_rate = learning_rate,271            weight_decay = weight_decay,272            adam_beta1 = adam_beta1,273            adam_beta2 = adam_beta2,274            adam_epsilon = adam_epsilon,275            max_grad_norm = max_grad_norm,276            num_train_epochs = num_train_epochs,277            max_steps = max_steps,278            lr_scheduler_type = lr_scheduler_type,279            warmup_ratio = warmup_ratio,280            warmup_steps = warmup_steps,281            log_level = log_level,282            log_level_replica = log_level_replica,283            log_on_each_node = log_on_each_node,284            logging_dir = logging_dir,285            logging_strategy = logging_strategy,286            logging_first_step = logging_first_step,287            logging_steps = logging_steps,288            logging_nan_inf_filter = logging_nan_inf_filter,289            save_strategy = save_strategy,290            save_steps = save_steps,291            save_total_limit = save_total_limit,292            save_safetensors = save_safetensors,293            save_on_each_node = save_on_each_node,294            save_only_model = save_only_model,295            restore_callback_states_from_checkpoint = restore_callback_states_from_checkpoint,296            no_cuda = no_cuda,297            use_cpu = use_cpu,298            use_mps_device = use_mps_device,299            seed = seed,300            data_seed = data_seed,301            jit_mode_eval = jit_mode_eval,302            use_ipex = use_ipex,303            bf16 = bf16,304            fp16 = fp16,305            fp16_opt_level = fp16_opt_level,306            half_precision_backend = half_precision_backend,307            bf16_full_eval = bf16_full_eval,308            fp16_full_eval = fp16_full_eval,309            tf32 = tf32,310            local_rank = local_rank,311            ddp_backend = ddp_backend,312            tpu_num_cores = tpu_num_cores,313            tpu_metrics_debug = tpu_metrics_debug,314            debug = debug,315            dataloader_drop_last = dataloader_drop_last,316            eval_steps = eval_steps,317            dataloader_num_workers = dataloader_num_workers,318            dataloader_prefetch_factor = dataloader_prefetch_factor,319            past_index = past_index,320            run_name = run_name,321            disable_tqdm = disable_tqdm,322            remove_unused_columns = remove_unused_columns,323            label_names = label_names,324            load_best_model_at_end = load_best_model_at_end,325            metric_for_best_model = metric_for_best_model,326            greater_is_better = greater_is_better,327            ignore_data_skip = ignore_data_skip,328            fsdp = fsdp,329            fsdp_min_num_params = fsdp_min_num_params,330            fsdp_config = fsdp_config,331            tp_size = tp_size,332            fsdp_transformer_layer_cls_to_wrap = fsdp_transformer_layer_cls_to_wrap,333            accelerator_config = accelerator_config,334            deepspeed = deepspeed,335            label_smoothing_factor = label_smoothing_factor,336            optim = optim,337            optim_args = optim_args,338            adafactor = adafactor,339            group_by_length = group_by_length,340            length_column_name = length_column_name,341            report_to = report_to,342            ddp_find_unused_parameters = ddp_find_unused_parameters,343            ddp_bucket_cap_mb = ddp_bucket_cap_mb,344            ddp_broadcast_buffers = ddp_broadcast_buffers,345            dataloader_pin_memory = dataloader_pin_memory,346            dataloader_persistent_workers = dataloader_persistent_workers,347            skip_memory_metrics = skip_memory_metrics,348            use_legacy_prediction_loop = use_legacy_prediction_loop,349            push_to_hub = push_to_hub,350            resume_from_checkpoint = resume_from_checkpoint,351            hub_model_id = hub_model_id,352            hub_strategy = hub_strategy,353            hub_token = hub_token,354            hub_private_repo = hub_private_repo,355            hub_always_push = hub_always_push,356            gradient_checkpointing = gradient_checkpointing,357            gradient_checkpointing_kwargs = gradient_checkpointing_kwargs,358            include_inputs_for_metrics = include_inputs_for_metrics,359            eval_do_concat_batches = eval_do_concat_batches,360            fp16_backend = fp16_backend,361            evaluation_strategy = evaluation_strategy,362            push_to_hub_model_id = push_to_hub_model_id,363            push_to_hub_organization = push_to_hub_organization,364            push_to_hub_token = push_to_hub_token,365            mp_parameters = mp_parameters,366            auto_find_batch_size = auto_find_batch_size,367            full_determinism = full_determinism,368            torchdynamo = torchdynamo,369            ray_scope = ray_scope,370            ddp_timeout = ddp_timeout,371            torch_compile = torch_compile,372            torch_compile_backend = torch_compile_backend,373            torch_compile_mode = torch_compile_mode,374            dispatch_batches = dispatch_batches,375            split_batches = split_batches,376            include_tokens_per_second = include_tokens_per_second,377            include_num_input_tokens_seen = include_num_input_tokens_seen,378            neftune_noise_alpha = neftune_noise_alpha,379            optim_target_modules = optim_target_modules,380            batch_eval_metrics = batch_eval_metrics,381            eval_on_start = eval_on_start,382            use_liger_kernel = use_liger_kernel,383            eval_use_gather_object = eval_use_gather_object,384            average_tokens_across_devices = average_tokens_across_devices,385            model_init_kwargs = model_init_kwargs,386            use_liger = use_liger,387            dataset_text_field = dataset_text_field,388            dataset_kwargs = dataset_kwargs,389            dataset_num_proc = dataset_num_proc,390            max_seq_length = max_seq_length,391            packing = packing,392            eval_packing = eval_packing,393            dataset_batch_size = dataset_batch_size,394            num_of_sequences = num_of_sequences,395            chars_per_token = chars_per_token,**kwargs)396        self.vllm_sampling_params = vllm_sampling_params397        self.unsloth_num_chunks = unsloth_num_chunks398pass399 400class _UnslothSFTTrainer(Trainer):401    """"""402 403    _tag_names = ["trl", "sft"]404 405    @deprecate_kwarg(406        "tokenizer", "0.16.0", "processing_class", warn_if_greater_or_equal_version=True, raise_if_both_names=True407    )408    def __init__(409        self,410        model: Union[str, nn.Module, PreTrainedModel],411        args: Optional[Union[SFTConfig, TrainingArguments]] = None,412        data_collator: Optional[DataCollator] = None,  # type: ignore413        train_dataset: Optional[Union[Dataset, IterableDataset]] = None,414        eval_dataset: Optional[Union[Dataset, dict[str, Dataset]]] = None,415        processing_class: Optional[416            Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin]417        ] = None,418        compute_loss_func: Optional[Callable] = None,419        compute_metrics: Optional[Callable[[EvalPrediction], dict]] = None,420        callbacks: Optional[list[TrainerCallback]] = None,421        optimizers: tuple[Optional[torch.optim.Optimizer], Optional[torch.optim.lr_scheduler.LambdaLR]] = (None, None),422        optimizer_cls_and_kwargs: Optional[tuple[Type[torch.optim.Optimizer], dict[str, Any]]] = None,423        preprocess_logits_for_metrics: Optional[Callable[[torch.Tensor, torch.Tensor], torch.Tensor]] = None,424        peft_config: Optional["PeftConfig"] = None,425        formatting_func: Optional[Union[Callable[[dict], str], Callable[[dict], list[str]]]] = None,426    ):427        # Args428        if args is None:429            model_name = model if isinstance(model, str) else model.config._name_or_path430            model_name = model_name.split("/")[-1]431            args = SFTConfig(f"{model_name}-SFT")432        elif isinstance(args, TrainingArguments) and not isinstance(args, SFTConfig):433            dict_args = args.to_dict()434            dict_args["hub_token"] = args.hub_token  # to_dict hides the hub_token435            dict_args.pop("push_to_hub_token")436            args = SFTConfig(**dict_args)437 438        # Model439        if args.model_init_kwargs is not None and not isinstance(model, str):440            warnings.warn(441                "You passed model_init_kwargs to the `SFTConfig`, but your model is already instantiated. "442                "The `model_init_kwargs` will be ignored."443            )444        if isinstance(model, str):445            model = self._create_model_from_path(model, args)446 447        # PEFT configuration and model wrapping448        if False:449            model = self._prepare_peft_model(model, peft_config, args)450 451        # Handle the tokenizer452        if processing_class is None:453            processing_class = AutoTokenizer.from_pretrained(model.config._name_or_path)454            if processing_class.pad_token is None:455                processing_class.pad_token = processing_class.eos_token  # required for padding when collating data456 457        # Dataset458        preprocess_dataset = args.dataset_kwargs is None or not args.dataset_kwargs.get("skip_prepare_dataset", False)459        if preprocess_dataset:460            train_dataset = self._prepare_dataset(461                train_dataset, processing_class, args, args.packing, formatting_func, "train"462            )463            if eval_dataset is not None:464                packing = args.packing if args.eval_packing is None else args.eval_packing465                if isinstance(eval_dataset, dict):466                    eval_dataset = {467                        key: self._prepare_dataset(dataset, processing_class, args, packing, formatting_func, key)468                        for key, dataset in eval_dataset.items()469                    }470                else:471                    eval_dataset = self._prepare_dataset(472                        eval_dataset, processing_class, args, packing, formatting_func, "eval"473                    )474 475        # Data collator476        if data_collator is None:477            data_collator = DataCollatorForLanguageModeling(tokenizer=processing_class, mlm=False)478 479        # Initialize the metrics480        self._metrics = defaultdict(list)481 482        # Initialize the Trainer. Parent class will handle:483        # - DeepSpeed configuration (through create_accelerator_and_postprocess)484        # - FSDP setup485        # - Distributed training setup486        # - Optimizer and scheduler creation487        # Some arguments are only available for transformers>=4.47.0. Can be removed when the min version is bumped.488        super_init_kwargs = {}489        if version.parse(transformers.__version__) >= version.parse("4.47.0.dev0"):490            super_init_kwargs["optimizer_cls_and_kwargs"] = optimizer_cls_and_kwargs491        else:492            if optimizer_cls_and_kwargs is not None:493                warnings.warn(494                    "The `optimizer_cls_and_kwargs` argument is only available for `transformers>=4.47.0`. "495                    "The default optimizer will be used. "496                    "Remove the `optimizer_cls_and_kwargs` or upgrade to `transformers>=4.47.0`."497                )498        super().__init__(499            model=model,500            args=args,501            data_collator=data_collator,502            train_dataset=train_dataset,503            eval_dataset=eval_dataset,504            processing_class=processing_class,505            compute_loss_func=compute_loss_func,506            compute_metrics=compute_metrics,507            callbacks=callbacks,508            optimizers=optimizers,509            preprocess_logits_for_metrics=preprocess_logits_for_metrics,510            **super_init_kwargs,511        )512 513        # Add tags for models that have been loaded with the correct transformers version514        if hasattr(self.model, "add_model_tags"):515            self.model.add_model_tags(self._tag_names)516 517    def _create_model_from_path(self, model_path: str, args: SFTConfig) -> PreTrainedModel:518        """Creates a model from a path or model identifier."""519        model_init_kwargs = args.model_init_kwargs or {}520        # Handle torch dtype521        torch_dtype = model_init_kwargs.get("torch_dtype")522        if isinstance(torch_dtype, torch.dtype) or torch_dtype == "auto" or torch_dtype is None:523            pass  # torch_dtype is already a torch.dtype or "auto" or None524        elif isinstance(torch_dtype, str):  # it's a str, but not "auto"525            torch_dtype = getattr(torch, torch_dtype)526            model_init_kwargs["torch_dtype"] = torch_dtype527        else:528            raise ValueError(529                "Invalid `torch_dtype` passed to `SFTConfig`. Expected either 'auto' or a string representing "530                f"a `torch.dtype` (e.g., 'float32'), but got {torch_dtype}."531            )532        # Disable caching if gradient checkpointing is enabled (not supported)533        if args.gradient_checkpointing:534            model_init_kwargs["use_cache"] = False535 536        # Create model537        if args.use_liger:538            if not is_liger_kernel_available():539                raise ImportError("Please install Liger-kernel for use_liger=True")540            model = AutoLigerKernelForCausalLM.from_pretrained(model_path, **model_init_kwargs)541        else:542            model = AutoModelForCausalLM.from_pretrained(model_path, **model_init_kwargs)543        return model544 545    def _prepare_peft_model(self, model: PreTrainedModel, peft_config: Any, args: SFTConfig) -> PreTrainedModel:546        """Prepares a model for PEFT training."""547        if not is_peft_available():548            raise ImportError("To use PeftModel, you need to install the `peft` library.")549 550        if not isinstance(peft_config, PeftConfig):551            raise ValueError(552                f"Expected PeftConfig object but got {type(peft_config)}. If you want to use the PeftModel, you need "553                "to pass a PeftConfig object to the SFTTrainer."554            )555 556        if isinstance(model, PeftModel):557            return model558 559        # Handle quantized models (QLoRA)560        is_qlora = getattr(model, "is_loaded_in_4bit", False) or getattr(model, "is_loaded_in_8bit", False)561 562        is_sharded_qlora = False563        if getattr(model, "is_loaded_in_4bit", False):564            # Check if model is sharded (FSDP/DS-Zero3)565            for _, param in model.named_parameters():566                if param.__class__.__name__ == "Params4bit":567                    is_sharded_qlora = param.data.device.type in {"cpu", "meta"}568                    break569 570        # Prepare model for kbit training if needed571        if is_qlora and not is_sharded_qlora:572            model = self._prepare_model_for_kbit_training(model, args)573            # Disable gradient checkpointing as it's handled by prepare_model_for_kbit_training574            args = dataclasses.replace(args, gradient_checkpointing=False)575        elif args.gradient_checkpointing:576            model = self._enable_gradient_checkpointing(model, args)577 578        # Create PEFT model579        if (580            version.parse(peft.__version__) >= version.parse("0.12")  # autocast_adapter_dtype introduced in 0.12581            and getattr(model, "is_loaded_in_4bit", False)582            and is_sharded_qlora583        ):584            model = get_peft_model(model, peft_config, autocast_adapter_dtype=False)585        else:586            model = get_peft_model(model, peft_config)587 588        # Handle bf16 casting for 4-bit models589        if args.bf16 and getattr(model, "is_loaded_in_4bit", False) and not is_sharded_qlora:590            peft_module_casting_to_bf16(model)591 592        return model593 594    def _prepare_model_for_kbit_training(self, model: PreTrainedModel, args: SFTConfig) -> PreTrainedModel:595        """Prepares a quantized model for kbit training."""596        prepare_model_kwargs = {597            "use_gradient_checkpointing": args.gradient_checkpointing,598            "gradient_checkpointing_kwargs": args.gradient_checkpointing_kwargs or {},599        }600 601        return prepare_model_for_kbit_training(model, **prepare_model_kwargs)602 603    def _enable_gradient_checkpointing(self, model: PreTrainedModel, args: SFTConfig) -> PreTrainedModel:604        """Enables gradient checkpointing for the model."""605        gradient_checkpointing_kwargs = args.gradient_checkpointing_kwargs or {}606        use_reentrant = (607            "use_reentrant" not in gradient_checkpointing_kwargs or gradient_checkpointing_kwargs["use_reentrant"]608        )609 610        if use_reentrant:611            if hasattr(model, "enable_input_require_grads"):612                model.enable_input_require_grads()613            else:614 615                def make_inputs_require_grad(module, input, output):616                    output.requires_grad_(True)617 618                model.get_input_embeddings().register_forward_hook(make_inputs_require_grad)619 620        return model621 622    def _prepare_dataset(623        self,624        dataset: Union[Dataset, IterableDataset],625        processing_class,626        args,627        packing: bool,628        formatting_func: Optional[Callable[[dict], str]],629        dataset_name: str,630    ) -> Union[Dataset, IterableDataset]:631        # All Unsloth Zoo code licensed under LGPLv3632        if isinstance(dataset, ConstantLengthDataset): return dataset633    634        map_kwargs = {}635        use_desc = isinstance(dataset, Dataset)636        is_vlm = hasattr(processing_class, "tokenizer")637        tokenizer = processing_class638        if is_vlm: tokenizer = processing_class.tokenizer639    640        # Get max length641        max_seq_length = getattr(args, "max_length", 0)642        if max_seq_length == 0: max_seq_length = getattr(args, "max_seq_length", 0)643        if max_seq_length == 0: max_seq_length = getattr(self, "max_seq_length", 0)644        if max_seq_length == 0: max_seq_length = getattr(self, "max_seq", 0)645        if max_seq_length == 0: raise RuntimeError("Unsloth: max_seq_length is 0! Please specify one!")646        dataset_text_field = getattr(args, "dataset_text_field", "text")647        do_truncation = max_seq_length != 0648        do_formatting_func = False649        do_tokenize = True650    651        # Get correct column names652        column_names = set(next(iter(dataset)).keys())653        used_column_names = ["input_ids"]654        if "attention_mask" in column_names:655            used_column_names.append("attention_mask")656    657        # Check if already tokenized so skip658        from transformers import DataCollatorForSeq2Seq, DataCollatorForLanguageModeling659        if "labels" in column_names:660            # Most likely forgot data collator!661            if is_vlm and not hasattr(tokenizer, "pad"):662                # Check if processing_class has a .pad, if not, use tokenizer.tokenizer663                raise RuntimeError(f"Unsloth: {processing_class.__class__} does not have .pad!")664            self.data_collator = DataCollatorForSeq2Seq(tokenizer)665            used_column_names.append("labels")666            do_tokenize = False667        elif "input_ids" in column_names:668            # Skip dataset prep, and set data collator669            if is_vlm and not hasattr(tokenizer, "pad"):670                # Check if processing_class has a .pad, if not, use tokenizer.tokenizer671                raise RuntimeError(f"Unsloth: {processing_class.__class__} does not have .pad!")672            self.data_collator = DataCollatorForLanguageModeling(tokenizer, mlm = False)673            do_tokenize = False674        elif dataset_text_field not in column_names:675            do_formatting_func = True676            if formatting_func is None:677                raise RuntimeError("Unsloth: You must specify a `formatting_func`")678        pass679    680        if do_tokenize:681            # Check double BOS tokens682            if do_formatting_func:683                test_text = formatting_func(dataset[0])684                if not isinstance(test_text, list):685                    raise ValueError(686                        "Unsloth: The `formatting_func` should return a list of processed strings."687                    )688                test_text = test_text[0]689            else:690                test_text = dataset[0][dataset_text_field]691    692            # Get chat template693            chat_template = getattr(processing_class, 'chat_template', '')694            if chat_template == '' and is_vlm:695                chat_template = getattr(tokenizer, 'chat_template', '')696            if chat_template is None:697                chat_template = ''698    699            # Get bos_token700            add_special_tokens = True701            bos_token_1 = getattr(processing_class, 'bos_token', None)702            bos_token_2 = getattr(tokenizer, 'bos_token', None)703            bos_token = bos_token_1 or bos_token_2704    705            if bos_token is not None:706                if test_text.startswith(bos_token) or bos_token in chat_template:707                    add_special_tokens = False708                    print("Unsloth: We found double BOS tokens - we shall remove one automatically.")709            pass710    711            # Create tokenize function712            def _tokenize(example):713                return tokenizer(714                    example[dataset_text_field] if not do_formatting_func else formatting_func(example),715                    truncation = do_truncation,716                    max_length = max_seq_length,717                    return_token_type_ids = False,718                    add_special_tokens = add_special_tokens,719                )720            pass721    722            map_kwargs["num_proc"] = getattr(args, "dataset_num_proc", 2)723            if use_desc: map_kwargs["desc"] = f'Unsloth: Tokenizing ["{dataset_text_field}"]'724            dataset = dataset.map(_tokenize, batched = True, **map_kwargs)725    726            # If VLM, switch data collator since .pad is needed!727            if is_vlm and not hasattr(processing_class, "pad"):728                data_collator = DataCollatorForLanguageModeling(tokenizer, mlm = False)729                self.data_collator = data_collator730            pass731        pass732        if packing:733            print("Unsloth: Hugging Face's packing is currently buggy - we're disabling it for now!")734            return dataset735    736            if max_seq_length == 0:737                raise ValueError("When packing is enabled, `max_seq_length` can't be `None`.")738    739            if use_desc: map_kwargs["desc"] = f"Unsloth: Packing {dataset_name} dataset"740            dataset = dataset.select_columns(used_column_names).map(741                pack_examples,742                batched = True,743                fn_kwargs = {"seq_length": max_seq_length,},744                **map_kwargs,745            )746        pass747        return dataset748    749    def compute_loss(self, model, inputs, return_outputs = False, num_items_in_batch = None):750        outputs = super().compute_loss(751            model,752            inputs,753            return_outputs = return_outputs,754            num_items_in_batch = num_items_in_batch,755        )756        return outputs757 758    def log(self, logs: dict[str, float], start_time: Optional[float] = None) -> None:759        metrics = {key: sum(val) / len(val) for key, val in self._metrics.items()}  # average the metrics760 761        # This method can be called both in training and evaluation. When called in evaluation, the keys in `logs`762        # start with "eval_". We need to add the prefix "eval_" to the keys in `metrics` to match the format.763        if next(iter(logs.keys())).startswith("eval_"):764            metrics = {f"eval_{key}": val for key, val in metrics.items()}765 766        logs = {**logs, **metrics}767        if version.parse(transformers.__version__) >= version.parse("4.47.0.dev0"):768            super().log(logs, start_time)769        else:  # transformers<=4.46770            super().log(logs)771        self._metrics.clear()772 773    def create_model_card(774        self,775        model_name: Optional[str] = None,776        dataset_name: Optional[str] = None,777        tags: Union[str, list[str], None] = None,778    ):779        """780        Creates a draft of a model card using the information available to the `Trainer`.781 782        Args:783            model_name (`str` or `None`, *optional*, defaults to `None`):784                Name of the model.785            dataset_name (`str` or `None`, *optional*, defaults to `None`):786                Name of the dataset used for training.787            tags (`str`, `list[str]` or `None`, *optional*, defaults to `None`):788                Tags to be associated with the model card.789        """790        if not self.is_world_process_zero():791            return792 793        if hasattr(self.model.config, "_name_or_path") and not os.path.isdir(self.model.config._name_or_path):794            base_model = self.model.config._name_or_path795        else:796            base_model = None797 798        tags = tags or []799        if isinstance(tags, str):800            tags = [tags]801 802        if hasattr(self.model.config, "unsloth_version"):803            tags.append("unsloth")804 805        model_card = generate_model_card(806            base_model=base_model,807            model_name=model_name,808            hub_model_id=self.hub_model_id,809            dataset_name=dataset_name,810            tags=tags,811            wandb_url=wandb.run.get_url() if is_wandb_available() and wandb.run is not None else None,812            comet_url=get_comet_experiment_url(),813            trainer_name="SFT",814        )815 816        model_card.save(os.path.join(self.args.output_dir, "README.md"))817class UnslothSFTTrainer(_UnslothSFTTrainer):818    """819    820    Trainer for Supervised Fine-Tuning (SFT) method.821 822    This class is a wrapper around the [`transformers.Trainer`] class and inherits all of its attributes and methods.823 824    Example:825 826    ```python827    from datasets import load_dataset828    from trl import SFTTrainer829 830    dataset = load_dataset("roneneldan/TinyStories", split="train[:1%]")831 832    trainer = SFTTrainer(model="Qwen/Qwen2-0.5B-Instruct", train_dataset=dataset)833    trainer.train()834    ```835 836    Args:837        model (`Union[str, PreTrainedModel]`):838            Model to be trained. Can be either:839 840            - A string, being the *model id* of a pretrained model hosted inside a model repo on huggingface.co, or841              a path to a *directory* containing model weights saved using842              [`~transformers.PreTrainedModel.save_pretrained`], e.g., `'./my_model_directory/'`. The model is843              loaded using [`~transformers.AutoModelForCausalLM.from_pretrained`] with the keywork arguments844              in `args.model_init_kwargs`.845            - A [`~transformers.PreTrainedModel`] object. Only causal language models are supported.846        args ([`SFTConfig`], *optional*, defaults to `None`):847            Configuration for this trainer. If `None`, a default configuration is used.848        data_collator (`DataCollator`, *optional*):849            Function to use to form a batch from a list of elements of the prcessed `train_dataset` or `eval_dataset`.850            Will default to [`~transformers.default_data_collator`] if no `processing_class` is provided, an instance851            of [`~transformers.DataCollatorWithPadding`] otherwise if the processing_class is a feature extractor or852            tokenizer.853        train_dataset ([`~datasets.Dataset`] or [`~datasets.IterableDataset`]):854            Dataset to use for training. SFT supports both [language modeling](#language-modeling) type and855            [prompt-completion](#prompt-completion) type. The format of the samples can be either:856 857            - [Standard](dataset_formats#standard): Each sample contains plain text.858            - [Conversational](dataset_formats#conversational): Each sample contains structured messages (e.g., role859              and content).860 861            The trainer also supports processed datasets (tokenized) as long as they contain an `input_ids` field.862        eval_dataset ([`~datasets.Dataset`], [`~datasets.IterableDataset`] or `dict[str, Union[Dataset, IterableDataset]]`):863            Dataset to use for evaluation. It must meet the same requirements as `train_dataset`.864        processing_class ([`~transformers.PreTrainedTokenizerBase`], *optional*, defaults to `None`):865            Processing class used to process the data. If `None`, the processing class is loaded from the model's name866            with [`~transformers.AutoTokenizer.from_pretrained`].867        callbacks (list of [`~transformers.TrainerCallback`], *optional*, defaults to `None`):868            List of callbacks to customize the training loop. Will add those to the list of default callbacks869            detailed in [here](https://huggingface.co/docs/transformers/main_classes/callback).870 871            If you want to remove one of the default callbacks used, use the [`~transformers.Trainer.remove_callback`]872            method.873        optimizers (`tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR]`, *optional*, defaults to `(None, None)`):874            A tuple containing the optimizer and the scheduler to use. Will default to an instance of [`AdamW`] on your875            model and a scheduler given by [`get_linear_schedule_with_warmup`] controlled by `args`.876        optimizer_cls_and_kwargs (`Tuple[Type[torch.optim.Optimizer], Dict[str, Any]]`, *optional*, defaults to `None`):877            A tuple containing the optimizer class and keyword arguments to use.878            Overrides `optim` and `optim_args` in `args`. Incompatible with the `optimizers` argument.879 880            Unlike `optimizers`, this argument avoids the need to place model parameters on the correct devices before initializing the Trainer.881        preprocess_logits_for_metrics (`Callable[[torch.Tensor, torch.Tensor], torch.Tensor]`, *optional*, defaults to `None`):882            A function that preprocess the logits right before caching them at each evaluation step. Must take two883            tensors, the logits and the labels, and return the logits once processed as desired. The modifications made884            by this function will be reflected in the predictions received by `compute_metrics`.885 886            Note that the labels (second parameter) will be `None` if the dataset does not have them.887        peft_config ([`~peft.PeftConfig`], *optional*, defaults to `None`):888            PEFT configuration used to wrap the model. If `None`, the model is not wrapped.889        formatting_func (`Optional[Callable]`):890            Formatting function applied to the dataset before tokenization.891    892    """893    def __init__(894        self,895        model,896        args = None,897        data_collator = None,898        train_dataset = None,899        eval_dataset = None,900        processing_class = None,901        compute_loss_func = None,902        compute_metrics = None,903        callbacks = None,904        optimizer_cls_and_kwargs = None,905        preprocess_logits_for_metrics = None,906        peft_config = None,907        formatting_func = None,908        **kwargs909    ):910        if args is None: args = UnslothSFTConfig()911        use_bf16 = getattr(args, 'bf16', False)912        use_fp16 = getattr(args, 'fp16', False)913        force_float32 = False914        if os.environ.get('UNSLOTH_FORCE_FLOAT32', '0') == '1':915            print('Unsloth: Switching to float32 training since model cannot work with float16')916            force_float32 = True917        mixed_precision_dtype = os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32')918        dtype = getattr(model.config, 'torch_dtype', None)919        if dtype is None: dtype = model.get_input_embeddings().dtype920        from unsloth_zoo.utils import _get_dtype921        dtype = _get_dtype(dtype)922        float16 = dtype == torch.float16923        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`')924        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`')925        if force_float32:926            args.fp16 = False927            args.bf16 = False928            os.environ['ACCELERATE_MIXED_PRECISION'] = 'no'929        elif (not use_bf16 and not use_fp16) and mixed_precision_dtype == 'float32':930            args.fp16 = float16931            args.bf16 = not float16932            os.environ['ACCELERATE_MIXED_PRECISION'] = 'fp16' if float16 else 'bf16'933        if getattr(args, 'eval_dataset', None) is not None and getattr(args, 'eval_strategy', 'no') == 'no':934            args.eval_strategy = 'steps'935            if getattr(args, 'eval_steps', None) is None: args.eval_steps = 0.1936        ga_steps = getattr(args, 'gradient_accumulation_steps', None)937        if ga_steps is not None and ga_steps > 1:938            from transformers import __version__ as transformers_version939            if Version(transformers_version) <= Version('4.45.2'):940                print('**** Unsloth: Please use our fixed gradient_accumulation_steps by updating transformers, TRL and Unsloth!\n'941                      '`pip install --upgrade --no-cache-dir --force-reinstall --no-deps unsloth transformers trl unsloth_zoo`')942        if getattr(args, 'eval_strategy', 'no') != 'no':943            eval_bsz = getattr(args, 'per_device_eval_batch_size', 8)944            if eval_bsz == 8 and args.per_device_train_batch_size < eval_bsz: args.per_device_eval_batch_size = args.per_device_train_batch_size945            if getattr(args, 'eval_accumulation_steps', None) is None and ga_steps is not None: args.eval_accumulation_steps = ga_steps946        fp16_full_eval = getattr(args, 'fp16_full_eval', False)947        bf16_full_eval = getattr(args, 'bf16_full_eval', False)948        if args.fp16 and bf16_full_eval: args.bf16_full_eval = False; args.fp16_full_eval = True949        if args.bf16 and fp16_full_eval: args.bf16_full_eval = True; args.fp16_full_eval = False950        if force_float32:951            args.bf16_full_eval = False952            args.fp16_full_eval = False953        elif os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32') == 'bfloat16':954            args.bf16_full_eval = True955            args.fp16_full_eval = False956        elif not bf16_full_eval and not fp16_full_eval:957            args.bf16_full_eval = args.bf16958            args.fp16_full_eval = args.fp16959        _output_logits = False960        if locals().get('compute_metrics', None) is not None: _output_logits = True961        if locals().get('preprocess_logits_for_metrics', None) is not None: _output_logits = True962        if _output_logits:963            os.environ['UNSLOTH_RETURN_LOGITS'] = '1'964        if 'max_seq_length' not in locals() and not hasattr(args, 'max_seq_length'):965            pass966        else:967            model_max_seq_length = getattr(model, 'max_seq_length', None)968            args_max_seq_length  = getattr(args,  'max_seq_length', None)969            if args_max_seq_length is None and model_max_seq_length is not None:970                max_seq_length = model.max_seq_length971                if hasattr(args, 'max_seq_length'): args.max_seq_length = max_seq_length972        if model is not None and hasattr(model, 'for_training'):973            model.for_training()974        if 'tokenizer' in locals() and hasattr(tokenizer, 'padding_side'): tokenizer.padding_side = 'right'975        if 'processing_class' in locals():976            if hasattr(processing_class, 'padding_side'): processing_class.padding_side = 'right'977            if hasattr(processing_class, 'tokenizer') and hasattr(processing_class.tokenizer, 'padding_side'): processing_class.tokenizer.padding_side = 'right'978        __tokenizer = processing_class if 'processing_class' in locals() else tokenizer979        from unsloth_zoo.vision_utils import UnslothVisionDataCollator980        if not isinstance(data_collator, UnslothVisionDataCollator):981            if isinstance(data_collator, DataCollatorForSeq2Seq) and 'labels' not in train_dataset.column_names:982                data_collator = DataCollatorForLanguageModeling(__tokenizer, mlm = False)983            elif isinstance(data_collator, DataCollatorForLanguageModeling) and 'labels' in train_dataset.column_names:984                data_collator = DataCollatorForSeq2Seq(__tokenizer)985        else:986            if hasattr(args, 'remove_unused_columns'): args.remove_unused_columns = False987            if hasattr(args, 'dataset_text_field'): args.dataset_text_field = ''988            if hasattr(args, 'dataset_kwargs'): args.dataset_kwargs = {'skip_prepare_dataset': True}989        if not isinstance(data_collator, UnslothVisionDataCollator):990            if not hasattr(__tokenizer, 'pad') and hasattr(__tokenizer, 'tokenizer'):991                if isinstance(data_collator, DataCollatorForSeq2Seq):992                    data_collator = DataCollatorForSeq2Seq(__tokenizer.tokenizer)993                else:994                    data_collator = DataCollatorForLanguageModeling(__tokenizer.tokenizer, mlm = False)995        other_metrics = []996        997        from unsloth_zoo.logging_utils import PatchRLStatistics998        PatchRLStatistics('sft_trainer', other_metrics)999        IGNORED_TOKENIZER_NAMES = os.environ.get('UNSLOTH_IGNORED_TOKENIZER_NAMES', '').split('\n')1000        from unsloth_zoo.tokenizer_utils import fix_untrained_tokens1001        from unsloth_zoo.training_utils  import fix_zero_training_loss1002        if 'tokenizer' not in locals(): tokenizer = processing_class1003        fix_untrained_tokens(model, tokenizer, train_dataset, IGNORED_TOKENIZER_NAMES, eps = 1e-16)1004        fix_zero_training_loss(model, tokenizer, train_dataset)1005        1006        super().__init__(1007            model = model,1008            args = args,1009            data_collator = data_collator,1010            train_dataset = train_dataset,1011            eval_dataset = eval_dataset,1012            processing_class = processing_class,1013            compute_loss_func = compute_loss_func,1014            compute_metrics = compute_metrics,1015            callbacks = callbacks,1016            optimizer_cls_and_kwargs = optimizer_cls_and_kwargs,1017            preprocess_logits_for_metrics = preprocess_logits_for_metrics,1018            peft_config = peft_config,1019            formatting_func = formatting_func,**kwargs)1020        if hasattr(self, 'neftune_hook_handle'):1021            self.neftune_hook_handle.remove()1022            if hasattr(self, 'neftune_hook_handle'): del self.neftune_hook_handle1023        if getattr(args, 'neftune_noise_alpha', None) is not None:1024            model.get_input_embeddings().neftune_noise_alpha = self.neftune_noise_alpha1025        pass1026        1027pass1028