220 lines
8.0 KiB
Python
220 lines
8.0 KiB
Python
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
|