first modification

This commit is contained in:
sylym
2023-03-28 00:34:18 +08:00
parent 686791ef70
commit 8a005a02ee
3 changed files with 67 additions and 130 deletions
+28
View File
@@ -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,
} }
+3 -2
View File
@@ -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
View File
@@ -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)