first modification
This commit is contained in:
+28
@@ -11,6 +11,8 @@ from .tuneavideo.util import ddim_inversion
|
|||||||
import comfy.utils
|
import comfy.utils
|
||||||
import folder_paths
|
import folder_paths
|
||||||
from einops import rearrange
|
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):
|
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"]
|
latent_image = latent["samples"]
|
||||||
@@ -299,6 +301,31 @@ class DdimInversionSequence:
|
|||||||
return (s,)
|
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 = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"LoadImageSequence": LoadImageSequence,
|
"LoadImageSequence": LoadImageSequence,
|
||||||
"VAEEncodeForInpaintSequence": VAEEncodeForInpaintSequence,
|
"VAEEncodeForInpaintSequence": VAEEncodeForInpaintSequence,
|
||||||
@@ -307,4 +334,5 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
"CheckpointLoaderSimpleSequence": CheckpointLoaderSimpleSequence,
|
"CheckpointLoaderSimpleSequence": CheckpointLoaderSimpleSequence,
|
||||||
"SetLatentNoiseSequence": SetLatentNoiseSequence,
|
"SetLatentNoiseSequence": SetLatentNoiseSequence,
|
||||||
"DdimInversionSequence": DdimInversionSequence,
|
"DdimInversionSequence": DdimInversionSequence,
|
||||||
|
"TrainUnetSequence": TrainUnetSequence,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 = instantiate_from_config(model_config)
|
||||||
model = load_model_weights(model, sd, verbose=False, load_state_dict_to=load_state_dict_to)
|
model = load_model_weights(model, sd, verbose=False, load_state_dict_to=load_state_dict_to)
|
||||||
|
with torch.inference_mode(mode=False):
|
||||||
model.model.diffusion_model = convert_unet_checkpoint(sd, OmegaConf.create({"model": model_config}))
|
model.model.diffusion_model = convert_unet_checkpoint(sd, OmegaConf.create({"model": model_config}))
|
||||||
|
|
||||||
if model_management.xformers_enabled():
|
if model_management.xformers_enabled():
|
||||||
model.model.diffusion_model.enable_xformers_memory_efficient_attention()
|
model.model.diffusion_model.enable_xformers_memory_efficient_attention()
|
||||||
|
|
||||||
if fp16:
|
#if fp16:
|
||||||
model = model.half()
|
# model = model.half()
|
||||||
|
|
||||||
return (ModelPatcher(model), clip, vae)
|
return (ModelPatcher(model), clip, vae)
|
||||||
|
|||||||
+33
-125
@@ -1,29 +1,20 @@
|
|||||||
import argparse
|
|
||||||
import inspect
|
import inspect
|
||||||
import math
|
import math
|
||||||
import os
|
from typing import Optional, Tuple
|
||||||
from typing import Dict, Optional, Tuple
|
|
||||||
from omegaconf import OmegaConf
|
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
import torch.utils.checkpoint
|
import torch.utils.checkpoint
|
||||||
|
from torch.utils.data import Dataset
|
||||||
|
|
||||||
from accelerate import Accelerator
|
from accelerate import Accelerator
|
||||||
from accelerate.logging import get_logger
|
from accelerate.logging import get_logger
|
||||||
from accelerate.utils import set_seed
|
from accelerate.utils import set_seed
|
||||||
from diffusers import AutoencoderKL, DDPMScheduler
|
|
||||||
from diffusers.optimization import get_scheduler
|
from diffusers.optimization import get_scheduler
|
||||||
from diffusers.utils import check_min_version
|
from diffusers.utils import check_min_version
|
||||||
from diffusers.utils.import_utils import is_xformers_available
|
|
||||||
from tqdm.auto import tqdm
|
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
|
from einops import rearrange
|
||||||
import shutil
|
|
||||||
|
|
||||||
|
|
||||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
# 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")
|
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(
|
example = {
|
||||||
pretrained_model_path: str,
|
"pixel_values": self.pixel_values,
|
||||||
pretrained_vae_path: str,
|
"prompt_ids": self.prompt_ids
|
||||||
output_dir: str,
|
}
|
||||||
train_data: Dict,
|
|
||||||
inference_data: Dict = None,
|
return example
|
||||||
|
|
||||||
|
def train(
|
||||||
|
model,
|
||||||
|
noise_scheduler,
|
||||||
|
samples,
|
||||||
|
context,
|
||||||
|
device,
|
||||||
trainable_modules: Tuple[str] = (
|
trainable_modules: Tuple[str] = (
|
||||||
"attn1.to_q",
|
"attn1.to_q",
|
||||||
"attn2.to_q",
|
"attn2.to_q",
|
||||||
@@ -56,13 +61,8 @@ def main(
|
|||||||
max_grad_norm: float = 1.0,
|
max_grad_norm: float = 1.0,
|
||||||
gradient_accumulation_steps: int = 1,
|
gradient_accumulation_steps: int = 1,
|
||||||
gradient_checkpointing: bool = True,
|
gradient_checkpointing: bool = True,
|
||||||
checkpointing_steps: int = 5000,
|
|
||||||
resume_from_checkpoint: Optional[str] = None,
|
|
||||||
mixed_precision: Optional[str] = "fp16",
|
mixed_precision: Optional[str] = "fp16",
|
||||||
use_8bit_adam: bool = False,
|
|
||||||
enable_xformers_memory_efficient_attention: bool = True,
|
|
||||||
seed: Optional[int] = None,
|
seed: Optional[int] = None,
|
||||||
gpu_id: str = "0",
|
|
||||||
):
|
):
|
||||||
*_, config = inspect.getargvalues(inspect.currentframe())
|
*_, config = inspect.getargvalues(inspect.currentframe())
|
||||||
|
|
||||||
@@ -76,34 +76,13 @@ def main(
|
|||||||
if seed is not None:
|
if seed is not None:
|
||||||
set_seed(seed)
|
set_seed(seed)
|
||||||
|
|
||||||
# Handle the output folder creation
|
unet = model.model.model.diffusion_model.to(device)
|
||||||
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.requires_grad_(False)
|
unet.requires_grad_(False)
|
||||||
for name, module in unet.named_modules():
|
for name, module in unet.named_modules():
|
||||||
if name.endswith(tuple(trainable_modules)):
|
if name.endswith(tuple(trainable_modules)):
|
||||||
for params in module.parameters():
|
for params in module.parameters():
|
||||||
params.requires_grad = True
|
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:
|
if gradient_checkpointing:
|
||||||
unet.enable_gradient_checkpointing()
|
unet.enable_gradient_checkpointing()
|
||||||
|
|
||||||
@@ -112,17 +91,6 @@ def main(
|
|||||||
learning_rate * gradient_accumulation_steps * train_batch_size * accelerator.num_processes
|
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(
|
optimizer = optimizer_cls(
|
||||||
@@ -134,12 +102,11 @@ def main(
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Get the training dataset
|
# Get the training dataset
|
||||||
train_dataset = TuneAVideoDataset(**train_data)
|
train_dataset = TuneAVideoDataset()
|
||||||
|
|
||||||
# Preprocessing the dataset
|
# Preprocessing the dataset
|
||||||
train_dataset.prompt_ids = tokenizer(
|
train_dataset.prompt_ids = context
|
||||||
train_dataset.prompt, max_length=tokenizer.model_max_length, padding="max_length", truncation=True, return_tensors="pt"
|
train_dataset.pixel_values = samples
|
||||||
).input_ids[0]
|
|
||||||
|
|
||||||
# DataLoaders creation:
|
# DataLoaders creation:
|
||||||
train_dataloader = torch.utils.data.DataLoader(
|
train_dataloader = torch.utils.data.DataLoader(
|
||||||
@@ -167,10 +134,6 @@ def main(
|
|||||||
elif accelerator.mixed_precision == "bf16":
|
elif accelerator.mixed_precision == "bf16":
|
||||||
weight_dtype = torch.bfloat16
|
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.
|
# 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)
|
num_update_steps_per_epoch = math.ceil(len(train_dataloader) / gradient_accumulation_steps)
|
||||||
# Afterwards we recalculate our number of training epochs
|
# Afterwards we recalculate our number of training epochs
|
||||||
@@ -185,23 +148,6 @@ def main(
|
|||||||
global_step = 0
|
global_step = 0
|
||||||
first_epoch = 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.
|
# 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 = tqdm(range(global_step, max_train_steps), disable=not accelerator.is_local_main_process)
|
||||||
progress_bar.set_description("Steps")
|
progress_bar.set_description("Steps")
|
||||||
@@ -210,20 +156,10 @@ def main(
|
|||||||
unet.train()
|
unet.train()
|
||||||
train_loss = 0.0
|
train_loss = 0.0
|
||||||
for step, batch in enumerate(train_dataloader):
|
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):
|
with accelerator.accumulate(unet):
|
||||||
# Convert videos to latent space
|
# Convert videos to latent space
|
||||||
pixel_values = batch["pixel_values"].to(weight_dtype).to(f"cuda:{gpu_id}")
|
latents = batch["pixel_values"].type(weight_dtype).to(device)
|
||||||
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
|
|
||||||
|
|
||||||
# Sample noise that we'll add to the latents
|
# Sample noise that we'll add to the latents
|
||||||
noise = torch.randn_like(latents)
|
noise = torch.randn_like(latents)
|
||||||
@@ -237,7 +173,7 @@ def main(
|
|||||||
noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps)
|
noisy_latents = noise_scheduler.add_noise(latents, noise, timesteps)
|
||||||
|
|
||||||
# Get the text embedding for conditioning
|
# 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
|
# Get the target for loss depending on the prediction type
|
||||||
if noise_scheduler.prediction_type == "epsilon":
|
if noise_scheduler.prediction_type == "epsilon":
|
||||||
@@ -248,7 +184,9 @@ def main(
|
|||||||
raise ValueError(f"Unknown prediction type {noise_scheduler.prediction_type}")
|
raise ValueError(f"Unknown prediction type {noise_scheduler.prediction_type}")
|
||||||
|
|
||||||
# Predict the noise residual and compute loss
|
# 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")
|
loss = F.mse_loss(model_pred.float(), target.float(), reduction="mean")
|
||||||
|
|
||||||
# Gather the losses across all processes for logging (if we use distributed training).
|
# 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)
|
accelerator.log({"train_loss": train_loss}, step=global_step)
|
||||||
train_loss = 0.0
|
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]}
|
logs = {"step_loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]}
|
||||||
progress_bar.set_postfix(**logs)
|
progress_bar.set_postfix(**logs)
|
||||||
|
|
||||||
@@ -284,31 +216,7 @@ def main(
|
|||||||
|
|
||||||
# Create the pipeline using the trained modules and save it.
|
# Create the pipeline using the trained modules and save it.
|
||||||
accelerator.wait_for_everyone()
|
accelerator.wait_for_everyone()
|
||||||
if accelerator.is_main_process:
|
|
||||||
unet = accelerator.unwrap_model(unet)
|
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)
|
|
||||||
|
|
||||||
accelerator.end_training()
|
accelerator.end_training()
|
||||||
|
model.model.model.diffusion_model = unet.to(torch.device("cpu"))
|
||||||
# 删除冗余文件夹
|
return model
|
||||||
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)
|
|
||||||
|
|||||||
Reference in New Issue
Block a user