Zwounds/Boolean_Search_Query_Model
0
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 