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.xpo_trainer import (Any, BaseImageProcessor, BasePairwiseJudge, Callable, Dataset, EvalPrediction, F, FeatureExtractionMixin, IterableDataset, OnlineDPOTrainer, OptimizerNames, Optional, PreTrainedModel, PreTrainedTokenizerBase, ProcessorMixin, SIMPLE_CHAT_TEMPLATE, TrainerCallback, Union, XPOConfig, XPOTrainer, empty_cache, generate_model_card, get_comet_experiment_url, get_reward, is_conversational, is_wandb_available, jinja2, maybe_apply_chat_template, nn, os, textwrap, torch, truncate_right, unwrap_model_for_generation)13 14 15import os16from typing import *17from dataclasses import dataclass, field18from packaging.version import Version19import torch20import numpy as np21from contextlib import nullcontext22from torch.nn import functional as F23from transformers import DataCollatorForSeq2Seq, DataCollatorForLanguageModeling24 25torch_compile_options = {26 "epilogue_fusion" : True,27 "max_autotune" : False,28 "shape_padding" : True,29 "trace.enabled" : False,30 "triton.cudagraphs" : False,31}32 33@torch.compile(dynamic = True, fullgraph = True, options = torch_compile_options,)34def selective_log_softmax(logits, index):35 logits = logits.to(torch.float32)36 selected_logits = torch.gather(logits, dim = -1, index = index.unsqueeze(-1)).squeeze(-1)37 # loop to reduce peak mem consumption38 # logsumexp_values = torch.stack([torch.logsumexp(lg, dim=-1) for lg in logits])39 logsumexp_values = torch.logsumexp(logits, dim = -1)40 per_token_logps = selected_logits - logsumexp_values # log_softmax(x_i) = x_i - logsumexp(x)41 return per_token_logps42@dataclass43class UnslothXPOConfig(XPOConfig):44 """45 46 Configuration class for the [`XPOTrainer`].47 48 Subclass of [`OnlineDPOConfig`] we can use all its arguments and add the following:49 50 Parameters:51 alpha (`float` or `list[float]`, *optional*, defaults to `1e-5`):52 Weight of the XPO loss term. If a list of floats is provided then the alpha is selected for each new epoch53 and the last alpha is used for the rest of the epochs.54 55 """56 vllm_sampling_params: Optional[Any] = field(57 default = None,58 metadata = {'help': 'vLLM SamplingParams'},59 )60 unsloth_num_chunks : Optional[int] = field(61 default = -1,62 metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},63 )64 def __init__(65 self,66 output_dir = None,67 overwrite_output_dir = None,68 do_train = False,69 do_eval = False,70 do_predict = False,71 eval_strategy = 'no',72 prediction_loss_only = False,73 per_device_train_batch_size = 4,74 per_device_eval_batch_size = 4,75 per_gpu_train_batch_size = None,76 per_gpu_eval_batch_size = None,77 gradient_accumulation_steps = 2,78 eval_accumulation_steps = 2,79 eval_delay = 0,80 torch_empty_cache_steps = 250,81 learning_rate = 5e-05,82 weight_decay = 0.01,83 adam_beta1 = 0.9,84 adam_beta2 = 0.999,85 adam_epsilon = 1e-08,86 max_grad_norm = 1.0,87 num_train_epochs = 3.0,88 max_steps = -1,89 lr_scheduler_type = 'linear',90 warmup_ratio = 0.1,91 warmup_steps = 0,92 log_level = 'passive',93 log_level_replica = 'warning',94 log_on_each_node = True,95 logging_dir = None,96 logging_strategy = 'steps',97 logging_first_step = False,98 logging_steps = 1,99 logging_nan_inf_filter = False,100 save_strategy = 'steps',101 save_steps = 500,102 save_total_limit = None,103 save_safetensors = True,104 save_on_each_node = False,105 save_only_model = False,106 restore_callback_states_from_checkpoint = False,107 no_cuda = False,108 use_cpu = False,109 use_mps_device = False,110 seed = 3407,111 data_seed = 3407,112 jit_mode_eval = False,113 use_ipex = False,114 bf16 = False,115 fp16 = False,116 fp16_opt_level = 'O1',117 half_precision_backend = 'auto',118 bf16_full_eval = False,119 fp16_full_eval = False,120 tf32 = None,121 local_rank = -1,122 ddp_backend = None,123 tpu_num_cores = None,124 tpu_metrics_debug = False,125 debug = '',126 dataloader_drop_last = False,127 eval_steps = None,128 dataloader_num_workers = 0,129 dataloader_prefetch_factor = None,130 past_index = -1,131 run_name = None,132 disable_tqdm = None,133 remove_unused_columns = True,134 label_names = None,135 load_best_model_at_end = False,136 metric_for_best_model = None,137 greater_is_better = None,138 ignore_data_skip = False,139 fsdp = '',140 fsdp_min_num_params = 0,141 fsdp_config = None,142 tp_size = 0,143 fsdp_transformer_layer_cls_to_wrap = None,144 accelerator_config = None,145 deepspeed = None,146 label_smoothing_factor = 0.0,147 optim = 'adamw_8bit',148 optim_args = None,149 adafactor = False,150 group_by_length = False,151 length_column_name = 'length',152 report_to = None,153 ddp_find_unused_parameters = None,154 ddp_bucket_cap_mb = None,155 ddp_broadcast_buffers = None,156 dataloader_pin_memory = True,157 dataloader_persistent_workers = False,158 skip_memory_metrics = True,159 use_legacy_prediction_loop = False,160 push_to_hub = False,161 resume_from_checkpoint = None,162 hub_model_id = None,163 hub_strategy = 'every_save',164 hub_token = None,165 hub_private_repo = None,166 hub_always_push = False,167 gradient_checkpointing = False,168 gradient_checkpointing_kwargs = None,169 include_inputs_for_metrics = False,170 eval_do_concat_batches = True,171 fp16_backend = 'auto',172 evaluation_strategy = None,173 push_to_hub_model_id = None,174 push_to_hub_organization = None,175 push_to_hub_token = None,176 mp_parameters = '',177 auto_find_batch_size = False,178 full_determinism = False,179 torchdynamo = None,180 ray_scope = 'last',181 ddp_timeout = 1800,182 torch_compile = False,183 torch_compile_backend = None,184 torch_compile_mode = None,185 dispatch_batches = None,186 split_batches = None,187 include_tokens_per_second = False,188 include_num_input_tokens_seen = False,189 neftune_noise_alpha = None,190 optim_target_modules = None,191 batch_eval_metrics = False,192 eval_on_start = False,193 use_liger_kernel = False,194 eval_use_gather_object = False,195 average_tokens_across_devices = False,196 reward_model_path = None,197 judge = None,198 max_new_tokens = 64,199 max_length = 512,200 temperature = 0.9,201 missing_eos_penalty = None,202 loss_type = 'sigmoid',203 dataset_num_proc = None,204 disable_dropout = True,205 use_vllm = False,206 ds3_gather_for_generation = True,207 vllm_sampling_params = None,208 unsloth_num_chunks = -1,209 **kwargs,210 ):211 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!')212 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!')213 if output_dir is None and save_strategy == 'steps' and save_steps == 500:214 output_dir = 'unsloth_training_checkpoints'215 save_strategy = 'no'216 if dataset_num_proc is None:217 from multiprocessing import cpu_count218 dataset_num_proc = cpu_count()219 220 super().__init__(221 output_dir = output_dir,222 overwrite_output_dir = overwrite_output_dir,223 do_train = do_train,224 do_eval = do_eval,225 do_predict = do_predict,226 eval_strategy = eval_strategy,227 prediction_loss_only = prediction_loss_only,228 per_device_train_batch_size = per_device_train_batch_size,229 per_device_eval_batch_size = per_device_eval_batch_size,230 per_gpu_train_batch_size = per_gpu_train_batch_size,231 per_gpu_eval_batch_size = per_gpu_eval_batch_size,232 gradient_accumulation_steps = gradient_accumulation_steps,233 eval_accumulation_steps = eval_accumulation_steps,234 eval_delay = eval_delay,235 torch_empty_cache_steps = torch_empty_cache_steps,236 learning_rate = learning_rate,237 weight_decay = weight_decay,238 adam_beta1 = adam_beta1,239 adam_beta2 = adam_beta2,240 adam_epsilon = adam_epsilon,241 max_grad_norm = max_grad_norm,242 num_train_epochs = num_train_epochs,243 max_steps = max_steps,244 lr_scheduler_type = lr_scheduler_type,245 warmup_ratio = warmup_ratio,246 warmup_steps = warmup_steps,247 log_level = log_level,248 log_level_replica = log_level_replica,249 log_on_each_node = log_on_each_node,250 logging_dir = logging_dir,251 logging_strategy = logging_strategy,252 logging_first_step = logging_first_step,253 logging_steps = logging_steps,254 logging_nan_inf_filter = logging_nan_inf_filter,255 save_strategy = save_strategy,256 save_steps = save_steps,257 save_total_limit = save_total_limit,258 save_safetensors = save_safetensors,259 save_on_each_node = save_on_each_node,260 save_only_model = save_only_model,261 restore_callback_states_from_checkpoint = restore_callback_states_from_checkpoint,262 no_cuda = no_cuda,263 use_cpu = use_cpu,264 use_mps_device = use_mps_device,265 seed = seed,266 data_seed = data_seed,267 jit_mode_eval = jit_mode_eval,268 use_ipex = use_ipex,269 bf16 = bf16,270 fp16 = fp16,271 fp16_opt_level = fp16_opt_level,272 half_precision_backend = half_precision_backend,273 bf16_full_eval = bf16_full_eval,274 fp16_full_eval = fp16_full_eval,275 tf32 = tf32,276 local_rank = local_rank,277 ddp_backend = ddp_backend,278 tpu_num_cores = tpu_num_cores,279 tpu_metrics_debug = tpu_metrics_debug,280 debug = debug,281 dataloader_drop_last = dataloader_drop_last,282 eval_steps = eval_steps,283 dataloader_num_workers = dataloader_num_workers,284 dataloader_prefetch_factor = dataloader_prefetch_factor,285 past_index = past_index,286 run_name = run_name,287 disable_tqdm = disable_tqdm,288 remove_unused_columns = remove_unused_columns,289 label_names = label_names,290 load_best_model_at_end = load_best_model_at_end,291 metric_for_best_model = metric_for_best_model,292 greater_is_better = greater_is_better,293 ignore_data_skip = ignore_data_skip,294 fsdp = fsdp,295 fsdp_min_num_params = fsdp_min_num_params,296 fsdp_config = fsdp_config,297 tp_size = tp_size,298 fsdp_transformer_layer_cls_to_wrap = fsdp_transformer_layer_cls_to_wrap,299 accelerator_config = accelerator_config,300 deepspeed = deepspeed,301 label_smoothing_factor = label_smoothing_factor,302 optim = optim,303 optim_args = optim_args,304 adafactor = adafactor,305 group_by_length = group_by_length,306 length_column_name = length_column_name,307 report_to = report_to,308 ddp_find_unused_parameters = ddp_find_unused_parameters,309 ddp_bucket_cap_mb = ddp_bucket_cap_mb,310 ddp_broadcast_buffers = ddp_broadcast_buffers,311 dataloader_pin_memory = dataloader_pin_memory,312 dataloader_persistent_workers = dataloader_persistent_workers,313 skip_memory_metrics = skip_memory_metrics,314 use_legacy_prediction_loop = use_legacy_prediction_loop,315 push_to_hub = push_to_hub,316 resume_from_checkpoint = resume_from_checkpoint,317 hub_model_id = hub_model_id,318 hub_strategy = hub_strategy,319 hub_token = hub_token,320 hub_private_repo = hub_private_repo,321 hub_always_push = hub_always_push,322 gradient_checkpointing = gradient_checkpointing,323 gradient_checkpointing_kwargs = gradient_checkpointing_kwargs,324 include_inputs_for_metrics = include_inputs_for_metrics,325 eval_do_concat_batches = eval_do_concat_batches,326 fp16_backend = fp16_backend,327 evaluation_strategy = evaluation_strategy,328 push_to_hub_model_id = push_to_hub_model_id,329 push_to_hub_organization = push_to_hub_organization,330 push_to_hub_token = push_to_hub_token,331 mp_parameters = mp_parameters,332 auto_find_batch_size = auto_find_batch_size,333 full_determinism = full_determinism,334 torchdynamo = torchdynamo,335 ray_scope = ray_scope,336 ddp_timeout = ddp_timeout,337 torch_compile = torch_compile,338 torch_compile_backend = torch_compile_backend,339 torch_compile_mode = torch_compile_mode,340 dispatch_batches = dispatch_batches,341 split_batches = split_batches,342 include_tokens_per_second = include_tokens_per_second,343 include_num_input_tokens_seen = include_num_input_tokens_seen,344 neftune_noise_alpha = neftune_noise_alpha,345 optim_target_modules = optim_target_modules,346 batch_eval_metrics = batch_eval_metrics,347 eval_on_start = eval_on_start,348 use_liger_kernel = use_liger_kernel,349 eval_use_gather_object = eval_use_gather_object,350 average_tokens_across_devices = average_tokens_across_devices,351 reward_model_path = reward_model_path,352 judge = judge,353 max_new_tokens = max_new_tokens,354 max_length = max_length,355 temperature = temperature,356 missing_eos_penalty = missing_eos_penalty,357 loss_type = loss_type,358 dataset_num_proc = dataset_num_proc,359 disable_dropout = disable_dropout,360 use_vllm = use_vllm,361 ds3_gather_for_generation = ds3_gather_for_generation,**kwargs)362 self.vllm_sampling_params = vllm_sampling_params363 self.unsloth_num_chunks = unsloth_num_chunks364pass365 366class _UnslothXPOTrainer(OnlineDPOTrainer):367 r""""""368 369 _tag_names = ["trl", "xpo"]370 371 def __init__(372 self,373 model: Union[PreTrainedModel, nn.Module] = None,374 ref_model: Union[PreTrainedModel, nn.Module] = None,375 reward_model: Optional[nn.Module] = None,376 judge: Optional[BasePairwiseJudge] = None,377 args: Optional[XPOConfig] = None,378 data_collator: Optional[Callable] = None,379 train_dataset: Optional[Union[Dataset, IterableDataset]] = None,380 eval_dataset: Optional[Union[Dataset, dict[str, Dataset]]] = None,381 processing_class: Optional[382 Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin]383 ] = None,384 peft_config: Optional[dict] = None,385 compute_metrics: Optional[Callable[[EvalPrediction], dict]] = None,386 callbacks: Optional[list[TrainerCallback]] = None,387 optimizers: tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),388 preprocess_logits_for_metrics: Optional[Callable[[torch.Tensor, torch.Tensor], torch.Tensor]] = None,389 ) -> None:390 super().__init__(391 model=model,392 ref_model=ref_model,393 judge=judge,394 reward_model=reward_model,395 args=args,396 data_collator=data_collator,397 train_dataset=train_dataset,398 eval_dataset=eval_dataset,399 processing_class=processing_class,400 reward_processing_class=processing_class, # for now, XPOTrainer can't use any reward model401 peft_config=peft_config,402 compute_metrics=compute_metrics,403 callbacks=callbacks,404 optimizers=optimizers,405 preprocess_logits_for_metrics=preprocess_logits_for_metrics,406 )407 408 self._alpha = self.args.alpha409 410 # Overwrite the stats dictionary to include XPO specific statistics411 self.stats = {412 # Remove "non_score_reward", "rlhf_reward", "scores"413 # Add "loss/dpo", "loss/xpo"414 "loss/dpo": [],415 "loss/xpo": [],416 "objective/kl": [],417 "objective/entropy": [],418 "rewards/chosen": [],419 "rewards/rejected": [],420 "rewards/accuracies": [],421 "rewards/margins": [],422 "logps/chosen": [],423 "logps/rejected": [],424 # Replace "contain_eos_token" by "model_contain_eos_token" and "ref_contain_eos_token"425 "val/model_contain_eos_token": [],426 "val/ref_contain_eos_token": [],427 "alpha": [],428 "beta": [],429 }430 if self.reward_model is not None:431 # Replace "scores" by "model_scores" and "ref_scores"432 self.stats["objective/model_scores"] = []433 self.stats["objective/ref_scores"] = []434 self.stats["objective/scores_margin"] = []435 436 @property437 def alpha(self):438 if isinstance(self._alpha, list):439 epoch = self.state.epoch440 return self._alpha[epoch] if epoch < len(self._alpha) else self._alpha[-1]441 else:442 return self._alpha443 444 def _generate_completions(self, prompts, model):445 with unwrap_model_for_generation(model, self.accelerator) as unwrapped_model:446 model_output = unwrapped_model.generate(447 input_ids=prompts["input_ids"],448 attention_mask=prompts["attention_mask"],449 generation_config=self.generation_config,450 )451 452 ref_model = model if self.ref_model is None else self.ref_model453 with torch.no_grad(), unwrap_model_for_generation(ref_model, self.accelerator) as unwrapped_ref_model:454 ref_output = unwrapped_ref_model.generate(455 input_ids=prompts["input_ids"],456 attention_mask=prompts["attention_mask"],457 generation_config=self.generation_config,458 )459 460 return model_output, ref_output461 462 def _process_completions(self, model_output, ref_output, prompts):463 context_length = prompts["input_ids"].shape[1]464 465 # Process model completions466 model_completion_ids = model_output[:, context_length:]467 model_completion_ids, model_completion_mask = truncate_right(468 model_completion_ids, self.processing_class.eos_token_id, self.processing_class.pad_token_id469 )470 model_data = {471 "input_ids": torch.cat((prompts["input_ids"], model_completion_ids), dim=1),472 "attention_mask": torch.cat((prompts["attention_mask"], model_completion_mask), dim=1),473 "raw": prompts["raw"],474 }475 476 # Process reference model completions477 ref_completion_ids = ref_output[:, context_length:]478 ref_completion_ids, ref_completion_mask = truncate_right(479 ref_completion_ids, self.processing_class.eos_token_id, self.processing_class.pad_token_id480 )481 ref_data = {482 "input_ids": torch.cat((prompts["input_ids"], ref_completion_ids), dim=1),483 "attention_mask": torch.cat((prompts["attention_mask"], ref_completion_mask), dim=1),484 "raw": prompts["raw"],485 }486 487 return model_data, ref_data488 489 def _compute_rewards(self, model_data, ref_data, context_length):490 with torch.no_grad():491 _, model_scores, _ = get_reward(492 self.reward_model, model_data["input_ids"], self.processing_class.pad_token_id, context_length493 )494 _, ref_scores, _ = get_reward(495 self.reward_model, ref_data["input_ids"], self.processing_class.pad_token_id, context_length496 )497 498 # Apply EOS penalty if needed499 if self.args.missing_eos_penalty is not None:500 model_contain_eos = torch.any(model_data["input_ids"] == self.processing_class.eos_token_id, dim=-1)501 ref_contain_eos = torch.any(ref_data["input_ids"] == self.processing_class.eos_token_id, dim=-1)502 model_scores[~model_contain_eos] -= self.args.missing_eos_penalty503 ref_scores[~ref_contain_eos] -= self.args.missing_eos_penalty504 505 return model_scores, ref_scores506 507 def _compute_judge(self, model_data, ref_data, context_length):508 prompts = model_data["raw"]509 model_data_completions = self.processing_class.batch_decode(510 model_data["input_ids"][:, context_length:], skip_special_tokens=True511 )512 model_data_completions = [completion.strip() for completion in model_data_completions]513 514 ref_data_completions = self.processing_class.batch_decode(515 ref_data["input_ids"][:, context_length:], skip_special_tokens=True516 )517 ref_data_completions = [completion.strip() for completion in ref_data_completions]518 519 if is_conversational({"prompt": prompts[0]}):520 model_data_completions = [521 [{"role": "assistant", "content": completion}] for completion in model_data_completions522 ]523 environment = jinja2.Environment()524 template = environment.from_string(SIMPLE_CHAT_TEMPLATE)525 prompts = [template.render(messages=message) for message in prompts]526 model_data_completions = [template.render(messages=completion) for completion in model_data_completions]527 528 ref_data_completions = [529 [{"role": "assistant", "content": completion}] for completion in ref_data_completions530 ]531 ref_data_completions = [template.render(messages=completion) for completion in ref_data_completions]532 533 ranks_of_first_completion = self.judge.judge(534 prompts,535 list(zip(model_data_completions, ref_data_completions)),536 )537 # convert ranks to a True/False mask:538 # when rank == 0, it means the first completion is the best539 # when rank == 1, it means the second completion is the best540 return torch.tensor([rank == 0 for rank in ranks_of_first_completion], device=model_data["input_ids"].device)541 542 def _compute_logprobs(self, model, model_data, ref_data, context_length):543 def compute_logprobs_for_data(m, data):544 output = m(data["input_ids"], attention_mask=data["attention_mask"])545 logits = output.logits[:, context_length - 1 : -1]546 token_logprobs = selective_log_softmax(logits, data["input_ids"][:, context_length:])547 return token_logprobs548 549 # Compute logprobs for model completions550 model_logprobs_model_data = compute_logprobs_for_data(model, model_data)551 # Compute logprobs for model on reference completions (for XPO loss)552 model_logprobs_ref_data = compute_logprobs_for_data(model, ref_data)553 554 # Compute logprobs for reference model completions555 with torch.no_grad():556 if self.ref_model is None:557 with model.disable_adapter():558 ref_logprobs_model_data = compute_logprobs_for_data(model, model_data)559 ref_logprobs_ref_data = compute_logprobs_for_data(model, ref_data)560 else:561 ref_logprobs_model_data = compute_logprobs_for_data(self.ref_model, model_data)562 ref_logprobs_ref_data = compute_logprobs_for_data(self.ref_model, ref_data)563 564 # Mask padding tokens565 model_padding_mask = model_data["attention_mask"][:, context_length:] == 0566 ref_padding_mask = ref_data["attention_mask"][:, context_length:] == 0567 model_logprobs_model_data = model_logprobs_model_data.masked_fill(model_padding_mask, 0.0)568 model_logprobs_ref_data = model_logprobs_ref_data.masked_fill(ref_padding_mask, 0.0)569 ref_logprobs_ref_data = ref_logprobs_ref_data.masked_fill(ref_padding_mask, 0.0)570 ref_logprobs_model_data = ref_logprobs_model_data.masked_fill(model_padding_mask, 0.0)571 572 return model_logprobs_model_data, model_logprobs_ref_data, ref_logprobs_ref_data, ref_logprobs_model_data573 574 def _compute_losses(575 self,576 model_logprobs_model_data,577 model_logprobs_ref_data,578 ref_logprobs_ref_data,579 ref_logprobs_model_data,580 chosen_mask,581 ):582 # Compute log probs583 model_logprobs_model_data_sum = model_logprobs_model_data.sum(1)584 model_logprobs_ref_data_sum = model_logprobs_ref_data.sum(1)585 ref_logprobs_ref_data_sum = ref_logprobs_ref_data.sum(1)586 ref_logprobs_model_data_sum = ref_logprobs_model_data.sum(1)587 588 chosen_model_logprobs = torch.where(chosen_mask, model_logprobs_model_data_sum, model_logprobs_ref_data_sum)589 chosen_ref_logprobs = torch.where(chosen_mask, ref_logprobs_model_data_sum, ref_logprobs_ref_data_sum)590 chosen_log_ratios = chosen_model_logprobs - chosen_ref_logprobs591 592 rejected_model_logprobs = torch.where(~chosen_mask, model_logprobs_model_data_sum, model_logprobs_ref_data_sum)593 rejected_ref_logprobs = torch.where(~chosen_mask, ref_logprobs_model_data_sum, ref_logprobs_ref_data_sum)594 rejected_log_ratios = rejected_model_logprobs - rejected_ref_logprobs595 596 # Compute logits as the difference between chosen and rejected log ratios597 logits = chosen_log_ratios - rejected_log_ratios598 599 if self.args.loss_type == "sigmoid":600 dpo_losses = -F.logsigmoid(self.beta * logits)601 elif self.args.loss_type == "ipo":602 dpo_losses = (logits - 1 / (2 * self.beta)) ** 2603 else:604 raise NotImplementedError(f"invalid loss type {self.args.loss_type}")605 606 # Compute XPO specific loss607 xpo_losses = self.alpha * model_logprobs_ref_data_sum608 609 # Total loss610 loss = (dpo_losses + xpo_losses).mean()611 612 return loss, dpo_losses, xpo_losses613 614 def _log_statistics(615 self,616 model_data,617 ref_data,618 model_logprobs_model_data,619 model_logprobs_ref_data,620 ref_logprobs_ref_data,621 ref_logprobs_model_data,622 chosen_mask,623 dpo_losses,624 xpo_losses,625 context_length,626 model_scores=None,627 ref_scores=None,628 ):629 # Helper function to gather and compute mean630 def gather_mean(tensor):631 return self.accelerator.gather_for_metrics(tensor).mean().item()632 633 # Log losses634 self.stats["loss/dpo"].append(gather_mean(dpo_losses))635 self.stats["loss/xpo"].append(gather_mean(xpo_losses))636 637 # Log scores638 if self.reward_model is not None:639 self.stats["objective/model_scores"].append(gather_mean(model_scores))640 self.stats["objective/ref_scores"].append(gather_mean(ref_scores))641 self.stats["objective/scores_margin"].append(gather_mean(model_scores - ref_scores))642 643 # Log logprobs644 model_logprobs_model_data_sum = model_logprobs_model_data.sum(1)645 model_logprobs_ref_data_sum = model_logprobs_ref_data.sum(1)646 ref_logprobs_ref_data_sum = ref_logprobs_ref_data.sum(1)647 ref_logprobs_model_data_sum = ref_logprobs_model_data.sum(1)648 649 chosen_model_logprobs = torch.where(chosen_mask, model_logprobs_model_data_sum, model_logprobs_ref_data_sum)650 chosen_ref_logprobs = torch.where(chosen_mask, ref_logprobs_model_data_sum, ref_logprobs_ref_data_sum)651 chosen_log_ratios = chosen_model_logprobs - chosen_ref_logprobs652 653 rejected_model_logprobs = torch.where(~chosen_mask, model_logprobs_model_data_sum, model_logprobs_ref_data_sum)654 rejected_ref_logprobs = torch.where(~chosen_mask, ref_logprobs_model_data_sum, ref_logprobs_ref_data_sum)655 rejected_log_ratios = rejected_model_logprobs - rejected_ref_logprobs656 657 self.stats["logps/chosen"].append(gather_mean(chosen_model_logprobs.mean() + chosen_ref_logprobs.mean()))658 self.stats["logps/rejected"].append(gather_mean(rejected_model_logprobs.mean() + rejected_ref_logprobs.mean()))659 660 # Log rewards661 # Compute various statistics662 chosen_rewards = chosen_log_ratios * self.beta663 rejected_rewards = rejected_log_ratios * self.beta664 self.stats["rewards/chosen"].append(gather_mean(chosen_rewards.mean()))665 self.stats["rewards/rejected"].append(gather_mean(rejected_rewards.mean()))666 667 # Calculate KL divergence for model and ref data668 kl_model_data = model_logprobs_model_data - ref_logprobs_model_data669 kl_ref_data = model_logprobs_ref_data - ref_logprobs_ref_data670 mean_kl = (kl_model_data.sum(1) + kl_ref_data.sum(1)).mean() / 2671 self.stats["objective/kl"].append(gather_mean(mean_kl))672 673 # Calculate entropy for model and ref data674 entropy_model_data = -model_logprobs_model_data.sum(1)675 entropy_ref_data = -model_logprobs_ref_data.sum(1)676 mean_entropy = (entropy_model_data.mean() + entropy_ref_data.mean()) / 2677 self.stats["objective/entropy"].append(gather_mean(mean_entropy))678 679 # Calculate margins680 margin = chosen_rewards - rejected_rewards681 self.stats["rewards/margins"].append(gather_mean(margin.mean()))682 683 # Calculate accuracy684 accuracy = (margin > 0).float()685 self.stats["rewards/accuracies"].append(gather_mean(accuracy.mean()))686 687 # Log EOS token statistics688 model_eos = (model_data["input_ids"][:, context_length:] == self.processing_class.eos_token_id).any(dim=1)689 ref_eos = (ref_data["input_ids"][:, context_length:] == self.processing_class.eos_token_id).any(dim=1)690 self.stats["val/model_contain_eos_token"].append(gather_mean(model_eos.float()))691 self.stats["val/ref_contain_eos_token"].append(gather_mean(ref_eos.float()))692 693 # Log alpha and beta694 self.stats["alpha"].append(self.alpha)695 self.stats["beta"].append(self.beta)696 697 def training_step(698 self, model: nn.Module, inputs: dict[str, Union[torch.Tensor, Any]], num_items_in_batch: Optional[int] = None699 ) -> torch.Tensor:700 model.train()701 702 # Apply chat template and tokenize the input703 batch_size = len(next(iter(inputs.values())))704 prompts = inputs["prompt"]705 inputs = [{k: v[i] for k, v in inputs.items()} for i in range(batch_size)]706 inputs = [maybe_apply_chat_template(x, self.processing_class) for x in inputs]707 inputs = [self.tokenize_row(x, self.model.config.is_encoder_decoder, self.processing_class) for x in inputs]708 inputs = self.data_collator(inputs)709 710 # need the prompt_ only711 inputs = self._prepare_inputs(inputs)712 context_length = inputs["prompt_input_ids"].shape[1]713 prompts = {714 "input_ids": inputs["prompt_input_ids"],715 "attention_mask": inputs["prompt_attention_mask"],716 "raw": prompts,717 }718 del inputs719 720 # Sample completions from both the model and the reference model721 model_output, ref_output = self._generate_completions(prompts, model)722 723 # Process model completions724 model_data, ref_data = self._process_completions(model_output, ref_output, prompts)725 726 # Compute rewards727 if self.reward_model is not None:728 model_scores, ref_scores = self._compute_rewards(model_data, ref_data, context_length)729 chosen_mask = model_scores >= ref_scores730 else:731 model_scores, ref_scores = None, None732 chosen_mask = self._compute_judge(model_data, ref_data, context_length)733 734 # Compute logprobs735 model_logprobs_model_data, model_logprobs_ref_data, ref_logprobs_ref_data, ref_logprobs_model_data = (736 self._compute_logprobs(model, model_data, ref_data, context_length)737 )738 739 # Compute loss740 loss, dpo_losses, xpo_losses = self._compute_losses(741 model_logprobs_model_data,742 model_logprobs_ref_data,743 ref_logprobs_ref_data,744 ref_logprobs_model_data,745 chosen_mask,746 )747 748 # Log everything749 self._log_statistics(750 model_data,751 ref_data,752 model_logprobs_model_data.detach(),753 model_logprobs_ref_data.detach(),754 ref_logprobs_ref_data,755 ref_logprobs_model_data,756 chosen_mask,757 dpo_losses.detach(),758 xpo_losses.detach(),759 context_length,760 model_scores,761 ref_scores,762 )763 764 if (765 self.args.torch_empty_cache_steps is not None766 and self.state.global_step % self.args.torch_empty_cache_steps == 0767 ):768 empty_cache()769 770 kwargs = {}771 # For LOMO optimizers you need to explicitly use the learning rate772 if self.args.optim in [OptimizerNames.LOMO, OptimizerNames.ADALOMO]:773 kwargs["learning_rate"] = self._get_learning_rate()774 775 if self.args.n_gpu > 1:776 loss = loss.mean() # mean() to average on multi-gpu parallel training777 778 if self.use_apex:779 with amp.scale_loss(loss, self.optimizer) as scaled_loss:780 scaled_loss.backward()781 else:782 self.accelerator.backward(loss, **kwargs)783 784 return loss.detach() / self.args.gradient_accumulation_steps785 786 def create_model_card(787 self,788 model_name: Optional[str] = None,789 dataset_name: Optional[str] = None,790 tags: Union[str, list[str], None] = None,791 ):792 """793 Creates a draft of a model card using the information available to the `Trainer`.794 795 Args:796 model_name (`str` or `None`, *optional*, defaults to `None`):797 Name of the model.798 dataset_name (`str` or `None`, *optional*, defaults to `None`):799 Name of the dataset used for training.800 tags (`str`, `list[str]` or `None`, *optional*, defaults to `None`):801 Tags to be associated with the model card.802 """803 if not self.is_world_process_zero():804 return805 806 if hasattr(self.model.config, "_name_or_path") and not os.path.isdir(self.model.config._name_or_path):807 base_model = self.model.config._name_or_path808 else:809 base_model = None810 811 tags = tags or []812 if isinstance(tags, str):813 tags = [tags]814 815 if hasattr(self.model.config, "unsloth_version"):816 tags.append("unsloth")817 818 citation = textwrap.dedent("""\819 @article{jung2024binary,820 title = {{Exploratory Preference Optimization: Harnessing Implicit Q*-Approximation for Sample-Efficient RLHF}},821 author = {Tengyang Xie and Dylan J. Foster and Akshay Krishnamurthy and Corby Rosset and Ahmed Awadallah and Alexander Rakhlin},822 year = 2024,823 eprint = {arXiv:2405.21046}824 }""")825 826 model_card = generate_model_card(827 base_model=base_model,828 model_name=model_name,829 hub_model_id=self.hub_model_id,830 dataset_name=dataset_name,831 tags=tags,832 wandb_url=wandb.run.get_url() if is_wandb_available() and wandb.run is not None else None,833 comet_url=get_comet_experiment_url(),834 trainer_name="XPO",835 trainer_citation=citation,836 paper_title="Exploratory Preference Optimization: Harnessing Implicit Q*-Approximation for Sample-Efficient RLHF",837 paper_id="2405.21046",838 )839 840 model_card.save(os.path.join(self.args.output_dir, "README.md"))841class UnslothXPOTrainer(_UnslothXPOTrainer):842 """843 844 Initialize XPOTrainer as a subclass of [`OnlineDPOConfig`].845 846 Args:847 model (`transformers.PreTrainedModel`):848 The model to train, preferably an `AutoModelForCausalLM`.849 ref_model (`PreTrainedModelWrapper`):850 Hugging Face transformer model with a casual language modelling head. Used for implicit reward computation and loss. If no851 reference model is provided, the trainer will create a reference model with the same architecture as the model to be optimized.852 reward_model (`transformers.PreTrainedModel`):853 The reward model to score completions with, preferably an `AutoModelForSequenceClassification`.854 judge (`BasePairwiseJudge`):855 The judge to use for pairwise comparison of model completions.856 args (`XPOConfig`):857 The XPO config arguments to use for training.858 data_collator (`transformers.DataCollator`):859 The data collator to use for training. If None is specified, the default data collator (`DPODataCollatorWithPadding`) will be used860 which will pad the sequences to the maximum length of the sequences in the batch, given a dataset of paired sequences.861 train_dataset (`datasets.Dataset`):862 The dataset to use for training.863 eval_dataset (`datasets.Dataset`):864 The dataset to use for evaluation.865 processing_class (`PreTrainedTokenizerBase` or `BaseImageProcessor` or `FeatureExtractionMixin` or `ProcessorMixin`, *optional*):866 Processing class used to process the data. If provided, will be used to automatically process the inputs867 for the model, and it will be saved along the model to make it easier to rerun an interrupted training or868 reuse the fine-tuned model.869 peft_config (`dict`):870 The peft config to use for training.871 compute_metrics (`Callable[[EvalPrediction], dict]`, *optional*):872 The function to use to compute the metrics. Must take a `EvalPrediction` and return873 a dictionary string to metric values.874 callbacks (`list[transformers.TrainerCallback]`):875 The callbacks to use for training.876 optimizers (`tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR]`):877 The optimizer and scheduler to use for training.878 preprocess_logits_for_metrics (`Callable[[torch.Tensor, torch.Tensor], torch.Tensor]`):879 The function to use to preprocess the logits before computing the metrics.880 881 """882 def __init__(883 self,884 model = None,885 ref_model = None,886 reward_model = None,887 judge = None,888 args = None,889 data_collator = None,890 train_dataset = None,891 eval_dataset = None,892 processing_class = None,893 peft_config = None,894 compute_metrics = None,895 callbacks = None,896 preprocess_logits_for_metrics = None,897 **kwargs898 ):899 if args is None: args = UnslothXPOConfig()900 use_bf16 = getattr(args, 'bf16', False)901 use_fp16 = getattr(args, 'fp16', False)902 force_float32 = False903 if os.environ.get('UNSLOTH_FORCE_FLOAT32', '0') == '1':904 print('Unsloth: Switching to float32 training since model cannot work with float16')905 force_float32 = True906 mixed_precision_dtype = os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32')907 dtype = getattr(model.config, 'torch_dtype', None)908 if dtype is None: dtype = model.get_input_embeddings().dtype909 from unsloth_zoo.utils import _get_dtype910 dtype = _get_dtype(dtype)911 float16 = dtype == torch.float16912 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`')913 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`')914 if force_float32:915 args.fp16 = False916 args.bf16 = False917 os.environ['ACCELERATE_MIXED_PRECISION'] = 'no'918 elif (not use_bf16 and not use_fp16) and mixed_precision_dtype == 'float32':919 args.fp16 = float16920 args.bf16 = not float16921 os.environ['ACCELERATE_MIXED_PRECISION'] = 'fp16' if float16 else 'bf16'922 if getattr(args, 'eval_dataset', None) is not None and getattr(args, 'eval_strategy', 'no') == 'no':923 args.eval_strategy = 'steps'924 if getattr(args, 'eval_steps', None) is None: args.eval_steps = 0.1925 ga_steps = getattr(args, 'gradient_accumulation_steps', None)926 if ga_steps is not None and ga_steps > 1:927 from transformers import __version__ as transformers_version928 if Version(transformers_version) <= Version('4.45.2'):929 print('**** Unsloth: Please use our fixed gradient_accumulation_steps by updating transformers, TRL and Unsloth!\n'930 '`pip install --upgrade --no-cache-dir --force-reinstall --no-deps unsloth transformers trl unsloth_zoo`')931 if getattr(args, 'eval_strategy', 'no') != 'no':932 eval_bsz = getattr(args, 'per_device_eval_batch_size', 8)933 if eval_bsz == 8 and args.per_device_train_batch_size < eval_bsz: args.per_device_eval_batch_size = args.per_device_train_batch_size934 if getattr(args, 'eval_accumulation_steps', None) is None and ga_steps is not None: args.eval_accumulation_steps = ga_steps935 fp16_full_eval = getattr(args, 'fp16_full_eval', False)936 bf16_full_eval = getattr(args, 'bf16_full_eval', False)937 if args.fp16 and bf16_full_eval: args.bf16_full_eval = False; args.fp16_full_eval = True938 if args.bf16 and fp16_full_eval: args.bf16_full_eval = True; args.fp16_full_eval = False939 if force_float32:940 args.bf16_full_eval = False941 args.fp16_full_eval = False942 elif os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32') == 'bfloat16':943 args.bf16_full_eval = True944 args.fp16_full_eval = False945 elif not bf16_full_eval and not fp16_full_eval:946 args.bf16_full_eval = args.bf16947 args.fp16_full_eval = args.fp16948 _output_logits = False949 if locals().get('compute_metrics', None) is not None: _output_logits = True950 if locals().get('preprocess_logits_for_metrics', None) is not None: _output_logits = True951 if _output_logits:952 os.environ['UNSLOTH_RETURN_LOGITS'] = '1'953 if 'max_seq_length' not in locals() and not hasattr(args, 'max_seq_length'):954 pass955 else:956 model_max_seq_length = getattr(model, 'max_seq_length', None)957 args_max_seq_length = getattr(args, 'max_seq_length', None)958 if args_max_seq_length is None and model_max_seq_length is not None:959 max_seq_length = model.max_seq_length960 if hasattr(args, 'max_seq_length'): args.max_seq_length = max_seq_length961 if model is not None and hasattr(model, 'for_training'):962 model.for_training()963 if 'tokenizer' in locals() and hasattr(tokenizer, 'padding_side'): tokenizer.padding_side = 'right'964 if 'processing_class' in locals():965 if hasattr(processing_class, 'padding_side'): processing_class.padding_side = 'right'966 if hasattr(processing_class, 'tokenizer') and hasattr(processing_class.tokenizer, 'padding_side'): processing_class.tokenizer.padding_side = 'right'967 __tokenizer = processing_class if 'processing_class' in locals() else tokenizer968 from unsloth_zoo.vision_utils import UnslothVisionDataCollator969 if not isinstance(data_collator, UnslothVisionDataCollator):970 if isinstance(data_collator, DataCollatorForSeq2Seq) and 'labels' not in train_dataset.column_names:971 data_collator = DataCollatorForLanguageModeling(__tokenizer, mlm = False)972 elif isinstance(data_collator, DataCollatorForLanguageModeling) and 'labels' in train_dataset.column_names:973 data_collator = DataCollatorForSeq2Seq(__tokenizer)974 else:975 if hasattr(args, 'remove_unused_columns'): args.remove_unused_columns = False976 if hasattr(args, 'dataset_text_field'): args.dataset_text_field = ''977 if hasattr(args, 'dataset_kwargs'): args.dataset_kwargs = {'skip_prepare_dataset': True}978 if not isinstance(data_collator, UnslothVisionDataCollator):979 if not hasattr(__tokenizer, 'pad') and hasattr(__tokenizer, 'tokenizer'):980 if isinstance(data_collator, DataCollatorForSeq2Seq):981 data_collator = DataCollatorForSeq2Seq(__tokenizer.tokenizer)982 else:983 data_collator = DataCollatorForLanguageModeling(__tokenizer.tokenizer, mlm = False)984 other_metrics = []985 986 from unsloth_zoo.logging_utils import PatchRLStatistics987 PatchRLStatistics('xpo_trainer', other_metrics)988 989 super().__init__(990 model = model,991 ref_model = ref_model,992 reward_model = reward_model,993 judge = judge,994 args = args,995 data_collator = data_collator,996 train_dataset = train_dataset,997 eval_dataset = eval_dataset,998 processing_class = processing_class,999 peft_config = peft_config,1000 compute_metrics = compute_metrics,1001 callbacks = callbacks,1002 preprocess_logits_for_metrics = preprocess_logits_for_metrics,**kwargs)1003 if hasattr(self, 'neftune_hook_handle'):1004 self.neftune_hook_handle.remove()1005 if hasattr(self, 'neftune_hook_handle'): del self.neftune_hook_handle1006 if getattr(args, 'neftune_noise_alpha', None) is not None:1007 model.get_input_embeddings().neftune_noise_alpha = self.neftune_noise_alpha1008 pass1009 1010pass1011 