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.alignprop_trainer import (Accelerator, AlignPropConfig, AlignPropTrainer, Any, Callable, DDPOStableDiffusionPipeline, Optional, ProjectConfiguration, PyTorchModelHubMixin, Union, defaultdict, generate_model_card, get_comet_experiment_url, is_wandb_available, logger, os, set_seed, textwrap, torch, warn)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 UnslothAlignPropConfig(AlignPropConfig):44 """45 46 Configuration class for the [`AlignPropTrainer`].47 48 Using [`~transformers.HfArgumentParser`] we can turn this class into49 [argparse](https://docs.python.org/3/library/argparse#module-argparse) arguments that can be specified on the50 command line.51 52 Parameters:53 exp_name (`str`, *optional*, defaults to `os.path.basename(sys.argv[0])[: -len(".py")]`):54 Name of this experiment (defaults to the file name without the extension).55 run_name (`str`, *optional*, defaults to `""`):56 Name of this run.57 seed (`int`, *optional*, defaults to `0`):58 Random seed for reproducibility.59 log_with (`str` or `None`, *optional*, defaults to `None`):60 Log with either `"wandb"` or `"tensorboard"`. Check61 [tracking](https://huggingface.co/docs/accelerate/usage_guides/tracking) for more details.62 log_image_freq (`int`, *optional*, defaults to `1`):63 Frequency for logging images.64 tracker_kwargs (`dict[str, Any]`, *optional*, defaults to `{}`):65 Keyword arguments for the tracker (e.g., `wandb_project`).66 accelerator_kwargs (`dict[str, Any]`, *optional*, defaults to `{}`):67 Keyword arguments for the accelerator.68 project_kwargs (`dict[str, Any]`, *optional*, defaults to `{}`):69 Keyword arguments for the accelerator project config (e.g., `logging_dir`).70 tracker_project_name (`str`, *optional*, defaults to `"trl"`):71 Name of project to use for tracking.72 logdir (`str`, *optional*, defaults to `"logs"`):73 Top-level logging directory for checkpoint saving.74 num_epochs (`int`, *optional*, defaults to `100`):75 Number of epochs to train.76 save_freq (`int`, *optional*, defaults to `1`):77 Number of epochs between saving model checkpoints.78 num_checkpoint_limit (`int`, *optional*, defaults to `5`):79 Number of checkpoints to keep before overwriting old ones.80 mixed_precision (`str`, *optional*, defaults to `"fp16"`):81 Mixed precision training.82 allow_tf32 (`bool`, *optional*, defaults to `True`):83 Allow `tf32` on Ampere GPUs.84 resume_from (`str`, *optional*, defaults to `""`):85 Path to resume training from a checkpoint.86 sample_num_steps (`int`, *optional*, defaults to `50`):87 Number of sampler inference steps.88 sample_eta (`float`, *optional*, defaults to `1.0`):89 Eta parameter for the DDIM sampler.90 sample_guidance_scale (`float`, *optional*, defaults to `5.0`):91 Classifier-free guidance weight.92 train_batch_size (`int`, *optional*, defaults to `1`):93 Batch size for training.94 train_use_8bit_adam (`bool`, *optional*, defaults to `False`):95 Whether to use the 8bit Adam optimizer from `bitsandbytes`.96 train_learning_rate (`float`, *optional*, defaults to `1e-3`):97 Learning rate.98 train_adam_beta1 (`float`, *optional*, defaults to `0.9`):99 Beta1 for Adam optimizer.100 train_adam_beta2 (`float`, *optional*, defaults to `0.999`):101 Beta2 for Adam optimizer.102 train_adam_weight_decay (`float`, *optional*, defaults to `1e-4`):103 Weight decay for Adam optimizer.104 train_adam_epsilon (`float`, *optional*, defaults to `1e-8`):105 Epsilon value for Adam optimizer.106 train_gradient_accumulation_steps (`int`, *optional*, defaults to `1`):107 Number of gradient accumulation steps.108 train_max_grad_norm (`float`, *optional*, defaults to `1.0`):109 Maximum gradient norm for gradient clipping.110 negative_prompts (`str` or `None`, *optional*, defaults to `None`):111 Comma-separated list of prompts to use as negative examples.112 truncated_backprop_rand (`bool`, *optional*, defaults to `True`):113 If `True`, randomized truncation to different diffusion timesteps is used.114 truncated_backprop_timestep (`int`, *optional*, defaults to `49`):115 Absolute timestep to which the gradients are backpropagated. Used only if `truncated_backprop_rand=False`.116 truncated_rand_backprop_minmax (`tuple[int, int]`, *optional*, defaults to `(0, 50)`):117 Range of diffusion timesteps for randomized truncated backpropagation.118 push_to_hub (`bool`, *optional*, defaults to `False`):119 Whether to push the final model to the Hub.120 121 """122 vllm_sampling_params: Optional[Any] = field(123 default = None,124 metadata = {'help': 'vLLM SamplingParams'},125 )126 unsloth_num_chunks : Optional[int] = field(127 default = -1,128 metadata = {'help': 'Chunk size to reduce memory usage. -1 is most efficient.'},129 )130 def __init__(131 self,132 exp_name = 'demo',133 run_name = '',134 seed = 3407,135 log_with = None,136 log_image_freq = 1,137 tracker_project_name = 'trl',138 logdir = 'logs',139 num_epochs = 100,140 save_freq = 1,141 num_checkpoint_limit = 5,142 mixed_precision = 'fp16',143 allow_tf32 = True,144 resume_from = '',145 sample_num_steps = 50,146 sample_eta = 1.0,147 sample_guidance_scale = 5.0,148 train_batch_size = 1,149 train_use_8bit_adam = False,150 train_learning_rate = 5e-05,151 train_adam_beta1 = 0.9,152 train_adam_beta2 = 0.999,153 train_adam_weight_decay = 0.01,154 train_adam_epsilon = 1e-08,155 train_gradient_accumulation_steps = 2,156 train_max_grad_norm = 1.0,157 negative_prompts = None,158 truncated_backprop_rand = True,159 truncated_backprop_timestep = 49,160 push_to_hub = False,161 vllm_sampling_params = None,162 unsloth_num_chunks = -1,163 **kwargs,164 ):165 166 super().__init__(167 exp_name = exp_name,168 run_name = run_name,169 seed = seed,170 log_with = log_with,171 log_image_freq = log_image_freq,172 tracker_project_name = tracker_project_name,173 logdir = logdir,174 num_epochs = num_epochs,175 save_freq = save_freq,176 num_checkpoint_limit = num_checkpoint_limit,177 mixed_precision = mixed_precision,178 allow_tf32 = allow_tf32,179 resume_from = resume_from,180 sample_num_steps = sample_num_steps,181 sample_eta = sample_eta,182 sample_guidance_scale = sample_guidance_scale,183 train_batch_size = train_batch_size,184 train_use_8bit_adam = train_use_8bit_adam,185 train_learning_rate = train_learning_rate,186 train_adam_beta1 = train_adam_beta1,187 train_adam_beta2 = train_adam_beta2,188 train_adam_weight_decay = train_adam_weight_decay,189 train_adam_epsilon = train_adam_epsilon,190 train_gradient_accumulation_steps = train_gradient_accumulation_steps,191 train_max_grad_norm = train_max_grad_norm,192 negative_prompts = negative_prompts,193 truncated_backprop_rand = truncated_backprop_rand,194 truncated_backprop_timestep = truncated_backprop_timestep,195 push_to_hub = push_to_hub,**kwargs)196 self.vllm_sampling_params = vllm_sampling_params197 self.unsloth_num_chunks = unsloth_num_chunks198pass199 200class _UnslothAlignPropTrainer(PyTorchModelHubMixin):201 """"""202 203 _tag_names = ["trl", "alignprop"]204 205 def __init__(206 self,207 config: AlignPropConfig,208 reward_function: Callable[[torch.Tensor, tuple[str], tuple[Any]], torch.Tensor],209 prompt_function: Callable[[], tuple[str, Any]],210 sd_pipeline: DDPOStableDiffusionPipeline,211 image_samples_hook: Optional[Callable[[Any, Any, Any], Any]] = None,212 ):213 if image_samples_hook is None:214 warn("No image_samples_hook provided; no images will be logged")215 216 self.prompt_fn = prompt_function217 self.reward_fn = reward_function218 self.config = config219 self.image_samples_callback = image_samples_hook220 221 accelerator_project_config = ProjectConfiguration(**self.config.project_kwargs)222 223 if self.config.resume_from:224 self.config.resume_from = os.path.normpath(os.path.expanduser(self.config.resume_from))225 if "checkpoint_" not in os.path.basename(self.config.resume_from):226 # get the most recent checkpoint in this directory227 checkpoints = list(228 filter(229 lambda x: "checkpoint_" in x,230 os.listdir(self.config.resume_from),231 )232 )233 if len(checkpoints) == 0:234 raise ValueError(f"No checkpoints found in {self.config.resume_from}")235 checkpoint_numbers = sorted([int(x.split("_")[-1]) for x in checkpoints])236 self.config.resume_from = os.path.join(237 self.config.resume_from,238 f"checkpoint_{checkpoint_numbers[-1]}",239 )240 241 accelerator_project_config.iteration = checkpoint_numbers[-1] + 1242 243 self.accelerator = Accelerator(244 log_with=self.config.log_with,245 mixed_precision=self.config.mixed_precision,246 project_config=accelerator_project_config,247 # we always accumulate gradients across timesteps; we want config.train.gradient_accumulation_steps to be the248 # number of *samples* we accumulate across, so we need to multiply by the number of training timesteps to get249 # the total number of optimizer steps to accumulate across.250 gradient_accumulation_steps=self.config.train_gradient_accumulation_steps,251 **self.config.accelerator_kwargs,252 )253 254 is_using_tensorboard = config.log_with is not None and config.log_with == "tensorboard"255 256 if self.accelerator.is_main_process:257 self.accelerator.init_trackers(258 self.config.tracker_project_name,259 config=dict(alignprop_trainer_config=config.to_dict())260 if not is_using_tensorboard261 else config.to_dict(),262 init_kwargs=self.config.tracker_kwargs,263 )264 265 logger.info(f"\n{config}")266 267 set_seed(self.config.seed, device_specific=True)268 269 self.sd_pipeline = sd_pipeline270 271 self.sd_pipeline.set_progress_bar_config(272 position=1,273 disable=not self.accelerator.is_local_main_process,274 leave=False,275 desc="Timestep",276 dynamic_ncols=True,277 )278 279 # For mixed precision training we cast all non-trainable weights (vae, non-lora text_encoder and non-lora unet) to half-precision280 # as these weights are only used for inference, keeping weights in full precision is not required.281 if self.accelerator.mixed_precision == "fp16":282 inference_dtype = torch.float16283 elif self.accelerator.mixed_precision == "bf16":284 inference_dtype = torch.bfloat16285 else:286 inference_dtype = torch.float32287 288 self.sd_pipeline.vae.to(self.accelerator.device, dtype=inference_dtype)289 self.sd_pipeline.text_encoder.to(self.accelerator.device, dtype=inference_dtype)290 self.sd_pipeline.unet.to(self.accelerator.device, dtype=inference_dtype)291 292 trainable_layers = self.sd_pipeline.get_trainable_layers()293 294 self.accelerator.register_save_state_pre_hook(self._save_model_hook)295 self.accelerator.register_load_state_pre_hook(self._load_model_hook)296 297 # Enable TF32 for faster training on Ampere GPUs,298 # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices299 if self.config.allow_tf32:300 torch.backends.cuda.matmul.allow_tf32 = True301 302 self.optimizer = self._setup_optimizer(303 trainable_layers.parameters() if not isinstance(trainable_layers, list) else trainable_layers304 )305 306 self.neg_prompt_embed = self.sd_pipeline.text_encoder(307 self.sd_pipeline.tokenizer(308 [""] if self.config.negative_prompts is None else self.config.negative_prompts,309 return_tensors="pt",310 padding="max_length",311 truncation=True,312 max_length=self.sd_pipeline.tokenizer.model_max_length,313 ).input_ids.to(self.accelerator.device)314 )[0]315 316 # NOTE: for some reason, autocast is necessary for non-lora training but for lora training it isn't necessary and it uses317 # more memory318 self.autocast = self.sd_pipeline.autocast or self.accelerator.autocast319 320 if hasattr(self.sd_pipeline, "use_lora") and self.sd_pipeline.use_lora:321 unet, self.optimizer = self.accelerator.prepare(trainable_layers, self.optimizer)322 self.trainable_layers = list(filter(lambda p: p.requires_grad, unet.parameters()))323 else:324 self.trainable_layers, self.optimizer = self.accelerator.prepare(trainable_layers, self.optimizer)325 326 if config.resume_from:327 logger.info(f"Resuming from {config.resume_from}")328 self.accelerator.load_state(config.resume_from)329 self.first_epoch = int(config.resume_from.split("_")[-1]) + 1330 else:331 self.first_epoch = 0332 333 def compute_rewards(self, prompt_image_pairs):334 reward, reward_metadata = self.reward_fn(335 prompt_image_pairs["images"], prompt_image_pairs["prompts"], prompt_image_pairs["prompt_metadata"]336 )337 return reward338 339 def step(self, epoch: int, global_step: int):340 """341 Perform a single step of training.342 343 Args:344 epoch (int): The current epoch.345 global_step (int): The current global step.346 347 Side Effects:348 - Model weights are updated349 - Logs the statistics to the accelerator trackers.350 - If `self.image_samples_callback` is not None, it will be called with the prompt_image_pairs, global_step, and the accelerator tracker.351 352 Returns:353 global_step (int): The updated global step.354 """355 info = defaultdict(list)356 357 self.sd_pipeline.unet.train()358 359 for _ in range(self.config.train_gradient_accumulation_steps):360 with self.accelerator.accumulate(self.sd_pipeline.unet), self.autocast(), torch.enable_grad():361 prompt_image_pairs = self._generate_samples(362 batch_size=self.config.train_batch_size,363 )364 365 rewards = self.compute_rewards(prompt_image_pairs)366 367 prompt_image_pairs["rewards"] = rewards368 369 rewards_vis = self.accelerator.gather(rewards).detach().cpu().numpy()370 371 loss = self.calculate_loss(rewards)372 373 self.accelerator.backward(loss)374 375 if self.accelerator.sync_gradients:376 self.accelerator.clip_grad_norm_(377 self.trainable_layers.parameters()378 if not isinstance(self.trainable_layers, list)379 else self.trainable_layers,380 self.config.train_max_grad_norm,381 )382 383 self.optimizer.step()384 self.optimizer.zero_grad()385 386 info["reward_mean"].append(rewards_vis.mean())387 info["reward_std"].append(rewards_vis.std())388 info["loss"].append(loss.item())389 390 # Checks if the accelerator has performed an optimization step behind the scenes391 if self.accelerator.sync_gradients:392 # log training-related stuff393 info = {k: torch.mean(torch.tensor(v)) for k, v in info.items()}394 info = self.accelerator.reduce(info, reduction="mean")395 info.update({"epoch": epoch})396 self.accelerator.log(info, step=global_step)397 global_step += 1398 info = defaultdict(list)399 else:400 raise ValueError(401 "Optimization step should have been performed by this point. Please check calculated gradient accumulation settings."402 )403 # Logs generated images404 if self.image_samples_callback is not None and global_step % self.config.log_image_freq == 0:405 self.image_samples_callback(prompt_image_pairs, global_step, self.accelerator.trackers[0])406 407 if epoch != 0 and epoch % self.config.save_freq == 0 and self.accelerator.is_main_process:408 self.accelerator.save_state()409 410 return global_step411 412 def calculate_loss(self, rewards):413 """414 Calculate the loss for a batch of an unpacked sample415 416 Args:417 rewards (torch.Tensor):418 Differentiable reward scalars for each generated image, shape: [batch_size]419 420 Returns:421 loss (torch.Tensor)422 (all of these are of shape (1,))423 """424 # Loss is specific to Aesthetic Reward function used in AlignProp (https://huggingface.co/papers/2310.03739)425 loss = 10.0 - (rewards).mean()426 return loss427 428 def loss(429 self,430 advantages: torch.Tensor,431 clip_range: float,432 ratio: torch.Tensor,433 ):434 unclipped_loss = -advantages * ratio435 clipped_loss = -advantages * torch.clamp(436 ratio,437 1.0 - clip_range,438 1.0 + clip_range,439 )440 return torch.mean(torch.maximum(unclipped_loss, clipped_loss))441 442 def _setup_optimizer(self, trainable_layers_parameters):443 if self.config.train_use_8bit_adam:444 import bitsandbytes445 446 optimizer_cls = bitsandbytes.optim.AdamW8bit447 else:448 optimizer_cls = torch.optim.AdamW449 450 return optimizer_cls(451 trainable_layers_parameters,452 lr=self.config.train_learning_rate,453 betas=(self.config.train_adam_beta1, self.config.train_adam_beta2),454 weight_decay=self.config.train_adam_weight_decay,455 eps=self.config.train_adam_epsilon,456 )457 458 def _save_model_hook(self, models, weights, output_dir):459 self.sd_pipeline.save_checkpoint(models, weights, output_dir)460 weights.pop() # ensures that accelerate doesn't try to handle saving of the model461 462 def _load_model_hook(self, models, input_dir):463 self.sd_pipeline.load_checkpoint(models, input_dir)464 models.pop() # ensures that accelerate doesn't try to handle loading of the model465 466 def _generate_samples(self, batch_size, with_grad=True, prompts=None):467 """468 Generate samples from the model469 470 Args:471 batch_size (int): Batch size to use for sampling472 with_grad (bool): Whether the generated RGBs should have gradients attached to it.473 474 Returns:475 prompt_image_pairs (dict[Any])476 """477 prompt_image_pairs = {}478 479 sample_neg_prompt_embeds = self.neg_prompt_embed.repeat(batch_size, 1, 1)480 481 if prompts is None:482 prompts, prompt_metadata = zip(*[self.prompt_fn() for _ in range(batch_size)])483 else:484 prompt_metadata = [{} for _ in range(batch_size)]485 486 prompt_ids = self.sd_pipeline.tokenizer(487 prompts,488 return_tensors="pt",489 padding="max_length",490 truncation=True,491 max_length=self.sd_pipeline.tokenizer.model_max_length,492 ).input_ids.to(self.accelerator.device)493 494 prompt_embeds = self.sd_pipeline.text_encoder(prompt_ids)[0]495 496 if with_grad:497 sd_output = self.sd_pipeline.rgb_with_grad(498 prompt_embeds=prompt_embeds,499 negative_prompt_embeds=sample_neg_prompt_embeds,500 num_inference_steps=self.config.sample_num_steps,501 guidance_scale=self.config.sample_guidance_scale,502 eta=self.config.sample_eta,503 truncated_backprop_rand=self.config.truncated_backprop_rand,504 truncated_backprop_timestep=self.config.truncated_backprop_timestep,505 truncated_rand_backprop_minmax=self.config.truncated_rand_backprop_minmax,506 output_type="pt",507 )508 else:509 sd_output = self.sd_pipeline(510 prompt_embeds=prompt_embeds,511 negative_prompt_embeds=sample_neg_prompt_embeds,512 num_inference_steps=self.config.sample_num_steps,513 guidance_scale=self.config.sample_guidance_scale,514 eta=self.config.sample_eta,515 output_type="pt",516 )517 518 images = sd_output.images519 520 prompt_image_pairs["images"] = images521 prompt_image_pairs["prompts"] = prompts522 prompt_image_pairs["prompt_metadata"] = prompt_metadata523 524 return prompt_image_pairs525 526 def train(self, epochs: Optional[int] = None):527 """528 Train the model for a given number of epochs529 """530 global_step = 0531 if epochs is None:532 epochs = self.config.num_epochs533 for epoch in range(self.first_epoch, epochs):534 global_step = self.step(epoch, global_step)535 536 def _save_pretrained(self, save_directory):537 self.sd_pipeline.save_pretrained(save_directory)538 self.create_model_card()539 540 def create_model_card(541 self,542 model_name: Optional[str] = None,543 dataset_name: Optional[str] = None,544 tags: Union[str, list[str], None] = None,545 ):546 """547 Creates a draft of a model card using the information available to the `Trainer`.548 549 Args:550 model_name (`str` or `None`, *optional*, defaults to `None`):551 Name of the model.552 dataset_name (`str` or `None`, *optional*, defaults to `None`):553 Name of the dataset used for training.554 tags (`str`, `list[str]` or `None`, *optional*, defaults to `None`):555 Tags to be associated with the model card.556 """557 if not self.is_world_process_zero():558 return559 560 if hasattr(self.model.config, "_name_or_path") and not os.path.isdir(self.model.config._name_or_path):561 base_model = self.model.config._name_or_path562 else:563 base_model = None564 565 tags = tags or []566 if isinstance(tags, str):567 tags = [tags]568 569 if hasattr(self.model.config, "unsloth_version"):570 tags.append("unsloth")571 572 citation = textwrap.dedent("""\573 @article{prabhudesai2024aligning,574 title = {{Aligning Text-to-Image Diffusion Models with Reward Backpropagation}},575 author = {Mihir Prabhudesai and Anirudh Goyal and Deepak Pathak and Katerina Fragkiadaki},576 year = 2024,577 eprint = {arXiv:2310.03739}578 }""")579 580 model_card = generate_model_card(581 base_model=base_model,582 model_name=model_name,583 hub_model_id=self.hub_model_id,584 dataset_name=dataset_name,585 tags=tags,586 wandb_url=wandb.run.get_url() if is_wandb_available() and wandb.run is not None else None,587 comet_url=get_comet_experiment_url(),588 trainer_name="AlignProp",589 trainer_citation=citation,590 paper_title="Aligning Text-to-Image Diffusion Models with Reward Backpropagation",591 paper_id="2310.03739",592 )593 594 model_card.save(os.path.join(self.args.output_dir, "README.md"))595class UnslothAlignPropTrainer(_UnslothAlignPropTrainer):596 """597 598 The AlignPropTrainer uses Deep Diffusion Policy Optimization to optimise diffusion models.599 Note, this trainer is heavily inspired by the work here: https://github.com/mihirp1998/AlignProp/600 As of now only Stable Diffusion based pipelines are supported601 602 Attributes:603 config (`AlignPropConfig`):604 Configuration object for AlignPropTrainer. Check the documentation of `PPOConfig` for more details.605 reward_function (`Callable[[torch.Tensor, tuple[str], tuple[Any]], torch.Tensor]`):606 Reward function to be used607 prompt_function (`Callable[[], tuple[str, Any]]`):608 Function to generate prompts to guide model609 sd_pipeline (`DDPOStableDiffusionPipeline`):610 Stable Diffusion pipeline to be used for training.611 image_samples_hook (`Optional[Callable[[Any, Any, Any], Any]]`):612 Hook to be called to log images613 614 """615 def __init__(616 self,617 config,618 reward_function,619 prompt_function,620 sd_pipeline,621 image_samples_hook = None,622 **kwargs623 ):624 if args is None: args = UnslothAlignPropConfig()625 other_metrics = []626 627 from unsloth_zoo.logging_utils import PatchRLStatistics628 PatchRLStatistics('alignprop_trainer', other_metrics)629 630 super().__init__(631 config = config,632 reward_function = reward_function,633 prompt_function = prompt_function,634 sd_pipeline = sd_pipeline,635 image_samples_hook = image_samples_hook,**kwargs)636 637pass638 