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