diff --git a/__init__.py b/__init__.py index 6a599e1..9b7aef5 100644 --- a/__init__.py +++ b/__init__.py @@ -11,6 +11,8 @@ from .tuneavideo.util import ddim_inversion import comfy.utils import folder_paths from einops import rearrange +from .train_tuneavideo import train +import copy def common_ksampler(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent, denoise=1.0, disable_noise=False, start_step=None, last_step=None, force_full_denoise=False): latent_image = latent["samples"] @@ -299,6 +301,31 @@ class DdimInversionSequence: return (s,) +class TrainUnetSequence: + @classmethod + def INPUT_TYPES(s): + return {"required": {"samples": ("LATENT",), + "model": ("MODEL",), + "context": ("CONDITIONING",), + "steps": ("INT", {"default": 20, "min": 0, "max": 10000}), + }} + + RETURN_TYPES = ("MODEL",) + FUNCTION = "train_unet" + + CATEGORY = "sampling" + + def train_unet(self, samples, model, context, steps): + device = model_management.get_torch_device() + noise_scheduler = convert_scheduler_checkpoint(model) + samples = rearrange(samples["samples"], "f c h w -> c f h w") + with torch.inference_mode(mode=False): + model_train = train(copy.deepcopy(model), noise_scheduler, samples, context[0][0].squeeze(0), device, max_train_steps=steps) + if model_management.should_use_fp16(): + model_train.model = model_train.model.half() + return (model_train,) + + NODE_CLASS_MAPPINGS = { "LoadImageSequence": LoadImageSequence, "VAEEncodeForInpaintSequence": VAEEncodeForInpaintSequence, @@ -307,4 +334,5 @@ NODE_CLASS_MAPPINGS = { "CheckpointLoaderSimpleSequence": CheckpointLoaderSimpleSequence, "SetLatentNoiseSequence": SetLatentNoiseSequence, "DdimInversionSequence": DdimInversionSequence, + "TrainUnetSequence": TrainUnetSequence, } diff --git a/sd.py b/sd.py index 332f4d6..1d2a9e5 100644 --- a/sd.py +++ b/sd.py @@ -102,12 +102,13 @@ def load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, e model = instantiate_from_config(model_config) model = load_model_weights(model, sd, verbose=False, load_state_dict_to=load_state_dict_to) - model.model.diffusion_model = convert_unet_checkpoint(sd, OmegaConf.create({"model": model_config})) + with torch.inference_mode(mode=False): + model.model.diffusion_model = convert_unet_checkpoint(sd, OmegaConf.create({"model": model_config})) if model_management.xformers_enabled(): model.model.diffusion_model.enable_xformers_memory_efficient_attention() - if fp16: - model = model.half() + #if fp16: + # model = model.half() return (ModelPatcher(model), clip, vae) diff --git a/train_tuneavideo.py b/train_tuneavideo.py index 2d174ff..88e4cb0 100644 --- a/train_tuneavideo.py +++ b/train_tuneavideo.py @@ -1,29 +1,20 @@ -import argparse import inspect import math -import os -from typing import Dict, Optional, Tuple -from omegaconf import OmegaConf +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.logging import get_logger from accelerate.utils import set_seed -from diffusers import AutoencoderKL, DDPMScheduler from diffusers.optimization import get_scheduler from diffusers.utils import check_min_version -from diffusers.utils.import_utils import is_xformers_available from tqdm.auto import tqdm -from transformers import CLIPTextModel, CLIPTokenizer -from tuneavideo.models.unet import UNet3DConditionModel -from tuneavideo.data.dataset import TuneAVideoDataset -from tuneavideo.pipelines.pipeline_tuneavideo import TuneAVideoPipeline from einops import rearrange -import shutil # Will error if the minimal version of diffusers is not installed. Remove at your own risks. @@ -31,13 +22,27 @@ check_min_version("0.10.0.dev0") logger = get_logger(__name__, log_level="INFO") +class TuneAVideoDataset(Dataset): + def __init__(self): + self.prompt_ids = None + self.pixel_values = None + def __len__(self): + return 1 + def __getitem__(self, index): -def main( - pretrained_model_path: str, - pretrained_vae_path: str, - output_dir: str, - train_data: Dict, - inference_data: Dict = None, + 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", @@ -56,13 +61,8 @@ def main( max_grad_norm: float = 1.0, gradient_accumulation_steps: int = 1, gradient_checkpointing: bool = True, - checkpointing_steps: int = 5000, - resume_from_checkpoint: Optional[str] = None, mixed_precision: Optional[str] = "fp16", - use_8bit_adam: bool = False, - enable_xformers_memory_efficient_attention: bool = True, seed: Optional[int] = None, - gpu_id: str = "0", ): *_, config = inspect.getargvalues(inspect.currentframe()) @@ -76,34 +76,13 @@ def main( if seed is not None: set_seed(seed) - # Handle the output folder creation - if accelerator.is_main_process: - os.makedirs(output_dir, exist_ok=True) - OmegaConf.save(config, os.path.join(output_dir, 'config.yaml')) - - # Load scheduler, tokenizer and models. - noise_scheduler = DDPMScheduler.from_pretrained(pretrained_model_path, subfolder="scheduler") - tokenizer = CLIPTokenizer.from_pretrained(pretrained_model_path, subfolder="tokenizer") - text_encoder = CLIPTextModel.from_pretrained(pretrained_model_path, subfolder="text_encoder") - vae = AutoencoderKL.from_pretrained(pretrained_vae_path, subfolder="vae") - unet = UNet3DConditionModel.from_pretrained_2d(pretrained_model_path, subfolder="unet").to(f"cuda:{gpu_id}") - - # Freeze vae and text_encoder - vae.requires_grad_(False) - text_encoder.requires_grad_(False) - + 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 enable_xformers_memory_efficient_attention: - if is_xformers_available(): - unet.enable_xformers_memory_efficient_attention() - else: - raise ValueError("xformers is not available. Make sure it is installed correctly") - if gradient_checkpointing: unet.enable_gradient_checkpointing() @@ -112,18 +91,7 @@ def main( learning_rate * gradient_accumulation_steps * train_batch_size * accelerator.num_processes ) - # Initialize the optimizer - if use_8bit_adam: - try: - import bitsandbytes as bnb - except ImportError: - raise ImportError( - "Please install bitsandbytes to use 8-bit Adam. You can do so by running `pip install bitsandbytes`" - ) - - optimizer_cls = bnb.optim.AdamW8bit - else: - optimizer_cls = torch.optim.AdamW + optimizer_cls = torch.optim.AdamW optimizer = optimizer_cls( unet.parameters(), @@ -134,12 +102,11 @@ def main( ) # Get the training dataset - train_dataset = TuneAVideoDataset(**train_data) + train_dataset = TuneAVideoDataset() # Preprocessing the dataset - train_dataset.prompt_ids = tokenizer( - train_dataset.prompt, max_length=tokenizer.model_max_length, padding="max_length", truncation=True, return_tensors="pt" - ).input_ids[0] + train_dataset.prompt_ids = context + train_dataset.pixel_values = samples # DataLoaders creation: train_dataloader = torch.utils.data.DataLoader( @@ -167,10 +134,6 @@ def main( elif accelerator.mixed_precision == "bf16": weight_dtype = torch.bfloat16 - # Move text_encode and vae to gpu and cast to weight_dtype - text_encoder.to(f"cuda:{gpu_id}", dtype=weight_dtype) - vae.to(f"cuda:{gpu_id}", dtype=weight_dtype) - # 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 @@ -185,23 +148,6 @@ def main( global_step = 0 first_epoch = 0 - # Potentially load in the weights and states from a previous save - if resume_from_checkpoint: - if resume_from_checkpoint != "latest": - path = os.path.basename(resume_from_checkpoint) - else: - # Get the most recent checkpoint - dirs = os.listdir(output_dir) - dirs = [d for d in dirs if d.startswith("checkpoint")] - dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) - path = dirs[-1] - accelerator.print(f"Resuming from checkpoint {path}") - accelerator.load_state(os.path.join(output_dir, path)) - global_step = int(path.split("-")[1]) - - first_epoch = global_step // num_update_steps_per_epoch - resume_step = global_step % num_update_steps_per_epoch - # 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") @@ -210,20 +156,10 @@ def main( unet.train() train_loss = 0.0 for step, batch in enumerate(train_dataloader): - # Skip steps until we reach the resumed step - if resume_from_checkpoint and epoch == first_epoch and step < resume_step: - if step % gradient_accumulation_steps == 0: - progress_bar.update(1) - continue with accelerator.accumulate(unet): # Convert videos to latent space - pixel_values = batch["pixel_values"].to(weight_dtype).to(f"cuda:{gpu_id}") - video_length = pixel_values.shape[1] - pixel_values = rearrange(pixel_values, "b f c h w -> (b f) c h w") - latents = vae.encode(pixel_values).latent_dist.sample() - latents = rearrange(latents, "(b f) c h w -> b c f h w", f=video_length) - latents = latents * 0.18215 + latents = batch["pixel_values"].type(weight_dtype).to(device) # Sample noise that we'll add to the latents noise = torch.randn_like(latents) @@ -237,7 +173,7 @@ def main( noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps) # Get the text embedding for conditioning - encoder_hidden_states = text_encoder(batch["prompt_ids"].to(f"cuda:{gpu_id}"))[0] + 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": @@ -248,7 +184,9 @@ def main( raise ValueError(f"Unknown prediction type {noise_scheduler.prediction_type}") # Predict the noise residual and compute loss - model_pred = unet(noisy_latents, timesteps, encoder_hidden_states).sample + 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). @@ -270,12 +208,6 @@ def main( accelerator.log({"train_loss": train_loss}, step=global_step) train_loss = 0.0 - if global_step % checkpointing_steps == 0: - if accelerator.is_main_process: - save_path = os.path.join(output_dir, f"checkpoint-{global_step}") - accelerator.save_state(save_path) - logger.info(f"Saved state to {save_path}") - logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} progress_bar.set_postfix(**logs) @@ -284,31 +216,7 @@ def main( # Create the pipeline using the trained modules and save it. accelerator.wait_for_everyone() - if accelerator.is_main_process: - unet = accelerator.unwrap_model(unet) - pipeline = TuneAVideoPipeline.from_pretrained( - pretrained_model_path, - text_encoder=text_encoder, - vae=vae, - unet=unet, - ) - pipeline.save_pretrained(output_dir) - + unet = accelerator.unwrap_model(unet) accelerator.end_training() - - # 删除冗余文件夹 - shutil.rmtree(os.path.join(output_dir, "scheduler")) - shutil.rmtree(os.path.join(output_dir, "text_encoder")) - shutil.rmtree(os.path.join(output_dir, "vae")) - shutil.rmtree(os.path.join(output_dir, "tokenizer")) - # 删除冗余文件 - os.remove(os.path.join(output_dir, "config.yaml")) - os.remove(os.path.join(output_dir, "model_index.json")) - - -if __name__ == "__main__": - parser = argparse.ArgumentParser() - parser.add_argument("--config", type=str, default="./configs/Continuous_frame_challenge.yaml") - parser.add_argument("--cuda", type=str, default="0") - args = parser.parse_args() - main(**OmegaConf.load(args.config), gpu_id=args.cuda) + model.model.model.diffusion_model = unet.to(torch.device("cpu")) + return model