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.nash_md_trainer import (Any, BaseImageProcessor, BasePairwiseJudge, Callable, Dataset, EvalPrediction, F, FeatureExtractionMixin, GeometricMixtureWrapper, IterableDataset, NashMDConfig, NashMDTrainer, OnlineDPOTrainer, OptimizerNames, Optional, PreTrainedModel, PreTrainedTokenizerBase, ProcessorMixin, SIMPLE_CHAT_TEMPLATE, TrainerCallback, Union, 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 UnslothNashMDConfig(NashMDConfig):44 """45 46 Configuration class for the [`NashMDTrainer`].47 48 Subclass of [`OnlineDPOConfig`] we can use all its arguments and add the following:49 50 Parameters:51 mixture_coef (`float` or `list[float]`, *optional*, defaults to `0.5`):52 Logit mixture coefficient for the model and reference model. If a list of floats is provided then the53 mixture coefficient is selected for each new epoch and the last coefficient is used for the rest of the54 epochs.55 56 """57 vllm_sampling_params: Optional[Any] = field(58 default = None,59 metadata = {'help': 'vLLM SamplingParams'},60 )61 unsloth_num_chunks : Optional[int] = field(62 default = -1,63 metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},64 )65 def __init__(66 self,67 output_dir = None,68 overwrite_output_dir = None,69 do_train = False,70 do_eval = False,71 do_predict = False,72 eval_strategy = 'no',73 prediction_loss_only = False,74 per_device_train_batch_size = 4,75 per_device_eval_batch_size = 4,76 per_gpu_train_batch_size = None,77 per_gpu_eval_batch_size = None,78 gradient_accumulation_steps = 2,79 eval_accumulation_steps = 2,80 eval_delay = 0,81 torch_empty_cache_steps = 250,82 learning_rate = 5e-05,83 weight_decay = 0.01,84 adam_beta1 = 0.9,85 adam_beta2 = 0.999,86 adam_epsilon = 1e-08,87 max_grad_norm = 1.0,88 num_train_epochs = 3.0,89 max_steps = -1,90 lr_scheduler_type = 'linear',91 warmup_ratio = 0.1,92 warmup_steps = 0,93 log_level = 'passive',94 log_level_replica = 'warning',95 log_on_each_node = True,96 logging_dir = None,97 logging_strategy = 'steps',98 logging_first_step = False,99 logging_steps = 1,100 logging_nan_inf_filter = False,101 save_strategy = 'steps',102 save_steps = 500,103 save_total_limit = None,104 save_safetensors = True,105 save_on_each_node = False,106 save_only_model = False,107 restore_callback_states_from_checkpoint = False,108 no_cuda = False,109 use_cpu = False,110 use_mps_device = False,111 seed = 3407,112 data_seed = 3407,113 jit_mode_eval = False,114 use_ipex = False,115 bf16 = False,116 fp16 = False,117 fp16_opt_level = 'O1',118 half_precision_backend = 'auto',119 bf16_full_eval = False,120 fp16_full_eval = False,121 tf32 = None,122 local_rank = -1,123 ddp_backend = None,124 tpu_num_cores = None,125 tpu_metrics_debug = False,126 debug = '',127 dataloader_drop_last = False,128 eval_steps = None,129 dataloader_num_workers = 0,130 dataloader_prefetch_factor = None,131 past_index = -1,132 run_name = None,133 disable_tqdm = None,134 remove_unused_columns = True,135 label_names = None,136 load_best_model_at_end = False,137 metric_for_best_model = None,138 greater_is_better = None,139 ignore_data_skip = False,140 fsdp = '',141 fsdp_min_num_params = 0,142 fsdp_config = None,143 tp_size = 0,144 fsdp_transformer_layer_cls_to_wrap = None,145 accelerator_config = None,146 deepspeed = None,147 label_smoothing_factor = 0.0,148 optim = 'adamw_8bit',149 optim_args = None,150 adafactor = False,151 group_by_length = False,152 length_column_name = 'length',153 report_to = None,154 ddp_find_unused_parameters = None,155 ddp_bucket_cap_mb = None,156 ddp_broadcast_buffers = None,157 dataloader_pin_memory = True,158 dataloader_persistent_workers = False,159 skip_memory_metrics = True,160 use_legacy_prediction_loop = False,161 push_to_hub = False,162 resume_from_checkpoint = None,163 hub_model_id = None,164 hub_strategy = 'every_save',165 hub_token = None,166 hub_private_repo = None,167 hub_always_push = False,168 gradient_checkpointing = False,169 gradient_checkpointing_kwargs = None,170 include_inputs_for_metrics = False,171 eval_do_concat_batches = True,172 fp16_backend = 'auto',173 evaluation_strategy = None,174 push_to_hub_model_id = None,175 push_to_hub_organization = None,176 push_to_hub_token = None,177 mp_parameters = '',178 auto_find_batch_size = False,179 full_determinism = False,180 torchdynamo = None,181 ray_scope = 'last',182 ddp_timeout = 1800,183 torch_compile = False,184 torch_compile_backend = None,185 torch_compile_mode = None,186 dispatch_batches = None,187 split_batches = None,188 include_tokens_per_second = False,189 include_num_input_tokens_seen = False,190 neftune_noise_alpha = None,191 optim_target_modules = None,192 batch_eval_metrics = False,193 eval_on_start = False,194 use_liger_kernel = False,195 eval_use_gather_object = False,196 average_tokens_across_devices = False,197 reward_model_path = None,198 judge = None,199 max_new_tokens = 64,200 max_length = 512,201 temperature = 0.9,202 missing_eos_penalty = None,203 loss_type = 'sigmoid',204 dataset_num_proc = None,205 disable_dropout = True,206 use_vllm = False,207 ds3_gather_for_generation = True,208 vllm_sampling_params = None,209 unsloth_num_chunks = -1,210 **kwargs,211 ):212 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!')213 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!')214 if output_dir is None and save_strategy == 'steps' and save_steps == 500:215 output_dir = 'unsloth_training_checkpoints'216 save_strategy = 'no'217 if dataset_num_proc is None:218 from multiprocessing import cpu_count219 dataset_num_proc = cpu_count()220 221 super().__init__(222 output_dir = output_dir,223 overwrite_output_dir = overwrite_output_dir,224 do_train = do_train,225 do_eval = do_eval,226 do_predict = do_predict,227 eval_strategy = eval_strategy,228 prediction_loss_only = prediction_loss_only,229 per_device_train_batch_size = per_device_train_batch_size,230 per_device_eval_batch_size = per_device_eval_batch_size,231 per_gpu_train_batch_size = per_gpu_train_batch_size,232 per_gpu_eval_batch_size = per_gpu_eval_batch_size,233 gradient_accumulation_steps = gradient_accumulation_steps,234 eval_accumulation_steps = eval_accumulation_steps,235 eval_delay = eval_delay,236 torch_empty_cache_steps = torch_empty_cache_steps,237 learning_rate = learning_rate,238 weight_decay = weight_decay,239 adam_beta1 = adam_beta1,240 adam_beta2 = adam_beta2,241 adam_epsilon = adam_epsilon,242 max_grad_norm = max_grad_norm,243 num_train_epochs = num_train_epochs,244 max_steps = max_steps,245 lr_scheduler_type = lr_scheduler_type,246 warmup_ratio = warmup_ratio,247 warmup_steps = warmup_steps,248 log_level = log_level,249 log_level_replica = log_level_replica,250 log_on_each_node = log_on_each_node,251 logging_dir = logging_dir,252 logging_strategy = logging_strategy,253 logging_first_step = logging_first_step,254 logging_steps = logging_steps,255 logging_nan_inf_filter = logging_nan_inf_filter,256 save_strategy = save_strategy,257 save_steps = save_steps,258 save_total_limit = save_total_limit,259 save_safetensors = save_safetensors,260 save_on_each_node = save_on_each_node,261 save_only_model = save_only_model,262 restore_callback_states_from_checkpoint = restore_callback_states_from_checkpoint,263 no_cuda = no_cuda,264 use_cpu = use_cpu,265 use_mps_device = use_mps_device,266 seed = seed,267 data_seed = data_seed,268 jit_mode_eval = jit_mode_eval,269 use_ipex = use_ipex,270 bf16 = bf16,271 fp16 = fp16,272 fp16_opt_level = fp16_opt_level,273 half_precision_backend = half_precision_backend,274 bf16_full_eval = bf16_full_eval,275 fp16_full_eval = fp16_full_eval,276 tf32 = tf32,277 local_rank = local_rank,278 ddp_backend = ddp_backend,279 tpu_num_cores = tpu_num_cores,280 tpu_metrics_debug = tpu_metrics_debug,281 debug = debug,282 dataloader_drop_last = dataloader_drop_last,283 eval_steps = eval_steps,284 dataloader_num_workers = dataloader_num_workers,285 dataloader_prefetch_factor = dataloader_prefetch_factor,286 past_index = past_index,287 run_name = run_name,288 disable_tqdm = disable_tqdm,289 remove_unused_columns = remove_unused_columns,290 label_names = label_names,291 load_best_model_at_end = load_best_model_at_end,292 metric_for_best_model = metric_for_best_model,293 greater_is_better = greater_is_better,294 ignore_data_skip = ignore_data_skip,295 fsdp = fsdp,296 fsdp_min_num_params = fsdp_min_num_params,297 fsdp_config = fsdp_config,298 tp_size = tp_size,299 fsdp_transformer_layer_cls_to_wrap = fsdp_transformer_layer_cls_to_wrap,300 accelerator_config = accelerator_config,301 deepspeed = deepspeed,302 label_smoothing_factor = label_smoothing_factor,303 optim = optim,304 optim_args = optim_args,305 adafactor = adafactor,306 group_by_length = group_by_length,307 length_column_name = length_column_name,308 report_to = report_to,309 ddp_find_unused_parameters = ddp_find_unused_parameters,310 ddp_bucket_cap_mb = ddp_bucket_cap_mb,311 ddp_broadcast_buffers = ddp_broadcast_buffers,312 dataloader_pin_memory = dataloader_pin_memory,313 dataloader_persistent_workers = dataloader_persistent_workers,314 skip_memory_metrics = skip_memory_metrics,315 use_legacy_prediction_loop = use_legacy_prediction_loop,316 push_to_hub = push_to_hub,317 resume_from_checkpoint = resume_from_checkpoint,318 hub_model_id = hub_model_id,319 hub_strategy = hub_strategy,320 hub_token = hub_token,321 hub_private_repo = hub_private_repo,322 hub_always_push = hub_always_push,323 gradient_checkpointing = gradient_checkpointing,324 gradient_checkpointing_kwargs = gradient_checkpointing_kwargs,325 include_inputs_for_metrics = include_inputs_for_metrics,326 eval_do_concat_batches = eval_do_concat_batches,327 fp16_backend = fp16_backend,328 evaluation_strategy = evaluation_strategy,329 push_to_hub_model_id = push_to_hub_model_id,330 push_to_hub_organization = push_to_hub_organization,331 push_to_hub_token = push_to_hub_token,332 mp_parameters = mp_parameters,333 auto_find_batch_size = auto_find_batch_size,334 full_determinism = full_determinism,335 torchdynamo = torchdynamo,336 ray_scope = ray_scope,337 ddp_timeout = ddp_timeout,338 torch_compile = torch_compile,339 torch_compile_backend = torch_compile_backend,340 torch_compile_mode = torch_compile_mode,341 dispatch_batches = dispatch_batches,342 split_batches = split_batches,343 include_tokens_per_second = include_tokens_per_second,344 include_num_input_tokens_seen = include_num_input_tokens_seen,345 neftune_noise_alpha = neftune_noise_alpha,346 optim_target_modules = optim_target_modules,347 batch_eval_metrics = batch_eval_metrics,348 eval_on_start = eval_on_start,349 use_liger_kernel = use_liger_kernel,350 eval_use_gather_object = eval_use_gather_object,351 average_tokens_across_devices = average_tokens_across_devices,352 reward_model_path = reward_model_path,353 judge = judge,354 max_new_tokens = max_new_tokens,355 max_length = max_length,356 temperature = temperature,357 missing_eos_penalty = missing_eos_penalty,358 loss_type = loss_type,359 dataset_num_proc = dataset_num_proc,360 disable_dropout = disable_dropout,361 use_vllm = use_vllm,362 ds3_gather_for_generation = ds3_gather_for_generation,**kwargs)363 self.vllm_sampling_params = vllm_sampling_params364 self.unsloth_num_chunks = unsloth_num_chunks365pass366 367class _UnslothNashMDTrainer(OnlineDPOTrainer):368 r""""""369 370 _tag_names = ["trl", "nash-md"]371 372 def __init__(373 self,374 model: Union[PreTrainedModel, nn.Module] = None,375 ref_model: Union[PreTrainedModel, nn.Module] = None,376 reward_model: Union[PreTrainedModel, nn.Module, None] = None,377 judge: Optional[BasePairwiseJudge] = None,378 args: Optional[NashMDConfig] = None,379 data_collator: Optional[Callable] = None,380 train_dataset: Optional[Union[Dataset, IterableDataset]] = None,381 eval_dataset: Optional[Union[Dataset, dict[str, Dataset]]] = None,382 processing_class: Optional[383 Union[PreTrainedTokenizerBase, BaseImageProcessor, FeatureExtractionMixin, ProcessorMixin]384 ] = None,385 peft_config: Optional[dict] = None,386 compute_metrics: Optional[Callable[[EvalPrediction], dict]] = None,387 callbacks: Optional[list[TrainerCallback]] = None,388 optimizers: tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR] = (None, None),389 preprocess_logits_for_metrics: Optional[Callable[[torch.Tensor, torch.Tensor], torch.Tensor]] = None,390 ) -> None:391 super().__init__(392 model=model,393 ref_model=ref_model,394 reward_model=reward_model,395 judge=judge,396 args=args,397 data_collator=data_collator,398 train_dataset=train_dataset,399 eval_dataset=eval_dataset,400 processing_class=processing_class,401 reward_processing_class=processing_class, # for now, NashMDTrainer can't use any reward model402 peft_config=peft_config,403 compute_metrics=compute_metrics,404 callbacks=callbacks,405 optimizers=optimizers,406 preprocess_logits_for_metrics=preprocess_logits_for_metrics,407 )408 409 self._mixture_coef = self.args.mixture_coef410 411 # Overwrite the stats dictionary to include NashMD specific statistics412 self.stats = {413 # Remove "non_score_reward", "rlhf_reward", "scores_margin"414 # Add "mixture_coef"415 "loss/kl": [],416 "objective/entropy": [],417 "loss/score": [],418 "rewards/probabilities": [],419 "rewards/accuracies": [],420 "rewards/margins": [],421 "logps/chosen": [],422 "logps/rejected": [],423 "val/model_contain_eos_token": [],424 "val/ref_contain_eos_token": [],425 "beta": [],426 "mixture_coef": [],427 }428 if self.reward_model is not None:429 self.stats["rewards/chosen"] = []430 self.stats["rewards/rejected"] = []431 432 @property433 def mixture_coef(self):434 if isinstance(self._mixture_coef, list):435 epoch = self.state.epoch436 return self._mixture_coef[epoch] if epoch < len(self._mixture_coef) else self._mixture_coef[-1]437 else:438 return self._mixture_coef439 440 def _generate_completions(self, model, prompts):441 with unwrap_model_for_generation(model, self.accelerator) as unwrapped_model:442 model_output = unwrapped_model.generate(443 input_ids=prompts["input_ids"],444 attention_mask=prompts["attention_mask"],445 generation_config=self.generation_config,446 )447 448 ref_model = model if self.ref_model is None else self.ref_model449 with torch.no_grad(), unwrap_model_for_generation(ref_model, self.accelerator) as unwrapped_ref_model:450 mixture_model = GeometricMixtureWrapper(451 model=unwrapped_model,452 ref_model=unwrapped_ref_model,453 generation_config=self.generation_config,454 mixture_coef=self.mixture_coef,455 device=self.accelerator.device,456 )457 458 mixture_output = mixture_model.generate(459 input_ids=prompts["input_ids"],460 attention_mask=prompts["attention_mask"],461 generation_config=self.generation_config,462 )463 464 return model_output, mixture_output465 466 def _process_completions(self, model_output, mixture_output, prompts):467 context_length = prompts["input_ids"].shape[1]468 469 # Process model completions470 model_completion_ids = model_output[:, context_length:]471 model_completion_ids, model_completion_mask = truncate_right(472 model_completion_ids, self.processing_class.eos_token_id, self.processing_class.pad_token_id473 )474 model_data = {475 "input_ids": torch.cat((prompts["input_ids"], model_completion_ids), dim=1),476 "attention_mask": torch.cat((prompts["attention_mask"], model_completion_mask), dim=1),477 "raw": prompts["raw"],478 }479 480 # Process reference model completions481 mixture_completion_ids = mixture_output[:, context_length:]482 mixture_completion_ids, mixture_completion_mask = truncate_right(483 mixture_completion_ids, self.processing_class.eos_token_id, self.processing_class.pad_token_id484 )485 mixture_data = {486 "input_ids": torch.cat((prompts["input_ids"], mixture_completion_ids), dim=1),487 "attention_mask": torch.cat((prompts["attention_mask"], mixture_completion_mask), dim=1),488 "raw": prompts["raw"],489 }490 491 return model_data, mixture_data492 493 def _compute_rewards(self, model_data, mixture_data, context_length):494 with torch.no_grad():495 _, model_scores, _ = get_reward(496 self.reward_model, model_data["input_ids"], self.processing_class.pad_token_id, context_length497 )498 _, mixture_scores, _ = get_reward(499 self.reward_model, mixture_data["input_ids"], self.processing_class.pad_token_id, context_length500 )501 502 # Apply EOS penalty if needed503 if self.args.missing_eos_penalty is not None:504 model_contain_eos = torch.any(model_data["input_ids"] == self.processing_class.eos_token_id, dim=-1)505 mixture_contain_eos = torch.any(mixture_data["input_ids"] == self.processing_class.eos_token_id, dim=-1)506 model_scores[~model_contain_eos] -= self.args.missing_eos_penalty507 mixture_scores[~mixture_contain_eos] -= self.args.missing_eos_penalty508 509 return model_scores, mixture_scores510 511 def _compute_judge(self, model_data, mixture_data, context_length):512 prompts = model_data["raw"]513 model_data_completions = self.processing_class.batch_decode(514 model_data["input_ids"][:, context_length:], skip_special_tokens=True515 )516 model_data_completions = [completion.strip() for completion in model_data_completions]517 518 mixture_data_completions = self.processing_class.batch_decode(519 mixture_data["input_ids"][:, context_length:], skip_special_tokens=True520 )521 mixture_data_completions = [completion.strip() for completion in mixture_data_completions]522 if is_conversational({"prompt": prompts[0]}):523 model_data_completions = [524 [{"role": "assistant", "content": completion}] for completion in model_data_completions525 ]526 environment = jinja2.Environment()527 template = environment.from_string(SIMPLE_CHAT_TEMPLATE)528 prompts = [template.render(messages=message) for message in prompts]529 model_data_completions = [template.render(messages=completion) for completion in model_data_completions]530 531 mixture_data_completions = [532 [{"role": "assistant", "content": completion}] for completion in mixture_data_completions533 ]534 mixture_data_completions = [535 template.render(messages=completion) for completion in mixture_data_completions536 ]537 538 probability = self.judge.judge(539 prompts,540 list(zip(model_data_completions, mixture_data_completions)),541 return_scores=True,542 )543 return torch.tensor(probability, device=model_data["input_ids"].device)544 545 def _compute_logprobs(self, model, model_data, context_length):546 def compute_logprobs_for_data(m, data):547 output = m(data["input_ids"], attention_mask=data["attention_mask"])548 logits = output.logits[:, context_length - 1 : -1]549 token_logprobs = selective_log_softmax(logits, data["input_ids"][:, context_length:])550 return token_logprobs551 552 # Compute logprobs for model completions under the model553 model_logprobs_model_data = compute_logprobs_for_data(model, model_data)554 555 # Compute logprobs of model completions under the reference model556 with torch.no_grad():557 if self.ref_model is None:558 with model.disable_adapter():559 ref_logprobs_model_data = compute_logprobs_for_data(model, model_data)560 else:561 ref_logprobs_model_data = compute_logprobs_for_data(self.ref_model, model_data)562 563 # Mask padding tokens564 model_padding_mask = model_data["attention_mask"][:, context_length:] == 0565 model_logprobs_model_data = model_logprobs_model_data.masked_fill(model_padding_mask, 0.0)566 ref_logprobs_model_data = ref_logprobs_model_data.masked_fill(model_padding_mask, 0.0)567 568 return (model_logprobs_model_data, ref_logprobs_model_data)569 570 def _compute_losses(571 self,572 model_logprobs_model_data,573 ref_logprobs_model_data,574 probability,575 ):576 # reinforce score where 0.5 is a control variate577 score = (probability - 0.5) * model_logprobs_model_data.sum(1)578 579 # kl divergence via reinforce580 with torch.no_grad():581 log_ratio = model_logprobs_model_data - ref_logprobs_model_data582 kl_div_log = log_ratio.sum(1)583 kl_div_loss = (log_ratio * model_logprobs_model_data).sum(1)584 585 # final loss586 loss = self.beta * kl_div_loss - score587 588 return loss.mean(), score, kl_div_log589 590 def _log_statistics(591 self,592 model_data,593 mixture_data,594 model_logprobs_model_data,595 ref_logprobs_model_data,596 probability,597 score,598 kl_div,599 context_length,600 model_scores=None,601 mixture_scores=None,602 ):603 # Helper function to gather and compute mean604 def gather_mean(tensor):605 return self.accelerator.gather_for_metrics(tensor).mean().item()606 607 # Log score608 self.stats["loss/score"].append(gather_mean(score))609 # Log KL divergence610 self.stats["loss/kl"].append(gather_mean(kl_div))611 612 # Log logprobs613 model_logprobs_model_data_sum = model_logprobs_model_data.sum(1)614 ref_logprobs_model_data_sum = ref_logprobs_model_data.sum(1)615 616 self.stats["logps/chosen"].append(gather_mean(model_logprobs_model_data_sum))617 self.stats["logps/rejected"].append(gather_mean(ref_logprobs_model_data_sum))618 619 # Log rewards620 if self.reward_model is not None:621 self.stats["rewards/chosen"].append(gather_mean(model_scores))622 self.stats["rewards/rejected"].append(gather_mean(mixture_scores))623 624 # Log probabilities625 self.stats["rewards/probabilities"].append(gather_mean(probability))626 627 # Calculate entropy for model data628 entropy_model_data = -model_logprobs_model_data.sum(1)629 self.stats["objective/entropy"].append(gather_mean(entropy_model_data))630 631 # Calculate margins632 margin = model_logprobs_model_data_sum - ref_logprobs_model_data_sum633 self.stats["rewards/margins"].append(gather_mean(margin))634 635 # Calculate accuracy636 accuracy = (margin > 0).float()637 self.stats["rewards/accuracies"].append(gather_mean(accuracy))638 639 # Log EOS token statistics640 model_eos = (model_data["input_ids"][:, context_length:] == self.processing_class.eos_token_id).any(dim=1)641 mixture_eos = (mixture_data["input_ids"][:, context_length:] == self.processing_class.eos_token_id).any(dim=1)642 self.stats["val/model_contain_eos_token"].append(gather_mean(model_eos.float()))643 self.stats["val/ref_contain_eos_token"].append(gather_mean(mixture_eos.float()))644 645 # Log beta and mixture coef646 self.stats["beta"].append(self.beta)647 self.stats["mixture_coef"].append(self.mixture_coef)648 649 def training_step(650 self, model: nn.Module, inputs: dict[str, Union[torch.Tensor, Any]], num_items_in_batch: Optional[int] = None651 ) -> torch.Tensor:652 model.train()653 654 # Apply chat template and tokenize the input655 batch_size = len(next(iter(inputs.values())))656 prompts = inputs["prompt"]657 inputs = [{k: v[i] for k, v in inputs.items()} for i in range(batch_size)]658 inputs = [maybe_apply_chat_template(x, self.processing_class) for x in inputs]659 inputs = [self.tokenize_row(x, self.model.config.is_encoder_decoder, self.processing_class) for x in inputs]660 inputs = self.data_collator(inputs)661 662 # need the prompt_ only663 inputs = self._prepare_inputs(inputs)664 context_length = inputs["prompt_input_ids"].shape[1]665 prompts = {666 "input_ids": inputs["prompt_input_ids"],667 "attention_mask": inputs["prompt_attention_mask"],668 "raw": prompts,669 }670 del inputs671 672 # Sample completions from both the model and the reference model673 model_output, mixture_output = self._generate_completions(model, prompts)674 675 # Process model completions676 model_data, mixture_data = self._process_completions(model_output, mixture_output, prompts)677 678 # Compute rewards679 if self.reward_model is not None:680 model_scores, mixture_scores = self._compute_rewards(model_data, mixture_data, context_length)681 # probability of the model data vs the mixture data682 probability = F.sigmoid(model_scores - mixture_scores)683 else:684 model_scores, mixture_scores = None, None685 probability = self._compute_judge(model_data, mixture_data, context_length)686 687 # Compute logprobs688 model_logprobs_model_data, ref_logprobs_model_data = self._compute_logprobs(model, model_data, context_length)689 690 # Compute loss691 loss, score, kl_div = self._compute_losses(model_logprobs_model_data, ref_logprobs_model_data, probability)692 693 # Log everything694 self._log_statistics(695 model_data,696 mixture_data,697 model_logprobs_model_data.detach(),698 ref_logprobs_model_data,699 probability,700 score.detach(),701 kl_div.detach(),702 context_length,703 model_scores,704 mixture_scores,705 )706 707 if (708 self.args.torch_empty_cache_steps is not None709 and self.state.global_step % self.args.torch_empty_cache_steps == 0710 ):711 empty_cache()712 713 kwargs = {}714 # For LOMO optimizers you need to explicitly use the learning rate715 if self.args.optim in [OptimizerNames.LOMO, OptimizerNames.ADALOMO]:716 kwargs["learning_rate"] = self._get_learning_rate()717 718 if self.args.n_gpu > 1:719 loss = loss.mean() # mean() to average on multi-gpu parallel training720 721 if self.use_apex:722 with amp.scale_loss(loss, self.optimizer) as scaled_loss:723 scaled_loss.backward()724 else:725 self.accelerator.backward(loss, **kwargs)726 727 return loss.detach() / self.args.gradient_accumulation_steps728 729 def create_model_card(730 self,731 model_name: Optional[str] = None,732 dataset_name: Optional[str] = None,733 tags: Union[str, list[str], None] = None,734 ):735 """736 Creates a draft of a model card using the information available to the `Trainer`.737 738 Args:739 model_name (`str` or `None`, *optional*, defaults to `None`):740 Name of the model.741 dataset_name (`str` or `None`, *optional*, defaults to `None`):742 Name of the dataset used for training.743 tags (`str`, `list[str]` or `None`, *optional*, defaults to `None`):744 Tags to be associated with the model card.745 """746 if not self.is_world_process_zero():747 return748 749 if hasattr(self.model.config, "_name_or_path") and not os.path.isdir(self.model.config._name_or_path):750 base_model = self.model.config._name_or_path751 else:752 base_model = None753 754 tags = tags or []755 if isinstance(tags, str):756 tags = [tags]757 758 if hasattr(self.model.config, "unsloth_version"):759 tags.append("unsloth")760 761 citation = textwrap.dedent("""\762 @inproceedings{munos2024nash,763 title = {{Nash Learning from Human Feedback}},764 author = {R{\'{e}}mi Munos and Michal Valko and Daniele Calandriello and Mohammad Gheshlaghi Azar and Mark Rowland and Zhaohan Daniel Guo and Yunhao Tang and Matthieu Geist and Thomas Mesnard and C{\\^{o}}me Fiegel and Andrea Michi and Marco Selvi and Sertan Girgin and Nikola Momchev and Olivier Bachem and Daniel J. Mankowitz and Doina Precup and Bilal Piot},765 year = 2024,766 booktitle = {Forty-first International Conference on Machine Learning, {ICML} 2024, Vienna, Austria, July 21-27, 2024},767 publisher = {OpenReview.net},768 url = {https://openreview.net/forum?id=Y5AmNYiyCQ}769 }""")770 771 model_card = generate_model_card(772 base_model=base_model,773 model_name=model_name,774 hub_model_id=self.hub_model_id,775 dataset_name=dataset_name,776 tags=tags,777 wandb_url=wandb.run.get_url() if is_wandb_available() and wandb.run is not None else None,778 comet_url=get_comet_experiment_url(),779 trainer_name="Nash-MD",780 trainer_citation=citation,781 paper_title="Nash Learning from Human Feedback",782 paper_id="2312.00886",783 )784 785 model_card.save(os.path.join(self.args.output_dir, "README.md"))786class UnslothNashMDTrainer(_UnslothNashMDTrainer):787 """788 789 Initialize NashMDTrainer as a subclass of [`OnlineDPOConfig`].790 791 Args:792 model (`transformers.PreTrainedModel`):793 The model to train, preferably an `AutoModelForCausalLM`.794 ref_model (`PreTrainedModelWrapper`):795 Hugging Face transformer model with a casual language modelling head. Used for implicit reward computation and loss. If no796 reference model is provided, the trainer will create a reference model with the same architecture as the model to be optimized.797 reward_model (`transformers.PreTrainedModel`):798 The reward model to score completions with, preferably an `AutoModelForSequenceClassification`.799 judge (`BasePairwiseJudge`):800 The judge to use for pairwise comparison of model completions.801 args (`NashMDConfig`):802 The NashMD config arguments to use for training.803 data_collator (`transformers.DataCollator`):804 The data collator to use for training. If None is specified, the default data collator (`DPODataCollatorWithPadding`) will be used805 which will pad the sequences to the maximum length of the sequences in the batch, given a dataset of paired sequences.806 train_dataset (`datasets.Dataset`):807 The dataset to use for training.808 eval_dataset (`datasets.Dataset`):809 The dataset to use for evaluation.810 processing_class (`PreTrainedTokenizerBase` or `BaseImageProcessor` or `FeatureExtractionMixin` or `ProcessorMixin`, *optional*):811 Processing class used to process the data. If provided, will be used to automatically process the inputs812 for the model, and it will be saved along the model to make it easier to rerun an interrupted training or813 reuse the fine-tuned model.814 peft_config (`dict`):815 The peft config to use for training.816 compute_metrics (`Callable[[EvalPrediction], dict]`, *optional*):817 The function to use to compute the metrics. Must take a `EvalPrediction` and return818 a dictionary string to metric values.819 callbacks (`list[transformers.TrainerCallback]`):820 The callbacks to use for training.821 optimizers (`tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR]`):822 The optimizer and scheduler to use for training.823 preprocess_logits_for_metrics (`Callable[[torch.Tensor, torch.Tensor], torch.Tensor]`):824 The function to use to preprocess the logits before computing the metrics.825 826 """827 def __init__(828 self,829 model = None,830 ref_model = None,831 reward_model = None,832 judge = None,833 args = None,834 data_collator = None,835 train_dataset = None,836 eval_dataset = None,837 processing_class = None,838 peft_config = None,839 compute_metrics = None,840 callbacks = None,841 preprocess_logits_for_metrics = None,842 **kwargs843 ):844 if args is None: args = UnslothNashMDConfig()845 use_bf16 = getattr(args, 'bf16', False)846 use_fp16 = getattr(args, 'fp16', False)847 force_float32 = False848 if os.environ.get('UNSLOTH_FORCE_FLOAT32', '0') == '1':849 print('Unsloth: Switching to float32 training since model cannot work with float16')850 force_float32 = True851 mixed_precision_dtype = os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32')852 dtype = getattr(model.config, 'torch_dtype', None)853 if dtype is None: dtype = model.get_input_embeddings().dtype854 from unsloth_zoo.utils import _get_dtype855 dtype = _get_dtype(dtype)856 float16 = dtype == torch.float16857 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`')858 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`')859 if force_float32:860 args.fp16 = False861 args.bf16 = False862 os.environ['ACCELERATE_MIXED_PRECISION'] = 'no'863 elif (not use_bf16 and not use_fp16) and mixed_precision_dtype == 'float32':864 args.fp16 = float16865 args.bf16 = not float16866 os.environ['ACCELERATE_MIXED_PRECISION'] = 'fp16' if float16 else 'bf16'867 if getattr(args, 'eval_dataset', None) is not None and getattr(args, 'eval_strategy', 'no') == 'no':868 args.eval_strategy = 'steps'869 if getattr(args, 'eval_steps', None) is None: args.eval_steps = 0.1870 ga_steps = getattr(args, 'gradient_accumulation_steps', None)871 if ga_steps is not None and ga_steps > 1:872 from transformers import __version__ as transformers_version873 if Version(transformers_version) <= Version('4.45.2'):874 print('**** Unsloth: Please use our fixed gradient_accumulation_steps by updating transformers, TRL and Unsloth!\n'875 '`pip install --upgrade --no-cache-dir --force-reinstall --no-deps unsloth transformers trl unsloth_zoo`')876 if getattr(args, 'eval_strategy', 'no') != 'no':877 eval_bsz = getattr(args, 'per_device_eval_batch_size', 8)878 if eval_bsz == 8 and args.per_device_train_batch_size < eval_bsz: args.per_device_eval_batch_size = args.per_device_train_batch_size879 if getattr(args, 'eval_accumulation_steps', None) is None and ga_steps is not None: args.eval_accumulation_steps = ga_steps880 fp16_full_eval = getattr(args, 'fp16_full_eval', False)881 bf16_full_eval = getattr(args, 'bf16_full_eval', False)882 if args.fp16 and bf16_full_eval: args.bf16_full_eval = False; args.fp16_full_eval = True883 if args.bf16 and fp16_full_eval: args.bf16_full_eval = True; args.fp16_full_eval = False884 if force_float32:885 args.bf16_full_eval = False886 args.fp16_full_eval = False887 elif os.environ.get('UNSLOTH_MIXED_PRECISION', 'float32') == 'bfloat16':888 args.bf16_full_eval = True889 args.fp16_full_eval = False890 elif not bf16_full_eval and not fp16_full_eval:891 args.bf16_full_eval = args.bf16892 args.fp16_full_eval = args.fp16893 _output_logits = False894 if locals().get('compute_metrics', None) is not None: _output_logits = True895 if locals().get('preprocess_logits_for_metrics', None) is not None: _output_logits = True896 if _output_logits:897 os.environ['UNSLOTH_RETURN_LOGITS'] = '1'898 if 'max_seq_length' not in locals() and not hasattr(args, 'max_seq_length'):899 pass900 else:901 model_max_seq_length = getattr(model, 'max_seq_length', None)902 args_max_seq_length = getattr(args, 'max_seq_length', None)903 if args_max_seq_length is None and model_max_seq_length is not None:904 max_seq_length = model.max_seq_length905 if hasattr(args, 'max_seq_length'): args.max_seq_length = max_seq_length906 if model is not None and hasattr(model, 'for_training'):907 model.for_training()908 if 'tokenizer' in locals() and hasattr(tokenizer, 'padding_side'): tokenizer.padding_side = 'right'909 if 'processing_class' in locals():910 if hasattr(processing_class, 'padding_side'): processing_class.padding_side = 'right'911 if hasattr(processing_class, 'tokenizer') and hasattr(processing_class.tokenizer, 'padding_side'): processing_class.tokenizer.padding_side = 'right'912 __tokenizer = processing_class if 'processing_class' in locals() else tokenizer913 from unsloth_zoo.vision_utils import UnslothVisionDataCollator914 if not isinstance(data_collator, UnslothVisionDataCollator):915 if isinstance(data_collator, DataCollatorForSeq2Seq) and 'labels' not in train_dataset.column_names:916 data_collator = DataCollatorForLanguageModeling(__tokenizer, mlm = False)917 elif isinstance(data_collator, DataCollatorForLanguageModeling) and 'labels' in train_dataset.column_names:918 data_collator = DataCollatorForSeq2Seq(__tokenizer)919 else:920 if hasattr(args, 'remove_unused_columns'): args.remove_unused_columns = False921 if hasattr(args, 'dataset_text_field'): args.dataset_text_field = ''922 if hasattr(args, 'dataset_kwargs'): args.dataset_kwargs = {'skip_prepare_dataset': True}923 if not isinstance(data_collator, UnslothVisionDataCollator):924 if not hasattr(__tokenizer, 'pad') and hasattr(__tokenizer, 'tokenizer'):925 if isinstance(data_collator, DataCollatorForSeq2Seq):926 data_collator = DataCollatorForSeq2Seq(__tokenizer.tokenizer)927 else:928 data_collator = DataCollatorForLanguageModeling(__tokenizer.tokenizer, mlm = False)929 other_metrics = []930 931 from unsloth_zoo.logging_utils import PatchRLStatistics932 PatchRLStatistics('nash_md_trainer', other_metrics)933 934 super().__init__(935 model = model,936 ref_model = ref_model,937 reward_model = reward_model,938 judge = judge,939 args = args,940 data_collator = data_collator,941 train_dataset = train_dataset,942 eval_dataset = eval_dataset,943 processing_class = processing_class,944 peft_config = peft_config,945 compute_metrics = compute_metrics,946 callbacks = callbacks,947 preprocess_logits_for_metrics = preprocess_logits_for_metrics,**kwargs)948 if hasattr(self, 'neftune_hook_handle'):949 self.neftune_hook_handle.remove()950 if hasattr(self, 'neftune_hook_handle'): del self.neftune_hook_handle951 if getattr(args, 'neftune_noise_alpha', None) is not None:952 model.get_input_embeddings().neftune_noise_alpha = self.neftune_noise_alpha953 pass954 955pass956 