import inspect import math from typing import Optional, Tuple import torch import torch.nn.functional as F import torch.utils.checkpoint from torch.utils.data import Dataset from accelerate import Accelerator from accelerate.utils import set_seed from diffusers.optimization import get_scheduler from diffusers.utils import check_min_version from tqdm.auto import tqdm from einops import rearrange # Will error if the minimal version of diffusers is not installed. Remove at your own risks. check_min_version("0.10.0.dev0") class TuneAVideoDataset(Dataset): def __init__(self): self.prompt_ids = None self.pixel_values = None def __len__(self): return 1 def __getitem__(self, index): example = { "pixel_values": self.pixel_values, "prompt_ids": self.prompt_ids } return example def train( model, noise_scheduler, samples, context, device, trainable_modules: Tuple[str] = ( "attn1.to_q", "attn2.to_q", "attn_temp", ), train_batch_size: int = 1, max_train_steps: int = 500, learning_rate: float = 3e-5, scale_lr: bool = False, lr_scheduler: str = "constant", lr_warmup_steps: int = 0, adam_beta1: float = 0.9, adam_beta2: float = 0.999, adam_weight_decay: float = 1e-2, adam_epsilon: float = 1e-08, max_grad_norm: float = 1.0, gradient_accumulation_steps: int = 1, gradient_checkpointing: bool = True, mixed_precision: Optional[str] = "fp16", seed: Optional[int] = None, ): *_, config = inspect.getargvalues(inspect.currentframe()) accelerator = Accelerator( gradient_accumulation_steps=gradient_accumulation_steps, mixed_precision=mixed_precision, device_placement=False ) # If passed along, set the training seed now. if seed is not None: set_seed(seed) unet = model.model.model.diffusion_model.to(device) unet.requires_grad_(False) for name, module in unet.named_modules(): if name.endswith(tuple(trainable_modules)): for params in module.parameters(): params.requires_grad = True if gradient_checkpointing: unet.enable_gradient_checkpointing() if scale_lr: learning_rate = ( learning_rate * gradient_accumulation_steps * train_batch_size * accelerator.num_processes ) optimizer_cls = torch.optim.AdamW optimizer = optimizer_cls( unet.parameters(), lr=learning_rate, betas=(adam_beta1, adam_beta2), weight_decay=adam_weight_decay, eps=adam_epsilon, ) # Get the training dataset train_dataset = TuneAVideoDataset() # Preprocessing the dataset train_dataset.prompt_ids = context train_dataset.pixel_values = samples # DataLoaders creation: train_dataloader = torch.utils.data.DataLoader( train_dataset, batch_size=train_batch_size ) # Scheduler lr_scheduler = get_scheduler( lr_scheduler, optimizer=optimizer, num_warmup_steps=lr_warmup_steps * gradient_accumulation_steps, num_training_steps=max_train_steps * gradient_accumulation_steps, ) # Prepare everything with our `accelerator`. unet, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( unet, optimizer, train_dataloader, lr_scheduler ) # For mixed precision training we cast the text_encoder and vae weights to half-precision # as these models are only used for inference, keeping weights in full precision is not required. weight_dtype = torch.float32 if accelerator.mixed_precision == "fp16": weight_dtype = torch.float16 elif accelerator.mixed_precision == "bf16": weight_dtype = torch.bfloat16 # We need to recalculate our total training steps as the size of the training dataloader may have changed. num_update_steps_per_epoch = math.ceil(len(train_dataloader) / gradient_accumulation_steps) # Afterwards we recalculate our number of training epochs num_train_epochs = math.ceil(max_train_steps / num_update_steps_per_epoch) # We need to initialize the trackers we use, and also store our configuration. # The trackers initializes automatically on the main process. if accelerator.is_main_process: accelerator.init_trackers("text2video-fine-tune") # Train! global_step = 0 first_epoch = 0 # Only show the progress bar once on each machine. progress_bar = tqdm(range(global_step, max_train_steps), disable=not accelerator.is_local_main_process) progress_bar.set_description("Steps") for epoch in range(first_epoch, num_train_epochs): unet.train() train_loss = 0.0 for step, batch in enumerate(train_dataloader): with accelerator.accumulate(unet): # Convert videos to latent space latents = batch["pixel_values"].type(weight_dtype).to(device) # Sample noise that we'll add to the latents noise = torch.randn_like(latents) bsz = latents.shape[0] # Sample a random timestep for each video timesteps = torch.randint(0, noise_scheduler.num_train_timesteps, (bsz,), device=latents.device) timesteps = timesteps.long() # Add noise to the latents according to the noise magnitude at each timestep # (this is the forward diffusion process) noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps) # Get the text embedding for conditioning encoder_hidden_states = batch["prompt_ids"].type(weight_dtype).to(device) # Get the target for loss depending on the prediction type if noise_scheduler.prediction_type == "epsilon": target = noise elif noise_scheduler.prediction_type == "v_prediction": target = noise_scheduler.get_velocity(latents, noise, timesteps) else: raise ValueError(f"Unknown prediction type {noise_scheduler.prediction_type}") # Predict the noise residual and compute loss noisy_latents = rearrange(noisy_latents.squeeze(0), "c f h w -> f c h w") model_pred = unet(noisy_latents, timesteps, encoder_hidden_states) model_pred = rearrange(model_pred.unsqueeze(0), "b f c h w -> b c f h w") loss = F.mse_loss(model_pred.float(), target.float(), reduction="mean") # Gather the losses across all processes for logging (if we use distributed training). avg_loss = accelerator.gather(loss.repeat(train_batch_size)).mean() train_loss += avg_loss.item() / gradient_accumulation_steps # Backpropagate accelerator.backward(loss) if accelerator.sync_gradients: accelerator.clip_grad_norm_(unet.parameters(), max_grad_norm) optimizer.step() lr_scheduler.step() optimizer.zero_grad() # Checks if the accelerator has performed an optimization step behind the scenes if accelerator.sync_gradients: progress_bar.update(1) global_step += 1 accelerator.log({"train_loss": train_loss}, step=global_step) train_loss = 0.0 logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) if global_step >= max_train_steps: break # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() unet = accelerator.unwrap_model(unet) accelerator.end_training() model.model.model.diffusion_model = unet.to(torch.device("cpu")) return model