Initial commit
This commit is contained in:
+26
@@ -0,0 +1,26 @@
|
|||||||
|
wandb/
|
||||||
|
*debug*
|
||||||
|
debugs/
|
||||||
|
outputs/
|
||||||
|
samples/
|
||||||
|
__pycache__/
|
||||||
|
ossutil_output/
|
||||||
|
.ossutil_checkpoint/
|
||||||
|
|
||||||
|
scripts/*
|
||||||
|
!scripts/animate.py
|
||||||
|
|
||||||
|
*.ipynb
|
||||||
|
*.safetensors
|
||||||
|
*.ckpt
|
||||||
|
|
||||||
|
models/*
|
||||||
|
!models/StableDiffusion/
|
||||||
|
models/StableDiffusion/*
|
||||||
|
!models/StableDiffusion/*.txt
|
||||||
|
!models/Motion_Module/
|
||||||
|
!models/Motion_Module/*.txt
|
||||||
|
!models/DreamBooth_LoRA/
|
||||||
|
!models/DreamBooth_LoRA/*.txt
|
||||||
|
!models/MotionLoRA/
|
||||||
|
!models/MotionLoRA/*.txt
|
||||||
@@ -0,0 +1,37 @@
|
|||||||
|
#HEAVILY WORK IN PROGRESS
|
||||||
|
|
||||||
|
[ComfyUI](https://github.com/comfyanonymous/ComfyUI) custom nodes for using [AnimateDiff-MotionDirector](https://github.com/ExponentialML/AnimateDiff-MotionDirector)
|
||||||
|
|
||||||
|
After training, the LoRAs are intended to be used with the ComfyUI Extension [ComfyUI-AnimateDiff-Evolved](https://github.com/Kosinkadink/ComfyUI-AnimateDiff-Evolved).
|
||||||
|
|
||||||
|
|
||||||
|
## BibTeX
|
||||||
|
|
||||||
|
```
|
||||||
|
@article{guo2023animatediff,
|
||||||
|
title={AnimateDiff: Animate Your Personalized Text-to-Image Diffusion Models without Specific Tuning},
|
||||||
|
author={Guo, Yuwei and Yang, Ceyuan and Rao, Anyi and Wang, Yaohui and Qiao, Yu and Lin, Dahua and Dai, Bo},
|
||||||
|
journal={arXiv preprint arXiv:2307.04725},
|
||||||
|
year={2023}
|
||||||
|
}
|
||||||
|
@article{zhao2023motiondirector,
|
||||||
|
title={MotionDirector: Motion Customization of Text-to-Video Diffusion Models},
|
||||||
|
author={Zhao, Rui and Gu, Yuchao and Wu, Jay Zhangjie and Zhang, David Junhao and Liu, Jiawei and Wu, Weijia and Keppo, Jussi and Shou, Mike Zheng},
|
||||||
|
journal={arXiv preprint arXiv:2310.08465},
|
||||||
|
year={2023}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
## Disclaimer
|
||||||
|
This project is released for academic use and creative usage. We disclaim responsibility for user-generated content. Users are solely liable for their actions. The project contributors are not legally affiliated with, nor accountable for, users' behaviors. Use the generative model responsibly, adhering to ethical and legal standards.
|
||||||
|
|
||||||
|
## Acknowledgements
|
||||||
|
Codebase built upon:
|
||||||
|
- [AnimateDiff-MotionDirector](https://github.com/ExponentialML/AnimateDiff-MotionDirector)
|
||||||
|
- [AnimateDiff](https://github.com/guoyww/AnimateDiff)
|
||||||
|
- [Tune-a-Video](https://github.com/showlab/Tune-A-Video).
|
||||||
|
- [MotionDirector](https://github.com/showlab/MotionDirector)
|
||||||
|
- [Text-To-Video-Finetuning](https://github.com/ExponentialML/Text-To-Video-Finetuning)
|
||||||
|
- [lora](https://github.com/cloneofsimo/lora)
|
||||||
|
- [ComfyUI](https://github.com/comfyanonymous/ComfyUI)
|
||||||
|
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||||
|
|
||||||
|
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||||
@@ -0,0 +1,98 @@
|
|||||||
|
import os, io, csv, math, random
|
||||||
|
import numpy as np
|
||||||
|
from einops import rearrange
|
||||||
|
from decord import VideoReader
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torchvision.transforms as transforms
|
||||||
|
from torch.utils.data.dataset import Dataset
|
||||||
|
from animatediff.utils.util import zero_rank_print
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
class WebVid10M(Dataset):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
csv_path, video_folder,
|
||||||
|
sample_size=256, sample_stride=4, sample_n_frames=16,
|
||||||
|
is_image=False,
|
||||||
|
):
|
||||||
|
zero_rank_print(f"loading annotations from {csv_path} ...")
|
||||||
|
with open(csv_path, 'r') as csvfile:
|
||||||
|
self.dataset = list(csv.DictReader(csvfile))
|
||||||
|
self.length = len(self.dataset)
|
||||||
|
zero_rank_print(f"data scale: {self.length}")
|
||||||
|
|
||||||
|
self.video_folder = video_folder
|
||||||
|
self.sample_stride = sample_stride
|
||||||
|
self.sample_n_frames = sample_n_frames
|
||||||
|
self.is_image = is_image
|
||||||
|
|
||||||
|
sample_size = tuple(sample_size) if not isinstance(sample_size, int) else (sample_size, sample_size)
|
||||||
|
self.pixel_transforms = transforms.Compose([
|
||||||
|
transforms.RandomHorizontalFlip(),
|
||||||
|
transforms.Resize(sample_size[0]),
|
||||||
|
transforms.CenterCrop(sample_size),
|
||||||
|
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
|
||||||
|
])
|
||||||
|
|
||||||
|
def get_batch(self, idx):
|
||||||
|
video_dict = self.dataset[idx]
|
||||||
|
videoid, name, page_dir = video_dict['videoid'], video_dict['name'], video_dict['page_dir']
|
||||||
|
|
||||||
|
video_dir = os.path.join(self.video_folder, f"{videoid}.mp4")
|
||||||
|
video_reader = VideoReader(video_dir)
|
||||||
|
video_length = len(video_reader)
|
||||||
|
|
||||||
|
if not self.is_image:
|
||||||
|
clip_length = min(video_length, (self.sample_n_frames - 1) * self.sample_stride + 1)
|
||||||
|
start_idx = random.randint(0, video_length - clip_length)
|
||||||
|
batch_index = np.linspace(start_idx, start_idx + clip_length - 1, self.sample_n_frames, dtype=int)
|
||||||
|
else:
|
||||||
|
batch_index = [random.randint(0, video_length - 1)]
|
||||||
|
|
||||||
|
pixel_values = torch.from_numpy(video_reader.get_batch(batch_index).asnumpy()).permute(0, 3, 1, 2).contiguous()
|
||||||
|
pixel_values = pixel_values / 255.
|
||||||
|
del video_reader
|
||||||
|
|
||||||
|
if self.is_image:
|
||||||
|
pixel_values = pixel_values[0]
|
||||||
|
|
||||||
|
return pixel_values, name
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return self.length
|
||||||
|
|
||||||
|
def __getitem__(self, idx):
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
pixel_values, name = self.get_batch(idx)
|
||||||
|
break
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
idx = random.randint(0, self.length-1)
|
||||||
|
|
||||||
|
pixel_values = self.pixel_transforms(pixel_values)
|
||||||
|
sample = dict(pixel_values=pixel_values, text=name)
|
||||||
|
return sample
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
from animatediff.utils.util import save_videos_grid
|
||||||
|
|
||||||
|
dataset = WebVid10M(
|
||||||
|
csv_path="/mnt/petrelfs/guoyuwei/projects/datasets/webvid/results_2M_val.csv",
|
||||||
|
video_folder="/mnt/petrelfs/guoyuwei/projects/datasets/webvid/2M_val",
|
||||||
|
sample_size=256,
|
||||||
|
sample_stride=4, sample_n_frames=16,
|
||||||
|
is_image=True,
|
||||||
|
)
|
||||||
|
import pdb
|
||||||
|
pdb.set_trace()
|
||||||
|
|
||||||
|
dataloader = torch.utils.data.DataLoader(dataset, batch_size=4, num_workers=16,)
|
||||||
|
for idx, batch in enumerate(dataloader):
|
||||||
|
print(batch["pixel_values"].shape, len(batch["text"]))
|
||||||
|
# for i in range(batch["pixel_values"].shape[0]):
|
||||||
|
# save_videos_grid(batch["pixel_values"][i:i+1].permute(0,2,1,3,4), os.path.join(".", f"{idx}-{i}.mp4"), rescale=True)
|
||||||
@@ -0,0 +1,300 @@
|
|||||||
|
# Adapted from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention.py
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||||
|
from diffusers import ModelMixin
|
||||||
|
from diffusers.utils import BaseOutput
|
||||||
|
from diffusers.utils.import_utils import is_xformers_available
|
||||||
|
from diffusers.models.attention import Attention, FeedForward, AdaLayerNorm
|
||||||
|
|
||||||
|
from einops import rearrange, repeat
|
||||||
|
import pdb
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class Transformer3DModelOutput(BaseOutput):
|
||||||
|
sample: torch.FloatTensor
|
||||||
|
|
||||||
|
|
||||||
|
if is_xformers_available():
|
||||||
|
import xformers
|
||||||
|
import xformers.ops
|
||||||
|
else:
|
||||||
|
xformers = None
|
||||||
|
|
||||||
|
|
||||||
|
class Transformer3DModel(ModelMixin, ConfigMixin):
|
||||||
|
@register_to_config
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
num_attention_heads: int = 16,
|
||||||
|
attention_head_dim: int = 88,
|
||||||
|
in_channels: Optional[int] = None,
|
||||||
|
num_layers: int = 1,
|
||||||
|
dropout: float = 0.0,
|
||||||
|
norm_num_groups: int = 32,
|
||||||
|
cross_attention_dim: Optional[int] = None,
|
||||||
|
attention_bias: bool = False,
|
||||||
|
activation_fn: str = "geglu",
|
||||||
|
num_embeds_ada_norm: Optional[int] = None,
|
||||||
|
use_linear_projection: bool = False,
|
||||||
|
only_cross_attention: bool = False,
|
||||||
|
upcast_attention: bool = False,
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=None,
|
||||||
|
unet_use_temporal_attention=None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.use_linear_projection = use_linear_projection
|
||||||
|
self.num_attention_heads = num_attention_heads
|
||||||
|
self.attention_head_dim = attention_head_dim
|
||||||
|
inner_dim = num_attention_heads * attention_head_dim
|
||||||
|
|
||||||
|
# Define input layers
|
||||||
|
self.in_channels = in_channels
|
||||||
|
|
||||||
|
self.norm = torch.nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True)
|
||||||
|
if use_linear_projection:
|
||||||
|
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||||
|
else:
|
||||||
|
self.proj_in = nn.Conv2d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0)
|
||||||
|
|
||||||
|
# Define transformers blocks
|
||||||
|
self.transformer_blocks = nn.ModuleList(
|
||||||
|
[
|
||||||
|
BasicTransformerBlock(
|
||||||
|
inner_dim,
|
||||||
|
num_attention_heads,
|
||||||
|
attention_head_dim,
|
||||||
|
dropout=dropout,
|
||||||
|
cross_attention_dim=cross_attention_dim,
|
||||||
|
activation_fn=activation_fn,
|
||||||
|
num_embeds_ada_norm=num_embeds_ada_norm,
|
||||||
|
attention_bias=attention_bias,
|
||||||
|
only_cross_attention=only_cross_attention,
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=unet_use_cross_frame_attention,
|
||||||
|
unet_use_temporal_attention=unet_use_temporal_attention,
|
||||||
|
)
|
||||||
|
for d in range(num_layers)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
# 4. Define output layers
|
||||||
|
if use_linear_projection:
|
||||||
|
self.proj_out = nn.Linear(in_channels, inner_dim)
|
||||||
|
else:
|
||||||
|
self.proj_out = nn.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0)
|
||||||
|
|
||||||
|
def forward(self, hidden_states, encoder_hidden_states=None, timestep=None, return_dict: bool = True):
|
||||||
|
# Input
|
||||||
|
assert hidden_states.dim() == 5, f"Expected hidden_states to have ndim=5, but got ndim={hidden_states.dim()}."
|
||||||
|
video_length = hidden_states.shape[2]
|
||||||
|
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||||
|
encoder_hidden_states = repeat(encoder_hidden_states, 'b n c -> (b f) n c', f=video_length)
|
||||||
|
|
||||||
|
batch, channel, height, weight = hidden_states.shape
|
||||||
|
residual = hidden_states
|
||||||
|
|
||||||
|
hidden_states = self.norm(hidden_states)
|
||||||
|
if not self.use_linear_projection:
|
||||||
|
hidden_states = self.proj_in(hidden_states)
|
||||||
|
inner_dim = hidden_states.shape[1]
|
||||||
|
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * weight, inner_dim)
|
||||||
|
else:
|
||||||
|
inner_dim = hidden_states.shape[1]
|
||||||
|
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * weight, inner_dim)
|
||||||
|
hidden_states = self.proj_in(hidden_states)
|
||||||
|
|
||||||
|
# Blocks
|
||||||
|
for block in self.transformer_blocks:
|
||||||
|
hidden_states = block(
|
||||||
|
hidden_states,
|
||||||
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
|
timestep=timestep,
|
||||||
|
video_length=video_length
|
||||||
|
)
|
||||||
|
|
||||||
|
# Output
|
||||||
|
if not self.use_linear_projection:
|
||||||
|
hidden_states = (
|
||||||
|
hidden_states.reshape(batch, height, weight, inner_dim).permute(0, 3, 1, 2).contiguous()
|
||||||
|
)
|
||||||
|
hidden_states = self.proj_out(hidden_states)
|
||||||
|
else:
|
||||||
|
hidden_states = self.proj_out(hidden_states)
|
||||||
|
hidden_states = (
|
||||||
|
hidden_states.reshape(batch, height, weight, inner_dim).permute(0, 3, 1, 2).contiguous()
|
||||||
|
)
|
||||||
|
|
||||||
|
output = hidden_states + residual
|
||||||
|
|
||||||
|
output = rearrange(output, "(b f) c h w -> b c f h w", f=video_length)
|
||||||
|
if not return_dict:
|
||||||
|
return (output,)
|
||||||
|
|
||||||
|
return Transformer3DModelOutput(sample=output)
|
||||||
|
|
||||||
|
|
||||||
|
class BasicTransformerBlock(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim: int,
|
||||||
|
num_attention_heads: int,
|
||||||
|
attention_head_dim: int,
|
||||||
|
dropout=0.0,
|
||||||
|
cross_attention_dim: Optional[int] = None,
|
||||||
|
activation_fn: str = "geglu",
|
||||||
|
num_embeds_ada_norm: Optional[int] = None,
|
||||||
|
attention_bias: bool = False,
|
||||||
|
only_cross_attention: bool = False,
|
||||||
|
upcast_attention: bool = False,
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention = None,
|
||||||
|
unet_use_temporal_attention = None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.only_cross_attention = only_cross_attention
|
||||||
|
self.use_ada_layer_norm = num_embeds_ada_norm is not None
|
||||||
|
self.unet_use_cross_frame_attention = unet_use_cross_frame_attention
|
||||||
|
self.unet_use_temporal_attention = unet_use_temporal_attention
|
||||||
|
|
||||||
|
# SC-Attn
|
||||||
|
assert unet_use_cross_frame_attention is not None
|
||||||
|
if unet_use_cross_frame_attention:
|
||||||
|
self.attn1 = SparseCausalAttention2D(
|
||||||
|
query_dim=dim,
|
||||||
|
heads=num_attention_heads,
|
||||||
|
dim_head=attention_head_dim,
|
||||||
|
dropout=dropout,
|
||||||
|
bias=attention_bias,
|
||||||
|
cross_attention_dim=cross_attention_dim if only_cross_attention else None,
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.attn1 = Attention(
|
||||||
|
query_dim=dim,
|
||||||
|
heads=num_attention_heads,
|
||||||
|
dim_head=attention_head_dim,
|
||||||
|
dropout=dropout,
|
||||||
|
bias=attention_bias,
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
)
|
||||||
|
self.norm1 = AdaLayerNorm(dim, num_embeds_ada_norm) if self.use_ada_layer_norm else nn.LayerNorm(dim)
|
||||||
|
|
||||||
|
# Cross-Attn
|
||||||
|
if cross_attention_dim is not None:
|
||||||
|
self.attn2 = Attention(
|
||||||
|
query_dim=dim,
|
||||||
|
cross_attention_dim=cross_attention_dim,
|
||||||
|
heads=num_attention_heads,
|
||||||
|
dim_head=attention_head_dim,
|
||||||
|
dropout=dropout,
|
||||||
|
bias=attention_bias,
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.attn2 = None
|
||||||
|
|
||||||
|
if cross_attention_dim is not None:
|
||||||
|
self.norm2 = AdaLayerNorm(dim, num_embeds_ada_norm) if self.use_ada_layer_norm else nn.LayerNorm(dim)
|
||||||
|
else:
|
||||||
|
self.norm2 = None
|
||||||
|
|
||||||
|
# Feed-forward
|
||||||
|
self.ff = FeedForward(dim, dropout=dropout, activation_fn=activation_fn)
|
||||||
|
self.norm3 = nn.LayerNorm(dim)
|
||||||
|
|
||||||
|
# Temp-Attn
|
||||||
|
assert unet_use_temporal_attention is not None
|
||||||
|
if unet_use_temporal_attention:
|
||||||
|
self.attn_temp = Attention(
|
||||||
|
query_dim=dim,
|
||||||
|
heads=num_attention_heads,
|
||||||
|
dim_head=attention_head_dim,
|
||||||
|
dropout=dropout,
|
||||||
|
bias=attention_bias,
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
)
|
||||||
|
nn.init.zeros_(self.attn_temp.to_out[0].weight.data)
|
||||||
|
self.norm_temp = AdaLayerNorm(dim, num_embeds_ada_norm) if self.use_ada_layer_norm else nn.LayerNorm(dim)
|
||||||
|
|
||||||
|
def set_use_memory_efficient_attention_xformers(self, use_memory_efficient_attention_xformers: bool, *args, **kwargs):
|
||||||
|
if not is_xformers_available():
|
||||||
|
print("Here is how to install it")
|
||||||
|
raise ModuleNotFoundError(
|
||||||
|
"Refer to https://github.com/facebookresearch/xformers for more information on how to install"
|
||||||
|
" xformers",
|
||||||
|
name="xformers",
|
||||||
|
)
|
||||||
|
elif not torch.cuda.is_available():
|
||||||
|
raise ValueError(
|
||||||
|
"torch.cuda.is_available() should be True but is False. xformers' memory efficient attention is only"
|
||||||
|
" available for GPU "
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
try:
|
||||||
|
# Make sure we can run the memory efficient attention
|
||||||
|
_ = xformers.ops.memory_efficient_attention(
|
||||||
|
torch.randn((1, 2, 40), device="cuda"),
|
||||||
|
torch.randn((1, 2, 40), device="cuda"),
|
||||||
|
torch.randn((1, 2, 40), device="cuda"),
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
raise e
|
||||||
|
self.attn1._use_memory_efficient_attention_xformers = use_memory_efficient_attention_xformers
|
||||||
|
if self.attn2 is not None:
|
||||||
|
self.attn2._use_memory_efficient_attention_xformers = use_memory_efficient_attention_xformers
|
||||||
|
# self.attn_temp._use_memory_efficient_attention_xformers = use_memory_efficient_attention_xformers
|
||||||
|
|
||||||
|
def forward(self, hidden_states, encoder_hidden_states=None, timestep=None, attention_mask=None, video_length=None):
|
||||||
|
# SparseCausal-Attention
|
||||||
|
norm_hidden_states = (
|
||||||
|
self.norm1(hidden_states, timestep) if self.use_ada_layer_norm else self.norm1(hidden_states)
|
||||||
|
)
|
||||||
|
|
||||||
|
# if self.only_cross_attention:
|
||||||
|
# hidden_states = (
|
||||||
|
# self.attn1(norm_hidden_states, encoder_hidden_states, attention_mask=attention_mask) + hidden_states
|
||||||
|
# )
|
||||||
|
# else:
|
||||||
|
# hidden_states = self.attn1(norm_hidden_states, attention_mask=attention_mask, video_length=video_length) + hidden_states
|
||||||
|
|
||||||
|
# pdb.set_trace()
|
||||||
|
if self.unet_use_cross_frame_attention:
|
||||||
|
hidden_states = self.attn1(norm_hidden_states, attention_mask=attention_mask, video_length=video_length) + hidden_states
|
||||||
|
else:
|
||||||
|
hidden_states = self.attn1(norm_hidden_states, attention_mask=attention_mask) + hidden_states
|
||||||
|
|
||||||
|
if self.attn2 is not None:
|
||||||
|
# Cross-Attention
|
||||||
|
norm_hidden_states = (
|
||||||
|
self.norm2(hidden_states, timestep) if self.use_ada_layer_norm else self.norm2(hidden_states)
|
||||||
|
)
|
||||||
|
hidden_states = (
|
||||||
|
self.attn2(
|
||||||
|
norm_hidden_states, encoder_hidden_states=encoder_hidden_states, attention_mask=attention_mask
|
||||||
|
)
|
||||||
|
+ hidden_states
|
||||||
|
)
|
||||||
|
|
||||||
|
# Feed-forward
|
||||||
|
hidden_states = self.ff(self.norm3(hidden_states)) + hidden_states
|
||||||
|
|
||||||
|
# Temporal-Attention
|
||||||
|
if self.unet_use_temporal_attention:
|
||||||
|
d = hidden_states.shape[1]
|
||||||
|
hidden_states = rearrange(hidden_states, "(b f) d c -> (b d) f c", f=video_length)
|
||||||
|
norm_hidden_states = (
|
||||||
|
self.norm_temp(hidden_states, timestep) if self.use_ada_layer_norm else self.norm_temp(hidden_states)
|
||||||
|
)
|
||||||
|
hidden_states = self.attn_temp(norm_hidden_states) + hidden_states
|
||||||
|
hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d)
|
||||||
|
|
||||||
|
return hidden_states
|
||||||
@@ -0,0 +1,353 @@
|
|||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import numpy as np
|
||||||
|
import torch.nn.functional as F
|
||||||
|
from torch import nn
|
||||||
|
import torchvision
|
||||||
|
import diffusers
|
||||||
|
from packaging import version
|
||||||
|
|
||||||
|
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||||
|
from diffusers import ModelMixin
|
||||||
|
from diffusers.utils import BaseOutput
|
||||||
|
from diffusers.utils.import_utils import is_xformers_available
|
||||||
|
from diffusers.models.attention import Attention, FeedForward
|
||||||
|
|
||||||
|
from einops import rearrange, repeat
|
||||||
|
import math
|
||||||
|
|
||||||
|
|
||||||
|
def zero_module(module):
|
||||||
|
# Zero out the parameters of a module and return it.
|
||||||
|
for p in module.parameters():
|
||||||
|
p.detach().zero_()
|
||||||
|
return module
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class TemporalTransformer3DModelOutput(BaseOutput):
|
||||||
|
sample: torch.FloatTensor
|
||||||
|
|
||||||
|
|
||||||
|
if is_xformers_available():
|
||||||
|
import xformers
|
||||||
|
import xformers.ops
|
||||||
|
else:
|
||||||
|
xformers = None
|
||||||
|
|
||||||
|
|
||||||
|
def get_motion_module(
|
||||||
|
in_channels,
|
||||||
|
motion_module_type: str,
|
||||||
|
motion_module_kwargs: dict
|
||||||
|
):
|
||||||
|
if motion_module_type == "Vanilla":
|
||||||
|
return VanillaTemporalModule(in_channels=in_channels, **motion_module_kwargs,)
|
||||||
|
else:
|
||||||
|
raise ValueError
|
||||||
|
|
||||||
|
|
||||||
|
class VanillaTemporalModule(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels,
|
||||||
|
num_attention_heads = 8,
|
||||||
|
num_transformer_block = 2,
|
||||||
|
attention_block_types =( "Temporal_Self", "Temporal_Self" ),
|
||||||
|
cross_frame_attention_mode = None,
|
||||||
|
temporal_position_encoding = False,
|
||||||
|
temporal_position_encoding_max_len = 24,
|
||||||
|
temporal_attention_dim_div = 1,
|
||||||
|
zero_initialize = True,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.temporal_transformer = TemporalTransformer3DModel(
|
||||||
|
in_channels=in_channels,
|
||||||
|
num_attention_heads=num_attention_heads,
|
||||||
|
attention_head_dim=in_channels // num_attention_heads // temporal_attention_dim_div,
|
||||||
|
num_layers=num_transformer_block,
|
||||||
|
attention_block_types=attention_block_types,
|
||||||
|
cross_frame_attention_mode=cross_frame_attention_mode,
|
||||||
|
temporal_position_encoding=temporal_position_encoding,
|
||||||
|
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||||
|
)
|
||||||
|
|
||||||
|
if zero_initialize:
|
||||||
|
self.temporal_transformer.proj_out = zero_module(self.temporal_transformer.proj_out)
|
||||||
|
|
||||||
|
def forward(self, input_tensor, temb, encoder_hidden_states, attention_mask=None, anchor_frame_idx=None):
|
||||||
|
video_length = input_tensor.shape[2]
|
||||||
|
|
||||||
|
if video_length > 1:
|
||||||
|
hidden_states = input_tensor
|
||||||
|
hidden_states = self.temporal_transformer(hidden_states, encoder_hidden_states, attention_mask)
|
||||||
|
output = hidden_states
|
||||||
|
else:
|
||||||
|
output = input_tensor
|
||||||
|
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
class TemporalTransformer3DModel(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels,
|
||||||
|
num_attention_heads,
|
||||||
|
attention_head_dim,
|
||||||
|
|
||||||
|
num_layers,
|
||||||
|
attention_block_types = ( "Temporal_Self", "Temporal_Self", ),
|
||||||
|
dropout = 0.0,
|
||||||
|
norm_num_groups = 32,
|
||||||
|
cross_attention_dim = 768,
|
||||||
|
activation_fn = "geglu",
|
||||||
|
attention_bias = False,
|
||||||
|
upcast_attention = False,
|
||||||
|
|
||||||
|
cross_frame_attention_mode = None,
|
||||||
|
temporal_position_encoding = False,
|
||||||
|
temporal_position_encoding_max_len = 24,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
inner_dim = num_attention_heads * attention_head_dim
|
||||||
|
|
||||||
|
self.norm = torch.nn.GroupNorm(num_groups=norm_num_groups, num_channels=in_channels, eps=1e-6, affine=True)
|
||||||
|
self.proj_in = nn.Linear(in_channels, inner_dim)
|
||||||
|
|
||||||
|
self.transformer_blocks = nn.ModuleList(
|
||||||
|
[
|
||||||
|
TemporalTransformerBlock(
|
||||||
|
dim=inner_dim,
|
||||||
|
num_attention_heads=num_attention_heads,
|
||||||
|
attention_head_dim=attention_head_dim,
|
||||||
|
attention_block_types=attention_block_types,
|
||||||
|
dropout=dropout,
|
||||||
|
norm_num_groups=norm_num_groups,
|
||||||
|
cross_attention_dim=cross_attention_dim,
|
||||||
|
activation_fn=activation_fn,
|
||||||
|
attention_bias=attention_bias,
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
cross_frame_attention_mode=cross_frame_attention_mode,
|
||||||
|
temporal_position_encoding=temporal_position_encoding,
|
||||||
|
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||||
|
)
|
||||||
|
for d in range(num_layers)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
self.proj_out = nn.Linear(inner_dim, in_channels)
|
||||||
|
|
||||||
|
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None):
|
||||||
|
assert hidden_states.dim() == 5, f"Expected hidden_states to have ndim=5, but got ndim={hidden_states.dim()}."
|
||||||
|
video_length = hidden_states.shape[2]
|
||||||
|
hidden_states = rearrange(hidden_states, "b c f h w -> (b f) c h w")
|
||||||
|
|
||||||
|
batch, channel, height, weight = hidden_states.shape
|
||||||
|
residual = hidden_states
|
||||||
|
|
||||||
|
hidden_states = self.norm(hidden_states)
|
||||||
|
inner_dim = hidden_states.shape[1]
|
||||||
|
hidden_states = hidden_states.permute(0, 2, 3, 1).reshape(batch, height * weight, inner_dim)
|
||||||
|
hidden_states = self.proj_in(hidden_states)
|
||||||
|
|
||||||
|
# Transformer Blocks
|
||||||
|
for block in self.transformer_blocks:
|
||||||
|
hidden_states = block(hidden_states, encoder_hidden_states=encoder_hidden_states, video_length=video_length)
|
||||||
|
|
||||||
|
# output
|
||||||
|
hidden_states = self.proj_out(hidden_states)
|
||||||
|
hidden_states = hidden_states.reshape(batch, height, weight, inner_dim).permute(0, 3, 1, 2).contiguous()
|
||||||
|
|
||||||
|
output = hidden_states + residual
|
||||||
|
output = rearrange(output, "(b f) c h w -> b c f h w", f=video_length)
|
||||||
|
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
class TemporalTransformerBlock(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
dim,
|
||||||
|
num_attention_heads,
|
||||||
|
attention_head_dim,
|
||||||
|
attention_block_types = ( "Temporal_Self", "Temporal_Self", ),
|
||||||
|
dropout = 0.0,
|
||||||
|
norm_num_groups = 32,
|
||||||
|
cross_attention_dim = 768,
|
||||||
|
activation_fn = "geglu",
|
||||||
|
attention_bias = False,
|
||||||
|
upcast_attention = False,
|
||||||
|
cross_frame_attention_mode = None,
|
||||||
|
temporal_position_encoding = False,
|
||||||
|
temporal_position_encoding_max_len = 24,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
attention_blocks = []
|
||||||
|
norms = []
|
||||||
|
|
||||||
|
for block_name in attention_block_types:
|
||||||
|
attention_blocks.append(
|
||||||
|
VersatileAttention(
|
||||||
|
attention_mode=block_name.split("_")[0],
|
||||||
|
cross_attention_dim=cross_attention_dim if block_name.endswith("_Cross") else None,
|
||||||
|
|
||||||
|
query_dim=dim,
|
||||||
|
heads=num_attention_heads,
|
||||||
|
dim_head=attention_head_dim,
|
||||||
|
dropout=dropout,
|
||||||
|
bias=attention_bias,
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
|
||||||
|
cross_frame_attention_mode=cross_frame_attention_mode,
|
||||||
|
temporal_position_encoding=temporal_position_encoding,
|
||||||
|
temporal_position_encoding_max_len=temporal_position_encoding_max_len,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
norms.append(nn.LayerNorm(dim))
|
||||||
|
|
||||||
|
self.attention_blocks = nn.ModuleList(attention_blocks)
|
||||||
|
self.norms = nn.ModuleList(norms)
|
||||||
|
|
||||||
|
self.ff = FeedForward(dim, dropout=dropout, activation_fn=activation_fn)
|
||||||
|
self.ff_norm = nn.LayerNorm(dim)
|
||||||
|
|
||||||
|
|
||||||
|
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, video_length=None):
|
||||||
|
for attention_block, norm in zip(self.attention_blocks, self.norms):
|
||||||
|
norm_hidden_states = norm(hidden_states)
|
||||||
|
hidden_states = attention_block(
|
||||||
|
norm_hidden_states,
|
||||||
|
encoder_hidden_states=encoder_hidden_states if attention_block.is_cross_attention else None,
|
||||||
|
video_length=video_length,
|
||||||
|
) + hidden_states
|
||||||
|
|
||||||
|
hidden_states = self.ff(self.ff_norm(hidden_states)) + hidden_states
|
||||||
|
|
||||||
|
output = hidden_states
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
class PositionalEncoding(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
d_model,
|
||||||
|
dropout = 0.,
|
||||||
|
max_len = 24
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.dropout = nn.Dropout(p=dropout)
|
||||||
|
position = torch.arange(max_len).unsqueeze(1)
|
||||||
|
div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
|
||||||
|
pe = torch.zeros(1, max_len, d_model)
|
||||||
|
pe[0, :, 0::2] = torch.sin(position * div_term)
|
||||||
|
pe[0, :, 1::2] = torch.cos(position * div_term)
|
||||||
|
self.register_buffer('pe', pe)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
x = x + self.pe[:, :x.size(1)]
|
||||||
|
return self.dropout(x)
|
||||||
|
|
||||||
|
|
||||||
|
class VersatileAttention(Attention):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
attention_mode = None,
|
||||||
|
cross_frame_attention_mode = None,
|
||||||
|
temporal_position_encoding = False,
|
||||||
|
temporal_position_encoding_max_len = 24,
|
||||||
|
*args, **kwargs
|
||||||
|
):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
assert attention_mode == "Temporal"
|
||||||
|
|
||||||
|
self.attention_mode = attention_mode
|
||||||
|
self.is_cross_attention = kwargs["cross_attention_dim"] is not None
|
||||||
|
|
||||||
|
self.pos_encoder = PositionalEncoding(
|
||||||
|
kwargs["query_dim"],
|
||||||
|
dropout=0.,
|
||||||
|
max_len=temporal_position_encoding_max_len
|
||||||
|
) if (temporal_position_encoding and attention_mode == "Temporal") else None
|
||||||
|
|
||||||
|
def extra_repr(self):
|
||||||
|
return f"(Module Info) Attention_Mode: {self.attention_mode}, Is_Cross_Attention: {self.is_cross_attention}"
|
||||||
|
|
||||||
|
def forward(self, hidden_states, encoder_hidden_states=None, attention_mask=None, video_length=None):
|
||||||
|
batch_size, sequence_length, _ = hidden_states.shape
|
||||||
|
|
||||||
|
if self.attention_mode == "Temporal":
|
||||||
|
d = hidden_states.shape[1]
|
||||||
|
hidden_states = rearrange(hidden_states, "(b f) d c -> (b d) f c", f=video_length)
|
||||||
|
|
||||||
|
if self.pos_encoder is not None:
|
||||||
|
hidden_states = self.pos_encoder(hidden_states)
|
||||||
|
|
||||||
|
encoder_hidden_states = repeat(encoder_hidden_states, "b n c -> (b d) n c", d=d) if encoder_hidden_states is not None else encoder_hidden_states
|
||||||
|
else:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
encoder_hidden_states = encoder_hidden_states
|
||||||
|
|
||||||
|
if version.parse(diffusers.__version__) > version.parse("0.11.1"):
|
||||||
|
hidden_states = self.processor(self, hidden_states, encoder_hidden_states)
|
||||||
|
else:
|
||||||
|
if self.group_norm is not None:
|
||||||
|
hidden_states = self.group_norm(hidden_states.transpose(1, 2)).transpose(1, 2)
|
||||||
|
|
||||||
|
query = self.to_q(hidden_states)
|
||||||
|
dim = query.shape[-1]
|
||||||
|
query = self.head_to_batch_dim(query)
|
||||||
|
|
||||||
|
if self.added_kv_proj_dim is not None:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
encoder_hidden_states = encoder_hidden_states if encoder_hidden_states is not None else hidden_states
|
||||||
|
key = self.to_k(encoder_hidden_states)
|
||||||
|
value = self.to_v(encoder_hidden_states)
|
||||||
|
|
||||||
|
key = self.head_to_batch_dim(key)
|
||||||
|
value = self.head_to_batch_dim(value)
|
||||||
|
|
||||||
|
if attention_mask is not None:
|
||||||
|
if attention_mask.shape[-1] != query.shape[1]:
|
||||||
|
target_length = query.shape[1]
|
||||||
|
attention_mask = F.pad(attention_mask, (0, target_length), value=0.0)
|
||||||
|
attention_mask = attention_mask.repeat_interleave(self.heads, dim=0)
|
||||||
|
|
||||||
|
# attention, what we cannot get enough of
|
||||||
|
|
||||||
|
if self._use_memory_efficient_attention_xformers:
|
||||||
|
hidden_states = self._memory_efficient_attention_xformers(query, key, value, attention_mask)
|
||||||
|
# Some versions of xformers return output in fp32, cast it back to the dtype of the input
|
||||||
|
hidden_states = hidden_states.to(query.dtype)
|
||||||
|
else:
|
||||||
|
if self._slice_size is None or query.shape[0] // self._slice_size == 1:
|
||||||
|
hidden_states = self._attention(query, key, value, attention_mask)
|
||||||
|
else:
|
||||||
|
hidden_states = self._sliced_attention(query, key, value, sequence_length, dim, attention_mask)
|
||||||
|
else:
|
||||||
|
#if "xformers" in self.processor.__class__.__name__.lower():
|
||||||
|
# hidden_states = xformers.ops.memory_efficient_attention(query, key, value, attention_mask)
|
||||||
|
# # Some versions of xformers return output in fp32, cast it back to the dtype of the input
|
||||||
|
# hidden_states = hidden_states.to(query.dtype)
|
||||||
|
#else:
|
||||||
|
hidden_states = F.scaled_dot_product_attention(
|
||||||
|
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||||
|
)
|
||||||
|
hidden_states = hidden_states.transpose(1, 2).reshape(batch_size, -1, attn.heads * head_dim)
|
||||||
|
hidden_states = hidden_states.to(query.dtype)
|
||||||
|
|
||||||
|
# linear proj
|
||||||
|
hidden_states = self.to_out[0](hidden_states)
|
||||||
|
|
||||||
|
# dropout
|
||||||
|
hidden_states = self.to_out[1](hidden_states)
|
||||||
|
|
||||||
|
if self.attention_mode == "Temporal":
|
||||||
|
hidden_states = rearrange(hidden_states, "(b d) f c -> (b f) d c", d=d)
|
||||||
|
|
||||||
|
return hidden_states
|
||||||
@@ -0,0 +1,217 @@
|
|||||||
|
# Adapted from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/resnet.py
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
from einops import rearrange
|
||||||
|
|
||||||
|
|
||||||
|
class InflatedConv3d(nn.Conv2d):
|
||||||
|
def forward(self, x):
|
||||||
|
video_length = x.shape[2]
|
||||||
|
|
||||||
|
x = rearrange(x, "b c f h w -> (b f) c h w")
|
||||||
|
x = super().forward(x)
|
||||||
|
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
|
||||||
|
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
class InflatedGroupNorm(nn.GroupNorm):
|
||||||
|
def forward(self, x):
|
||||||
|
video_length = x.shape[2]
|
||||||
|
|
||||||
|
x = rearrange(x, "b c f h w -> (b f) c h w")
|
||||||
|
x = super().forward(x)
|
||||||
|
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
|
||||||
|
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
class Upsample3D(nn.Module):
|
||||||
|
def __init__(self, channels, use_conv=False, use_conv_transpose=False, out_channels=None, name="conv"):
|
||||||
|
super().__init__()
|
||||||
|
self.channels = channels
|
||||||
|
self.out_channels = out_channels or channels
|
||||||
|
self.use_conv = use_conv
|
||||||
|
self.use_conv_transpose = use_conv_transpose
|
||||||
|
self.name = name
|
||||||
|
|
||||||
|
conv = None
|
||||||
|
if use_conv_transpose:
|
||||||
|
raise NotImplementedError
|
||||||
|
elif use_conv:
|
||||||
|
self.conv = InflatedConv3d(self.channels, self.out_channels, 3, padding=1)
|
||||||
|
|
||||||
|
def forward(self, hidden_states, output_size=None):
|
||||||
|
assert hidden_states.shape[1] == self.channels
|
||||||
|
|
||||||
|
if self.use_conv_transpose:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
# Cast to float32 to as 'upsample_nearest2d_out_frame' op does not support bfloat16
|
||||||
|
dtype = hidden_states.dtype
|
||||||
|
if dtype == torch.bfloat16:
|
||||||
|
hidden_states = hidden_states.to(torch.float32)
|
||||||
|
|
||||||
|
# upsample_nearest_nhwc fails with large batch sizes. see https://github.com/huggingface/diffusers/issues/984
|
||||||
|
if hidden_states.shape[0] >= 64:
|
||||||
|
hidden_states = hidden_states.contiguous()
|
||||||
|
|
||||||
|
# if `output_size` is passed we force the interpolation output
|
||||||
|
# size and do not make use of `scale_factor=2`
|
||||||
|
if output_size is None:
|
||||||
|
hidden_states = F.interpolate(hidden_states, scale_factor=[1.0, 2.0, 2.0], mode="nearest")
|
||||||
|
else:
|
||||||
|
hidden_states = F.interpolate(hidden_states, size=output_size, mode="nearest")
|
||||||
|
|
||||||
|
# If the input is bfloat16, we cast back to bfloat16
|
||||||
|
if dtype == torch.bfloat16:
|
||||||
|
hidden_states = hidden_states.to(dtype)
|
||||||
|
|
||||||
|
# if self.use_conv:
|
||||||
|
# if self.name == "conv":
|
||||||
|
# hidden_states = self.conv(hidden_states)
|
||||||
|
# else:
|
||||||
|
# hidden_states = self.Conv2d_0(hidden_states)
|
||||||
|
hidden_states = self.conv(hidden_states)
|
||||||
|
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
|
class Downsample3D(nn.Module):
|
||||||
|
def __init__(self, channels, use_conv=False, out_channels=None, padding=1, name="conv"):
|
||||||
|
super().__init__()
|
||||||
|
self.channels = channels
|
||||||
|
self.out_channels = out_channels or channels
|
||||||
|
self.use_conv = use_conv
|
||||||
|
self.padding = padding
|
||||||
|
stride = 2
|
||||||
|
self.name = name
|
||||||
|
|
||||||
|
if use_conv:
|
||||||
|
self.conv = InflatedConv3d(self.channels, self.out_channels, 3, stride=stride, padding=padding)
|
||||||
|
else:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def forward(self, hidden_states):
|
||||||
|
assert hidden_states.shape[1] == self.channels
|
||||||
|
if self.use_conv and self.padding == 0:
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
assert hidden_states.shape[1] == self.channels
|
||||||
|
hidden_states = self.conv(hidden_states)
|
||||||
|
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
|
class ResnetBlock3D(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
in_channels,
|
||||||
|
out_channels=None,
|
||||||
|
conv_shortcut=False,
|
||||||
|
dropout=0.0,
|
||||||
|
temb_channels=512,
|
||||||
|
groups=32,
|
||||||
|
groups_out=None,
|
||||||
|
pre_norm=True,
|
||||||
|
eps=1e-6,
|
||||||
|
non_linearity="swish",
|
||||||
|
time_embedding_norm="default",
|
||||||
|
output_scale_factor=1.0,
|
||||||
|
use_in_shortcut=None,
|
||||||
|
use_inflated_groupnorm=False,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.pre_norm = pre_norm
|
||||||
|
self.pre_norm = True
|
||||||
|
self.in_channels = in_channels
|
||||||
|
out_channels = in_channels if out_channels is None else out_channels
|
||||||
|
self.out_channels = out_channels
|
||||||
|
self.use_conv_shortcut = conv_shortcut
|
||||||
|
self.time_embedding_norm = time_embedding_norm
|
||||||
|
self.output_scale_factor = output_scale_factor
|
||||||
|
|
||||||
|
if groups_out is None:
|
||||||
|
groups_out = groups
|
||||||
|
|
||||||
|
assert use_inflated_groupnorm != None
|
||||||
|
if use_inflated_groupnorm:
|
||||||
|
self.norm1 = InflatedGroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
|
||||||
|
else:
|
||||||
|
self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
|
||||||
|
|
||||||
|
self.conv1 = InflatedConv3d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||||
|
|
||||||
|
if temb_channels is not None:
|
||||||
|
if self.time_embedding_norm == "default":
|
||||||
|
time_emb_proj_out_channels = out_channels
|
||||||
|
elif self.time_embedding_norm == "scale_shift":
|
||||||
|
time_emb_proj_out_channels = out_channels * 2
|
||||||
|
else:
|
||||||
|
raise ValueError(f"unknown time_embedding_norm : {self.time_embedding_norm} ")
|
||||||
|
|
||||||
|
self.time_emb_proj = torch.nn.Linear(temb_channels, time_emb_proj_out_channels)
|
||||||
|
else:
|
||||||
|
self.time_emb_proj = None
|
||||||
|
|
||||||
|
if use_inflated_groupnorm:
|
||||||
|
self.norm2 = InflatedGroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
|
||||||
|
else:
|
||||||
|
self.norm2 = torch.nn.GroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
|
||||||
|
|
||||||
|
self.dropout = torch.nn.Dropout(dropout)
|
||||||
|
self.conv2 = InflatedConv3d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
|
||||||
|
|
||||||
|
if non_linearity == "swish":
|
||||||
|
self.nonlinearity = lambda x: F.silu(x)
|
||||||
|
elif non_linearity == "mish":
|
||||||
|
self.nonlinearity = Mish()
|
||||||
|
elif non_linearity == "silu":
|
||||||
|
self.nonlinearity = nn.SiLU()
|
||||||
|
|
||||||
|
self.use_in_shortcut = self.in_channels != self.out_channels if use_in_shortcut is None else use_in_shortcut
|
||||||
|
|
||||||
|
self.conv_shortcut = None
|
||||||
|
if self.use_in_shortcut:
|
||||||
|
self.conv_shortcut = InflatedConv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
|
||||||
|
|
||||||
|
def forward(self, input_tensor, temb):
|
||||||
|
hidden_states = input_tensor
|
||||||
|
|
||||||
|
hidden_states = self.norm1(hidden_states)
|
||||||
|
hidden_states = self.nonlinearity(hidden_states)
|
||||||
|
|
||||||
|
hidden_states = self.conv1(hidden_states)
|
||||||
|
|
||||||
|
if temb is not None:
|
||||||
|
temb = self.time_emb_proj(self.nonlinearity(temb))[:, :, None, None, None]
|
||||||
|
|
||||||
|
if temb is not None and self.time_embedding_norm == "default":
|
||||||
|
hidden_states = hidden_states + temb
|
||||||
|
|
||||||
|
hidden_states = self.norm2(hidden_states)
|
||||||
|
|
||||||
|
if temb is not None and self.time_embedding_norm == "scale_shift":
|
||||||
|
scale, shift = torch.chunk(temb, 2, dim=1)
|
||||||
|
hidden_states = hidden_states * (1 + scale) + shift
|
||||||
|
|
||||||
|
hidden_states = self.nonlinearity(hidden_states)
|
||||||
|
|
||||||
|
hidden_states = self.dropout(hidden_states)
|
||||||
|
hidden_states = self.conv2(hidden_states)
|
||||||
|
|
||||||
|
if self.conv_shortcut is not None:
|
||||||
|
input_tensor = self.conv_shortcut(input_tensor)
|
||||||
|
|
||||||
|
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
|
||||||
|
|
||||||
|
return output_tensor
|
||||||
|
|
||||||
|
|
||||||
|
class Mish(torch.nn.Module):
|
||||||
|
def forward(self, hidden_states):
|
||||||
|
return hidden_states * torch.tanh(torch.nn.functional.softplus(hidden_states))
|
||||||
@@ -0,0 +1,587 @@
|
|||||||
|
# Copyright 2023 The HuggingFace Team. All rights reserved.
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
#
|
||||||
|
# Changes were made to this source code by Yuwei Guo.
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
from torch.nn import functional as F
|
||||||
|
|
||||||
|
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||||
|
from diffusers.utils import BaseOutput, logging
|
||||||
|
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
|
||||||
|
from diffusers import ModelMixin
|
||||||
|
|
||||||
|
|
||||||
|
from .unet_blocks import (
|
||||||
|
CrossAttnDownBlock3D,
|
||||||
|
DownBlock3D,
|
||||||
|
UNetMidBlock3DCrossAttn,
|
||||||
|
get_down_block,
|
||||||
|
)
|
||||||
|
from einops import repeat, rearrange
|
||||||
|
from .resnet import InflatedConv3d
|
||||||
|
|
||||||
|
from diffusers.models.unets.unet_2d_condition import UNet2DConditionModel
|
||||||
|
|
||||||
|
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class SparseControlNetOutput(BaseOutput):
|
||||||
|
down_block_res_samples: Tuple[torch.Tensor]
|
||||||
|
mid_block_res_sample: torch.Tensor
|
||||||
|
|
||||||
|
|
||||||
|
class SparseControlNetConditioningEmbedding(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
conditioning_embedding_channels: int,
|
||||||
|
conditioning_channels: int = 3,
|
||||||
|
block_out_channels: Tuple[int] = (16, 32, 96, 256),
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.conv_in = InflatedConv3d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1)
|
||||||
|
|
||||||
|
self.blocks = nn.ModuleList([])
|
||||||
|
|
||||||
|
for i in range(len(block_out_channels) - 1):
|
||||||
|
channel_in = block_out_channels[i]
|
||||||
|
channel_out = block_out_channels[i + 1]
|
||||||
|
self.blocks.append(InflatedConv3d(channel_in, channel_in, kernel_size=3, padding=1))
|
||||||
|
self.blocks.append(InflatedConv3d(channel_in, channel_out, kernel_size=3, padding=1, stride=2))
|
||||||
|
|
||||||
|
self.conv_out = zero_module(
|
||||||
|
InflatedConv3d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1)
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, conditioning):
|
||||||
|
embedding = self.conv_in(conditioning)
|
||||||
|
embedding = F.silu(embedding)
|
||||||
|
|
||||||
|
for block in self.blocks:
|
||||||
|
embedding = block(embedding)
|
||||||
|
embedding = F.silu(embedding)
|
||||||
|
|
||||||
|
embedding = self.conv_out(embedding)
|
||||||
|
|
||||||
|
return embedding
|
||||||
|
|
||||||
|
|
||||||
|
class SparseControlNetModel(ModelMixin, ConfigMixin):
|
||||||
|
_supports_gradient_checkpointing = True
|
||||||
|
|
||||||
|
@register_to_config
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels: int = 4,
|
||||||
|
conditioning_channels: int = 3,
|
||||||
|
flip_sin_to_cos: bool = True,
|
||||||
|
freq_shift: int = 0,
|
||||||
|
down_block_types: Tuple[str] = (
|
||||||
|
"CrossAttnDownBlock2D",
|
||||||
|
"CrossAttnDownBlock2D",
|
||||||
|
"CrossAttnDownBlock2D",
|
||||||
|
"DownBlock2D",
|
||||||
|
),
|
||||||
|
only_cross_attention: Union[bool, Tuple[bool]] = False,
|
||||||
|
block_out_channels: Tuple[int] = (320, 640, 1280, 1280),
|
||||||
|
layers_per_block: int = 2,
|
||||||
|
downsample_padding: int = 1,
|
||||||
|
mid_block_scale_factor: float = 1,
|
||||||
|
act_fn: str = "silu",
|
||||||
|
norm_num_groups: Optional[int] = 32,
|
||||||
|
norm_eps: float = 1e-5,
|
||||||
|
cross_attention_dim: int = 1280,
|
||||||
|
attention_head_dim: Union[int, Tuple[int]] = 8,
|
||||||
|
num_attention_heads: Optional[Union[int, Tuple[int]]] = None,
|
||||||
|
use_linear_projection: bool = False,
|
||||||
|
class_embed_type: Optional[str] = None,
|
||||||
|
num_class_embeds: Optional[int] = None,
|
||||||
|
upcast_attention: bool = False,
|
||||||
|
resnet_time_scale_shift: str = "default",
|
||||||
|
projection_class_embeddings_input_dim: Optional[int] = None,
|
||||||
|
controlnet_conditioning_channel_order: str = "rgb",
|
||||||
|
conditioning_embedding_out_channels: Optional[Tuple[int]] = (16, 32, 96, 256),
|
||||||
|
global_pool_conditions: bool = False,
|
||||||
|
|
||||||
|
use_motion_module = True,
|
||||||
|
motion_module_resolutions = ( 1,2,4,8 ),
|
||||||
|
motion_module_mid_block = False,
|
||||||
|
motion_module_type = "Vanilla",
|
||||||
|
motion_module_kwargs = {
|
||||||
|
"num_attention_heads": 8,
|
||||||
|
"num_transformer_block": 1,
|
||||||
|
"attention_block_types": ["Temporal_Self"],
|
||||||
|
"temporal_position_encoding": True,
|
||||||
|
"temporal_position_encoding_max_len": 32,
|
||||||
|
"temporal_attention_dim_div": 1,
|
||||||
|
"causal_temporal_attention": False,
|
||||||
|
},
|
||||||
|
|
||||||
|
concate_conditioning_mask: bool = True,
|
||||||
|
use_simplified_condition_embedding: bool = False,
|
||||||
|
|
||||||
|
set_noisy_sample_input_to_zero: bool = False,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
# If `num_attention_heads` is not defined (which is the case for most models)
|
||||||
|
# it will default to `attention_head_dim`. This looks weird upon first reading it and it is.
|
||||||
|
# The reason for this behavior is to correct for incorrectly named variables that were introduced
|
||||||
|
# when this library was created. The incorrect naming was only discovered much later in https://github.com/huggingface/diffusers/issues/2011#issuecomment-1547958131
|
||||||
|
# Changing `attention_head_dim` to `num_attention_heads` for 40,000+ configurations is too backwards breaking
|
||||||
|
# which is why we correct for the naming here.
|
||||||
|
num_attention_heads = num_attention_heads or attention_head_dim
|
||||||
|
|
||||||
|
# Check inputs
|
||||||
|
if len(block_out_channels) != len(down_block_types):
|
||||||
|
raise ValueError(
|
||||||
|
f"Must provide the same number of `block_out_channels` as `down_block_types`. `block_out_channels`: {block_out_channels}. `down_block_types`: {down_block_types}."
|
||||||
|
)
|
||||||
|
|
||||||
|
if not isinstance(only_cross_attention, bool) and len(only_cross_attention) != len(down_block_types):
|
||||||
|
raise ValueError(
|
||||||
|
f"Must provide the same number of `only_cross_attention` as `down_block_types`. `only_cross_attention`: {only_cross_attention}. `down_block_types`: {down_block_types}."
|
||||||
|
)
|
||||||
|
|
||||||
|
if not isinstance(num_attention_heads, int) and len(num_attention_heads) != len(down_block_types):
|
||||||
|
raise ValueError(
|
||||||
|
f"Must provide the same number of `num_attention_heads` as `down_block_types`. `num_attention_heads`: {num_attention_heads}. `down_block_types`: {down_block_types}."
|
||||||
|
)
|
||||||
|
|
||||||
|
# input
|
||||||
|
self.set_noisy_sample_input_to_zero = set_noisy_sample_input_to_zero
|
||||||
|
|
||||||
|
conv_in_kernel = 3
|
||||||
|
conv_in_padding = (conv_in_kernel - 1) // 2
|
||||||
|
self.conv_in = InflatedConv3d(
|
||||||
|
in_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding
|
||||||
|
)
|
||||||
|
|
||||||
|
if concate_conditioning_mask:
|
||||||
|
conditioning_channels = conditioning_channels + 1
|
||||||
|
self.concate_conditioning_mask = concate_conditioning_mask
|
||||||
|
|
||||||
|
# control net conditioning embedding
|
||||||
|
if use_simplified_condition_embedding:
|
||||||
|
self.controlnet_cond_embedding = zero_module(
|
||||||
|
InflatedConv3d(conditioning_channels, block_out_channels[0], kernel_size=conv_in_kernel, padding=conv_in_padding)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.controlnet_cond_embedding = SparseControlNetConditioningEmbedding(
|
||||||
|
conditioning_embedding_channels=block_out_channels[0],
|
||||||
|
block_out_channels=conditioning_embedding_out_channels,
|
||||||
|
conditioning_channels=conditioning_channels,
|
||||||
|
)
|
||||||
|
self.use_simplified_condition_embedding = use_simplified_condition_embedding
|
||||||
|
|
||||||
|
# time
|
||||||
|
time_embed_dim = block_out_channels[0] * 4
|
||||||
|
|
||||||
|
self.time_proj = Timesteps(block_out_channels[0], flip_sin_to_cos, freq_shift)
|
||||||
|
timestep_input_dim = block_out_channels[0]
|
||||||
|
|
||||||
|
self.time_embedding = TimestepEmbedding(
|
||||||
|
timestep_input_dim,
|
||||||
|
time_embed_dim,
|
||||||
|
act_fn=act_fn,
|
||||||
|
)
|
||||||
|
|
||||||
|
# class embedding
|
||||||
|
if class_embed_type is None and num_class_embeds is not None:
|
||||||
|
self.class_embedding = nn.Embedding(num_class_embeds, time_embed_dim)
|
||||||
|
elif class_embed_type == "timestep":
|
||||||
|
self.class_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim)
|
||||||
|
elif class_embed_type == "identity":
|
||||||
|
self.class_embedding = nn.Identity(time_embed_dim, time_embed_dim)
|
||||||
|
elif class_embed_type == "projection":
|
||||||
|
if projection_class_embeddings_input_dim is None:
|
||||||
|
raise ValueError(
|
||||||
|
"`class_embed_type`: 'projection' requires `projection_class_embeddings_input_dim` be set"
|
||||||
|
)
|
||||||
|
# The projection `class_embed_type` is the same as the timestep `class_embed_type` except
|
||||||
|
# 1. the `class_labels` inputs are not first converted to sinusoidal embeddings
|
||||||
|
# 2. it projects from an arbitrary input dimension.
|
||||||
|
#
|
||||||
|
# Note that `TimestepEmbedding` is quite general, being mainly linear layers and activations.
|
||||||
|
# When used for embedding actual timesteps, the timesteps are first converted to sinusoidal embeddings.
|
||||||
|
# As a result, `TimestepEmbedding` can be passed arbitrary vectors.
|
||||||
|
self.class_embedding = TimestepEmbedding(projection_class_embeddings_input_dim, time_embed_dim)
|
||||||
|
else:
|
||||||
|
self.class_embedding = None
|
||||||
|
|
||||||
|
|
||||||
|
self.down_blocks = nn.ModuleList([])
|
||||||
|
self.controlnet_down_blocks = nn.ModuleList([])
|
||||||
|
|
||||||
|
if isinstance(only_cross_attention, bool):
|
||||||
|
only_cross_attention = [only_cross_attention] * len(down_block_types)
|
||||||
|
|
||||||
|
if isinstance(attention_head_dim, int):
|
||||||
|
attention_head_dim = (attention_head_dim,) * len(down_block_types)
|
||||||
|
|
||||||
|
if isinstance(num_attention_heads, int):
|
||||||
|
num_attention_heads = (num_attention_heads,) * len(down_block_types)
|
||||||
|
|
||||||
|
# down
|
||||||
|
output_channel = block_out_channels[0]
|
||||||
|
|
||||||
|
controlnet_block = InflatedConv3d(output_channel, output_channel, kernel_size=1)
|
||||||
|
controlnet_block = zero_module(controlnet_block)
|
||||||
|
self.controlnet_down_blocks.append(controlnet_block)
|
||||||
|
|
||||||
|
for i, down_block_type in enumerate(down_block_types):
|
||||||
|
res = 2 ** i
|
||||||
|
input_channel = output_channel
|
||||||
|
output_channel = block_out_channels[i]
|
||||||
|
is_final_block = i == len(block_out_channels) - 1
|
||||||
|
|
||||||
|
down_block = get_down_block(
|
||||||
|
down_block_type,
|
||||||
|
num_layers=layers_per_block,
|
||||||
|
in_channels=input_channel,
|
||||||
|
out_channels=output_channel,
|
||||||
|
temb_channels=time_embed_dim,
|
||||||
|
add_downsample=not is_final_block,
|
||||||
|
resnet_eps=norm_eps,
|
||||||
|
resnet_act_fn=act_fn,
|
||||||
|
resnet_groups=norm_num_groups,
|
||||||
|
cross_attention_dim=cross_attention_dim,
|
||||||
|
attn_num_head_channels=attention_head_dim[i] if attention_head_dim[i] is not None else output_channel,
|
||||||
|
downsample_padding=downsample_padding,
|
||||||
|
use_linear_projection=use_linear_projection,
|
||||||
|
only_cross_attention=only_cross_attention[i],
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||||
|
|
||||||
|
use_inflated_groupnorm=True,
|
||||||
|
|
||||||
|
use_motion_module=use_motion_module and (res in motion_module_resolutions),
|
||||||
|
motion_module_type=motion_module_type,
|
||||||
|
motion_module_kwargs=motion_module_kwargs,
|
||||||
|
)
|
||||||
|
self.down_blocks.append(down_block)
|
||||||
|
|
||||||
|
for _ in range(layers_per_block):
|
||||||
|
controlnet_block = InflatedConv3d(output_channel, output_channel, kernel_size=1)
|
||||||
|
controlnet_block = zero_module(controlnet_block)
|
||||||
|
self.controlnet_down_blocks.append(controlnet_block)
|
||||||
|
|
||||||
|
if not is_final_block:
|
||||||
|
controlnet_block = InflatedConv3d(output_channel, output_channel, kernel_size=1)
|
||||||
|
controlnet_block = zero_module(controlnet_block)
|
||||||
|
self.controlnet_down_blocks.append(controlnet_block)
|
||||||
|
|
||||||
|
# mid
|
||||||
|
mid_block_channel = block_out_channels[-1]
|
||||||
|
|
||||||
|
controlnet_block = InflatedConv3d(mid_block_channel, mid_block_channel, kernel_size=1)
|
||||||
|
controlnet_block = zero_module(controlnet_block)
|
||||||
|
self.controlnet_mid_block = controlnet_block
|
||||||
|
|
||||||
|
self.mid_block = UNetMidBlock3DCrossAttn(
|
||||||
|
in_channels=mid_block_channel,
|
||||||
|
temb_channels=time_embed_dim,
|
||||||
|
resnet_eps=norm_eps,
|
||||||
|
resnet_act_fn=act_fn,
|
||||||
|
output_scale_factor=mid_block_scale_factor,
|
||||||
|
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||||
|
cross_attention_dim=cross_attention_dim,
|
||||||
|
attn_num_head_channels=num_attention_heads[-1],
|
||||||
|
resnet_groups=norm_num_groups,
|
||||||
|
use_linear_projection=use_linear_projection,
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
|
||||||
|
use_inflated_groupnorm=True,
|
||||||
|
use_motion_module=use_motion_module and motion_module_mid_block,
|
||||||
|
motion_module_type=motion_module_type,
|
||||||
|
motion_module_kwargs=motion_module_kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_unet(
|
||||||
|
cls,
|
||||||
|
unet: UNet2DConditionModel,
|
||||||
|
controlnet_conditioning_channel_order: str = "rgb",
|
||||||
|
conditioning_embedding_out_channels: Optional[Tuple[int]] = (16, 32, 96, 256),
|
||||||
|
load_weights_from_unet: bool = True,
|
||||||
|
|
||||||
|
controlnet_additional_kwargs: dict = {},
|
||||||
|
):
|
||||||
|
controlnet = cls(
|
||||||
|
in_channels=unet.config.in_channels,
|
||||||
|
flip_sin_to_cos=unet.config.flip_sin_to_cos,
|
||||||
|
freq_shift=unet.config.freq_shift,
|
||||||
|
down_block_types=unet.config.down_block_types,
|
||||||
|
only_cross_attention=unet.config.only_cross_attention,
|
||||||
|
block_out_channels=unet.config.block_out_channels,
|
||||||
|
layers_per_block=unet.config.layers_per_block,
|
||||||
|
downsample_padding=unet.config.downsample_padding,
|
||||||
|
mid_block_scale_factor=unet.config.mid_block_scale_factor,
|
||||||
|
act_fn=unet.config.act_fn,
|
||||||
|
norm_num_groups=unet.config.norm_num_groups,
|
||||||
|
norm_eps=unet.config.norm_eps,
|
||||||
|
cross_attention_dim=unet.config.cross_attention_dim,
|
||||||
|
attention_head_dim=unet.config.attention_head_dim,
|
||||||
|
num_attention_heads=unet.config.num_attention_heads,
|
||||||
|
use_linear_projection=unet.config.use_linear_projection,
|
||||||
|
class_embed_type=unet.config.class_embed_type,
|
||||||
|
num_class_embeds=unet.config.num_class_embeds,
|
||||||
|
upcast_attention=unet.config.upcast_attention,
|
||||||
|
resnet_time_scale_shift=unet.config.resnet_time_scale_shift,
|
||||||
|
projection_class_embeddings_input_dim=unet.config.projection_class_embeddings_input_dim,
|
||||||
|
controlnet_conditioning_channel_order=controlnet_conditioning_channel_order,
|
||||||
|
conditioning_embedding_out_channels=conditioning_embedding_out_channels,
|
||||||
|
|
||||||
|
**controlnet_additional_kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
if load_weights_from_unet:
|
||||||
|
m, u = controlnet.conv_in.load_state_dict(cls.image_layer_filter(unet.conv_in.state_dict()), strict=False)
|
||||||
|
assert len(u) == 0
|
||||||
|
m, u = controlnet.time_proj.load_state_dict(cls.image_layer_filter(unet.time_proj.state_dict()), strict=False)
|
||||||
|
assert len(u) == 0
|
||||||
|
m, u = controlnet.time_embedding.load_state_dict(cls.image_layer_filter(unet.time_embedding.state_dict()), strict=False)
|
||||||
|
assert len(u) == 0
|
||||||
|
|
||||||
|
if controlnet.class_embedding:
|
||||||
|
m, u = controlnet.class_embedding.load_state_dict(cls.image_layer_filter(unet.class_embedding.state_dict()), strict=False)
|
||||||
|
assert len(u) == 0
|
||||||
|
m, u = controlnet.down_blocks.load_state_dict(cls.image_layer_filter(unet.down_blocks.state_dict()), strict=False)
|
||||||
|
assert len(u) == 0
|
||||||
|
m, u = controlnet.mid_block.load_state_dict(cls.image_layer_filter(unet.mid_block.state_dict()), strict=False)
|
||||||
|
assert len(u) == 0
|
||||||
|
|
||||||
|
return controlnet
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def image_layer_filter(state_dict):
|
||||||
|
new_state_dict = {}
|
||||||
|
for name, param in state_dict.items():
|
||||||
|
if "motion_modules." in name or "lora" in name: continue
|
||||||
|
new_state_dict[name] = param
|
||||||
|
return new_state_dict
|
||||||
|
|
||||||
|
# Copied from diffusers.models.unet_2d_condition.UNet2DConditionModel.set_attention_slice
|
||||||
|
def set_attention_slice(self, slice_size):
|
||||||
|
r"""
|
||||||
|
Enable sliced attention computation.
|
||||||
|
|
||||||
|
When this option is enabled, the attention module splits the input tensor in slices to compute attention in
|
||||||
|
several steps. This is useful for saving some memory in exchange for a small decrease in speed.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
slice_size (`str` or `int` or `list(int)`, *optional*, defaults to `"auto"`):
|
||||||
|
When `"auto"`, input to the attention heads is halved, so attention is computed in two steps. If
|
||||||
|
`"max"`, maximum amount of memory is saved by running only one slice at a time. If a number is
|
||||||
|
provided, uses as many slices as `attention_head_dim // slice_size`. In this case, `attention_head_dim`
|
||||||
|
must be a multiple of `slice_size`.
|
||||||
|
"""
|
||||||
|
sliceable_head_dims = []
|
||||||
|
|
||||||
|
def fn_recursive_retrieve_sliceable_dims(module: torch.nn.Module):
|
||||||
|
if hasattr(module, "set_attention_slice"):
|
||||||
|
sliceable_head_dims.append(module.sliceable_head_dim)
|
||||||
|
|
||||||
|
for child in module.children():
|
||||||
|
fn_recursive_retrieve_sliceable_dims(child)
|
||||||
|
|
||||||
|
# retrieve number of attention layers
|
||||||
|
for module in self.children():
|
||||||
|
fn_recursive_retrieve_sliceable_dims(module)
|
||||||
|
|
||||||
|
num_sliceable_layers = len(sliceable_head_dims)
|
||||||
|
|
||||||
|
if slice_size == "auto":
|
||||||
|
# half the attention head size is usually a good trade-off between
|
||||||
|
# speed and memory
|
||||||
|
slice_size = [dim // 2 for dim in sliceable_head_dims]
|
||||||
|
elif slice_size == "max":
|
||||||
|
# make smallest slice possible
|
||||||
|
slice_size = num_sliceable_layers * [1]
|
||||||
|
|
||||||
|
slice_size = num_sliceable_layers * [slice_size] if not isinstance(slice_size, list) else slice_size
|
||||||
|
|
||||||
|
if len(slice_size) != len(sliceable_head_dims):
|
||||||
|
raise ValueError(
|
||||||
|
f"You have provided {len(slice_size)}, but {self.config} has {len(sliceable_head_dims)} different"
|
||||||
|
f" attention layers. Make sure to match `len(slice_size)` to be {len(sliceable_head_dims)}."
|
||||||
|
)
|
||||||
|
|
||||||
|
for i in range(len(slice_size)):
|
||||||
|
size = slice_size[i]
|
||||||
|
dim = sliceable_head_dims[i]
|
||||||
|
if size is not None and size > dim:
|
||||||
|
raise ValueError(f"size {size} has to be smaller or equal to {dim}.")
|
||||||
|
|
||||||
|
# Recursively walk through all the children.
|
||||||
|
# Any children which exposes the set_attention_slice method
|
||||||
|
# gets the message
|
||||||
|
def fn_recursive_set_attention_slice(module: torch.nn.Module, slice_size: List[int]):
|
||||||
|
if hasattr(module, "set_attention_slice"):
|
||||||
|
module.set_attention_slice(slice_size.pop())
|
||||||
|
|
||||||
|
for child in module.children():
|
||||||
|
fn_recursive_set_attention_slice(child, slice_size)
|
||||||
|
|
||||||
|
reversed_slice_size = list(reversed(slice_size))
|
||||||
|
for module in self.children():
|
||||||
|
fn_recursive_set_attention_slice(module, reversed_slice_size)
|
||||||
|
|
||||||
|
def _set_gradient_checkpointing(self, module, value=False):
|
||||||
|
if isinstance(module, (CrossAttnDownBlock2D, DownBlock2D)):
|
||||||
|
module.gradient_checkpointing = value
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
sample: torch.FloatTensor,
|
||||||
|
timestep: Union[torch.Tensor, float, int],
|
||||||
|
encoder_hidden_states: torch.Tensor,
|
||||||
|
|
||||||
|
controlnet_cond: torch.FloatTensor,
|
||||||
|
conditioning_mask: Optional[torch.FloatTensor] = None,
|
||||||
|
|
||||||
|
conditioning_scale: float = 1.0,
|
||||||
|
class_labels: Optional[torch.Tensor] = None,
|
||||||
|
attention_mask: Optional[torch.Tensor] = None,
|
||||||
|
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||||
|
guess_mode: bool = False,
|
||||||
|
return_dict: bool = True,
|
||||||
|
) -> Union[SparseControlNetOutput, Tuple]:
|
||||||
|
|
||||||
|
# set input noise to zero
|
||||||
|
if self.set_noisy_sample_input_to_zero:
|
||||||
|
sample = torch.zeros_like(sample).to(sample.device)
|
||||||
|
|
||||||
|
# prepare attention_mask
|
||||||
|
if attention_mask is not None:
|
||||||
|
attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0
|
||||||
|
attention_mask = attention_mask.unsqueeze(1)
|
||||||
|
|
||||||
|
# 1. time
|
||||||
|
timesteps = timestep
|
||||||
|
if not torch.is_tensor(timesteps):
|
||||||
|
# TODO: this requires sync between CPU and GPU. So try to pass timesteps as tensors if you can
|
||||||
|
# This would be a good case for the `match` statement (Python 3.10+)
|
||||||
|
is_mps = sample.device.type == "mps"
|
||||||
|
if isinstance(timestep, float):
|
||||||
|
dtype = torch.float32 if is_mps else torch.float64
|
||||||
|
else:
|
||||||
|
dtype = torch.int32 if is_mps else torch.int64
|
||||||
|
timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device)
|
||||||
|
elif len(timesteps.shape) == 0:
|
||||||
|
timesteps = timesteps[None].to(sample.device)
|
||||||
|
|
||||||
|
timesteps = timesteps.repeat(sample.shape[0] // timesteps.shape[0])
|
||||||
|
encoder_hidden_states = encoder_hidden_states.repeat(sample.shape[0] // encoder_hidden_states.shape[0], 1, 1)
|
||||||
|
|
||||||
|
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||||
|
timesteps = timesteps.expand(sample.shape[0])
|
||||||
|
|
||||||
|
t_emb = self.time_proj(timesteps)
|
||||||
|
|
||||||
|
# timesteps does not contain any weights and will always return f32 tensors
|
||||||
|
# but time_embedding might actually be running in fp16. so we need to cast here.
|
||||||
|
# there might be better ways to encapsulate this.
|
||||||
|
t_emb = t_emb.to(dtype=self.dtype)
|
||||||
|
emb = self.time_embedding(t_emb)
|
||||||
|
|
||||||
|
if self.class_embedding is not None:
|
||||||
|
if class_labels is None:
|
||||||
|
raise ValueError("class_labels should be provided when num_class_embeds > 0")
|
||||||
|
|
||||||
|
if self.config.class_embed_type == "timestep":
|
||||||
|
class_labels = self.time_proj(class_labels)
|
||||||
|
|
||||||
|
class_emb = self.class_embedding(class_labels).to(dtype=self.dtype)
|
||||||
|
emb = emb + class_emb
|
||||||
|
|
||||||
|
# 2. pre-process
|
||||||
|
sample = self.conv_in(sample)
|
||||||
|
|
||||||
|
if self.concate_conditioning_mask:
|
||||||
|
controlnet_cond = torch.cat([controlnet_cond, conditioning_mask], dim=1)
|
||||||
|
controlnet_cond = self.controlnet_cond_embedding(controlnet_cond)
|
||||||
|
|
||||||
|
sample = sample + controlnet_cond
|
||||||
|
|
||||||
|
# 3. down
|
||||||
|
down_block_res_samples = (sample,)
|
||||||
|
for downsample_block in self.down_blocks:
|
||||||
|
if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention:
|
||||||
|
sample, res_samples = downsample_block(
|
||||||
|
hidden_states=sample,
|
||||||
|
temb=emb,
|
||||||
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
|
attention_mask=attention_mask,
|
||||||
|
# cross_attention_kwargs=cross_attention_kwargs,
|
||||||
|
)
|
||||||
|
else: sample, res_samples = downsample_block(hidden_states=sample, temb=emb)
|
||||||
|
|
||||||
|
down_block_res_samples += res_samples
|
||||||
|
|
||||||
|
# 4. mid
|
||||||
|
if self.mid_block is not None:
|
||||||
|
sample = self.mid_block(
|
||||||
|
sample,
|
||||||
|
emb,
|
||||||
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
|
attention_mask=attention_mask,
|
||||||
|
# cross_attention_kwargs=cross_attention_kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 5. controlnet blocks
|
||||||
|
controlnet_down_block_res_samples = ()
|
||||||
|
|
||||||
|
for down_block_res_sample, controlnet_block in zip(down_block_res_samples, self.controlnet_down_blocks):
|
||||||
|
down_block_res_sample = controlnet_block(down_block_res_sample)
|
||||||
|
controlnet_down_block_res_samples = controlnet_down_block_res_samples + (down_block_res_sample,)
|
||||||
|
|
||||||
|
down_block_res_samples = controlnet_down_block_res_samples
|
||||||
|
|
||||||
|
mid_block_res_sample = self.controlnet_mid_block(sample)
|
||||||
|
|
||||||
|
# 6. scaling
|
||||||
|
if guess_mode and not self.config.global_pool_conditions:
|
||||||
|
scales = torch.logspace(-1, 0, len(down_block_res_samples) + 1, device=sample.device) # 0.1 to 1.0
|
||||||
|
|
||||||
|
scales = scales * conditioning_scale
|
||||||
|
down_block_res_samples = [sample * scale for sample, scale in zip(down_block_res_samples, scales)]
|
||||||
|
mid_block_res_sample = mid_block_res_sample * scales[-1] # last one
|
||||||
|
else:
|
||||||
|
down_block_res_samples = [sample * conditioning_scale for sample in down_block_res_samples]
|
||||||
|
mid_block_res_sample = mid_block_res_sample * conditioning_scale
|
||||||
|
|
||||||
|
if self.config.global_pool_conditions:
|
||||||
|
down_block_res_samples = [
|
||||||
|
torch.mean(sample, dim=(2, 3), keepdim=True) for sample in down_block_res_samples
|
||||||
|
]
|
||||||
|
mid_block_res_sample = torch.mean(mid_block_res_sample, dim=(2, 3), keepdim=True)
|
||||||
|
|
||||||
|
if not return_dict:
|
||||||
|
return (down_block_res_samples, mid_block_res_sample)
|
||||||
|
|
||||||
|
return SparseControlNetOutput(
|
||||||
|
down_block_res_samples=down_block_res_samples, mid_block_res_sample=mid_block_res_sample
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def zero_module(module):
|
||||||
|
for p in module.parameters():
|
||||||
|
nn.init.zeros_(p)
|
||||||
|
return module
|
||||||
@@ -0,0 +1,526 @@
|
|||||||
|
# Adapted from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/unet_2d_condition.py
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
import os
|
||||||
|
import json
|
||||||
|
import pdb
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.utils.checkpoint
|
||||||
|
|
||||||
|
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||||
|
from diffusers import ModelMixin
|
||||||
|
from diffusers.utils import BaseOutput, logging
|
||||||
|
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
|
||||||
|
from .unet_blocks import (
|
||||||
|
CrossAttnDownBlock3D,
|
||||||
|
CrossAttnUpBlock3D,
|
||||||
|
DownBlock3D,
|
||||||
|
UNetMidBlock3DCrossAttn,
|
||||||
|
UpBlock3D,
|
||||||
|
get_down_block,
|
||||||
|
get_up_block,
|
||||||
|
)
|
||||||
|
from .resnet import InflatedConv3d, InflatedGroupNorm
|
||||||
|
|
||||||
|
|
||||||
|
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class UNet3DConditionOutput(BaseOutput):
|
||||||
|
sample: torch.FloatTensor
|
||||||
|
|
||||||
|
|
||||||
|
class UNet3DConditionModel(ModelMixin, ConfigMixin):
|
||||||
|
_supports_gradient_checkpointing = True
|
||||||
|
|
||||||
|
@register_to_config
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
sample_size: Optional[int] = None,
|
||||||
|
in_channels: int = 4,
|
||||||
|
out_channels: int = 4,
|
||||||
|
center_input_sample: bool = False,
|
||||||
|
flip_sin_to_cos: bool = True,
|
||||||
|
freq_shift: int = 0,
|
||||||
|
down_block_types: Tuple[str] = (
|
||||||
|
"CrossAttnDownBlock3D",
|
||||||
|
"CrossAttnDownBlock3D",
|
||||||
|
"CrossAttnDownBlock3D",
|
||||||
|
"DownBlock3D",
|
||||||
|
),
|
||||||
|
mid_block_type: str = "UNetMidBlock3DCrossAttn",
|
||||||
|
up_block_types: Tuple[str] = (
|
||||||
|
"UpBlock3D",
|
||||||
|
"CrossAttnUpBlock3D",
|
||||||
|
"CrossAttnUpBlock3D",
|
||||||
|
"CrossAttnUpBlock3D"
|
||||||
|
),
|
||||||
|
only_cross_attention: Union[bool, Tuple[bool]] = False,
|
||||||
|
block_out_channels: Tuple[int] = (320, 640, 1280, 1280),
|
||||||
|
layers_per_block: int = 2,
|
||||||
|
downsample_padding: int = 1,
|
||||||
|
mid_block_scale_factor: float = 1,
|
||||||
|
act_fn: str = "silu",
|
||||||
|
norm_num_groups: int = 32,
|
||||||
|
norm_eps: float = 1e-5,
|
||||||
|
cross_attention_dim: int = 1280,
|
||||||
|
attention_head_dim: Union[int, Tuple[int]] = 8,
|
||||||
|
dual_cross_attention: bool = False,
|
||||||
|
use_linear_projection: bool = False,
|
||||||
|
class_embed_type: Optional[str] = None,
|
||||||
|
num_class_embeds: Optional[int] = None,
|
||||||
|
upcast_attention: bool = False,
|
||||||
|
resnet_time_scale_shift: str = "default",
|
||||||
|
|
||||||
|
use_inflated_groupnorm=False,
|
||||||
|
|
||||||
|
# Additional
|
||||||
|
use_motion_module = False,
|
||||||
|
motion_module_resolutions = ( 1,2,4,8 ),
|
||||||
|
motion_module_mid_block = False,
|
||||||
|
motion_module_decoder_only = False,
|
||||||
|
motion_module_type = None,
|
||||||
|
motion_module_kwargs = {},
|
||||||
|
unet_use_cross_frame_attention = False,
|
||||||
|
unet_use_temporal_attention = False,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.sample_size = sample_size
|
||||||
|
time_embed_dim = block_out_channels[0] * 4
|
||||||
|
|
||||||
|
# input
|
||||||
|
self.conv_in = InflatedConv3d(in_channels, block_out_channels[0], kernel_size=3, padding=(1, 1))
|
||||||
|
|
||||||
|
# time
|
||||||
|
self.time_proj = Timesteps(block_out_channels[0], flip_sin_to_cos, freq_shift)
|
||||||
|
timestep_input_dim = block_out_channels[0]
|
||||||
|
|
||||||
|
self.time_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim)
|
||||||
|
|
||||||
|
# class embedding
|
||||||
|
if class_embed_type is None and num_class_embeds is not None:
|
||||||
|
self.class_embedding = nn.Embedding(num_class_embeds, time_embed_dim)
|
||||||
|
elif class_embed_type == "timestep":
|
||||||
|
self.class_embedding = TimestepEmbedding(timestep_input_dim, time_embed_dim)
|
||||||
|
elif class_embed_type == "identity":
|
||||||
|
self.class_embedding = nn.Identity(time_embed_dim, time_embed_dim)
|
||||||
|
else:
|
||||||
|
self.class_embedding = None
|
||||||
|
|
||||||
|
self.down_blocks = nn.ModuleList([])
|
||||||
|
self.mid_block = None
|
||||||
|
self.up_blocks = nn.ModuleList([])
|
||||||
|
|
||||||
|
if isinstance(only_cross_attention, bool):
|
||||||
|
only_cross_attention = [only_cross_attention] * len(down_block_types)
|
||||||
|
|
||||||
|
if isinstance(attention_head_dim, int):
|
||||||
|
attention_head_dim = (attention_head_dim,) * len(down_block_types)
|
||||||
|
|
||||||
|
# down
|
||||||
|
output_channel = block_out_channels[0]
|
||||||
|
for i, down_block_type in enumerate(down_block_types):
|
||||||
|
res = 2 ** i
|
||||||
|
input_channel = output_channel
|
||||||
|
output_channel = block_out_channels[i]
|
||||||
|
is_final_block = i == len(block_out_channels) - 1
|
||||||
|
|
||||||
|
down_block = get_down_block(
|
||||||
|
down_block_type,
|
||||||
|
num_layers=layers_per_block,
|
||||||
|
in_channels=input_channel,
|
||||||
|
out_channels=output_channel,
|
||||||
|
temb_channels=time_embed_dim,
|
||||||
|
add_downsample=not is_final_block,
|
||||||
|
resnet_eps=norm_eps,
|
||||||
|
resnet_act_fn=act_fn,
|
||||||
|
resnet_groups=norm_num_groups,
|
||||||
|
cross_attention_dim=cross_attention_dim,
|
||||||
|
attn_num_head_channels=attention_head_dim[i],
|
||||||
|
downsample_padding=downsample_padding,
|
||||||
|
dual_cross_attention=dual_cross_attention,
|
||||||
|
use_linear_projection=use_linear_projection,
|
||||||
|
only_cross_attention=only_cross_attention[i],
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=unet_use_cross_frame_attention,
|
||||||
|
unet_use_temporal_attention=unet_use_temporal_attention,
|
||||||
|
use_inflated_groupnorm=use_inflated_groupnorm,
|
||||||
|
|
||||||
|
use_motion_module=use_motion_module and (res in motion_module_resolutions) and (not motion_module_decoder_only),
|
||||||
|
motion_module_type=motion_module_type,
|
||||||
|
motion_module_kwargs=motion_module_kwargs,
|
||||||
|
)
|
||||||
|
self.down_blocks.append(down_block)
|
||||||
|
|
||||||
|
# mid
|
||||||
|
if mid_block_type == "UNetMidBlock3DCrossAttn":
|
||||||
|
self.mid_block = UNetMidBlock3DCrossAttn(
|
||||||
|
in_channels=block_out_channels[-1],
|
||||||
|
temb_channels=time_embed_dim,
|
||||||
|
resnet_eps=norm_eps,
|
||||||
|
resnet_act_fn=act_fn,
|
||||||
|
output_scale_factor=mid_block_scale_factor,
|
||||||
|
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||||
|
cross_attention_dim=cross_attention_dim,
|
||||||
|
attn_num_head_channels=attention_head_dim[-1],
|
||||||
|
resnet_groups=norm_num_groups,
|
||||||
|
dual_cross_attention=dual_cross_attention,
|
||||||
|
use_linear_projection=use_linear_projection,
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=unet_use_cross_frame_attention,
|
||||||
|
unet_use_temporal_attention=unet_use_temporal_attention,
|
||||||
|
use_inflated_groupnorm=use_inflated_groupnorm,
|
||||||
|
|
||||||
|
use_motion_module=use_motion_module and motion_module_mid_block,
|
||||||
|
motion_module_type=motion_module_type,
|
||||||
|
motion_module_kwargs=motion_module_kwargs,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"unknown mid_block_type : {mid_block_type}")
|
||||||
|
|
||||||
|
# count how many layers upsample the videos
|
||||||
|
self.num_upsamplers = 0
|
||||||
|
|
||||||
|
# up
|
||||||
|
reversed_block_out_channels = list(reversed(block_out_channels))
|
||||||
|
reversed_attention_head_dim = list(reversed(attention_head_dim))
|
||||||
|
only_cross_attention = list(reversed(only_cross_attention))
|
||||||
|
output_channel = reversed_block_out_channels[0]
|
||||||
|
for i, up_block_type in enumerate(up_block_types):
|
||||||
|
res = 2 ** (3 - i)
|
||||||
|
is_final_block = i == len(block_out_channels) - 1
|
||||||
|
|
||||||
|
prev_output_channel = output_channel
|
||||||
|
output_channel = reversed_block_out_channels[i]
|
||||||
|
input_channel = reversed_block_out_channels[min(i + 1, len(block_out_channels) - 1)]
|
||||||
|
|
||||||
|
# add upsample block for all BUT final layer
|
||||||
|
if not is_final_block:
|
||||||
|
add_upsample = True
|
||||||
|
self.num_upsamplers += 1
|
||||||
|
else:
|
||||||
|
add_upsample = False
|
||||||
|
|
||||||
|
up_block = get_up_block(
|
||||||
|
up_block_type,
|
||||||
|
num_layers=layers_per_block + 1,
|
||||||
|
in_channels=input_channel,
|
||||||
|
out_channels=output_channel,
|
||||||
|
prev_output_channel=prev_output_channel,
|
||||||
|
temb_channels=time_embed_dim,
|
||||||
|
add_upsample=add_upsample,
|
||||||
|
resnet_eps=norm_eps,
|
||||||
|
resnet_act_fn=act_fn,
|
||||||
|
resnet_groups=norm_num_groups,
|
||||||
|
cross_attention_dim=cross_attention_dim,
|
||||||
|
attn_num_head_channels=reversed_attention_head_dim[i],
|
||||||
|
dual_cross_attention=dual_cross_attention,
|
||||||
|
use_linear_projection=use_linear_projection,
|
||||||
|
only_cross_attention=only_cross_attention[i],
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=unet_use_cross_frame_attention,
|
||||||
|
unet_use_temporal_attention=unet_use_temporal_attention,
|
||||||
|
use_inflated_groupnorm=use_inflated_groupnorm,
|
||||||
|
|
||||||
|
use_motion_module=use_motion_module and (res in motion_module_resolutions),
|
||||||
|
motion_module_type=motion_module_type,
|
||||||
|
motion_module_kwargs=motion_module_kwargs,
|
||||||
|
)
|
||||||
|
self.up_blocks.append(up_block)
|
||||||
|
prev_output_channel = output_channel
|
||||||
|
|
||||||
|
# out
|
||||||
|
if use_inflated_groupnorm:
|
||||||
|
self.conv_norm_out = InflatedGroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=norm_eps)
|
||||||
|
else:
|
||||||
|
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=norm_eps)
|
||||||
|
self.conv_act = nn.SiLU()
|
||||||
|
self.conv_out = InflatedConv3d(block_out_channels[0], out_channels, kernel_size=3, padding=1)
|
||||||
|
|
||||||
|
def set_attention_slice(self, slice_size):
|
||||||
|
r"""
|
||||||
|
Enable sliced attention computation.
|
||||||
|
|
||||||
|
When this option is enabled, the attention module will split the input tensor in slices, to compute attention
|
||||||
|
in several steps. This is useful to save some memory in exchange for a small speed decrease.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
slice_size (`str` or `int` or `list(int)`, *optional*, defaults to `"auto"`):
|
||||||
|
When `"auto"`, halves the input to the attention heads, so attention will be computed in two steps. If
|
||||||
|
`"max"`, maxium amount of memory will be saved by running only one slice at a time. If a number is
|
||||||
|
provided, uses as many slices as `attention_head_dim // slice_size`. In this case, `attention_head_dim`
|
||||||
|
must be a multiple of `slice_size`.
|
||||||
|
"""
|
||||||
|
sliceable_head_dims = []
|
||||||
|
|
||||||
|
def fn_recursive_retrieve_slicable_dims(module: torch.nn.Module):
|
||||||
|
if hasattr(module, "set_attention_slice"):
|
||||||
|
sliceable_head_dims.append(module.sliceable_head_dim)
|
||||||
|
|
||||||
|
for child in module.children():
|
||||||
|
fn_recursive_retrieve_slicable_dims(child)
|
||||||
|
|
||||||
|
# retrieve number of attention layers
|
||||||
|
for module in self.children():
|
||||||
|
fn_recursive_retrieve_slicable_dims(module)
|
||||||
|
|
||||||
|
num_slicable_layers = len(sliceable_head_dims)
|
||||||
|
|
||||||
|
if slice_size == "auto":
|
||||||
|
# half the attention head size is usually a good trade-off between
|
||||||
|
# speed and memory
|
||||||
|
slice_size = [dim // 2 for dim in sliceable_head_dims]
|
||||||
|
elif slice_size == "max":
|
||||||
|
# make smallest slice possible
|
||||||
|
slice_size = num_slicable_layers * [1]
|
||||||
|
|
||||||
|
slice_size = num_slicable_layers * [slice_size] if not isinstance(slice_size, list) else slice_size
|
||||||
|
|
||||||
|
if len(slice_size) != len(sliceable_head_dims):
|
||||||
|
raise ValueError(
|
||||||
|
f"You have provided {len(slice_size)}, but {self.config} has {len(sliceable_head_dims)} different"
|
||||||
|
f" attention layers. Make sure to match `len(slice_size)` to be {len(sliceable_head_dims)}."
|
||||||
|
)
|
||||||
|
|
||||||
|
for i in range(len(slice_size)):
|
||||||
|
size = slice_size[i]
|
||||||
|
dim = sliceable_head_dims[i]
|
||||||
|
if size is not None and size > dim:
|
||||||
|
raise ValueError(f"size {size} has to be smaller or equal to {dim}.")
|
||||||
|
|
||||||
|
# Recursively walk through all the children.
|
||||||
|
# Any children which exposes the set_attention_slice method
|
||||||
|
# gets the message
|
||||||
|
def fn_recursive_set_attention_slice(module: torch.nn.Module, slice_size: List[int]):
|
||||||
|
if hasattr(module, "set_attention_slice"):
|
||||||
|
module.set_attention_slice(slice_size.pop())
|
||||||
|
|
||||||
|
for child in module.children():
|
||||||
|
fn_recursive_set_attention_slice(child, slice_size)
|
||||||
|
|
||||||
|
reversed_slice_size = list(reversed(slice_size))
|
||||||
|
for module in self.children():
|
||||||
|
fn_recursive_set_attention_slice(module, reversed_slice_size)
|
||||||
|
|
||||||
|
def _set_gradient_checkpointing(self, module, value=False):
|
||||||
|
if isinstance(module, (CrossAttnDownBlock3D, DownBlock3D, CrossAttnUpBlock3D, UpBlock3D)):
|
||||||
|
module.gradient_checkpointing = value
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
sample: torch.FloatTensor,
|
||||||
|
timestep: Union[torch.Tensor, float, int],
|
||||||
|
encoder_hidden_states: torch.Tensor,
|
||||||
|
class_labels: Optional[torch.Tensor] = None,
|
||||||
|
attention_mask: Optional[torch.Tensor] = None,
|
||||||
|
|
||||||
|
# support controlnet
|
||||||
|
down_block_additional_residuals: Optional[Tuple[torch.Tensor]] = None,
|
||||||
|
mid_block_additional_residual: Optional[torch.Tensor] = None,
|
||||||
|
|
||||||
|
return_dict: bool = True,
|
||||||
|
) -> Union[UNet3DConditionOutput, Tuple]:
|
||||||
|
r"""
|
||||||
|
Args:
|
||||||
|
sample (`torch.FloatTensor`): (batch, channel, height, width) noisy inputs tensor
|
||||||
|
timestep (`torch.FloatTensor` or `float` or `int`): (batch) timesteps
|
||||||
|
encoder_hidden_states (`torch.FloatTensor`): (batch, sequence_length, feature_dim) encoder hidden states
|
||||||
|
return_dict (`bool`, *optional*, defaults to `True`):
|
||||||
|
Whether or not to return a [`models.unet_2d_condition.UNet2DConditionOutput`] instead of a plain tuple.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[`~models.unet_2d_condition.UNet2DConditionOutput`] or `tuple`:
|
||||||
|
[`~models.unet_2d_condition.UNet2DConditionOutput`] if `return_dict` is True, otherwise a `tuple`. When
|
||||||
|
returning a tuple, the first element is the sample tensor.
|
||||||
|
"""
|
||||||
|
# By default samples have to be AT least a multiple of the overall upsampling factor.
|
||||||
|
# The overall upsampling factor is equal to 2 ** (# num of upsampling layears).
|
||||||
|
# However, the upsampling interpolation output size can be forced to fit any upsampling size
|
||||||
|
# on the fly if necessary.
|
||||||
|
default_overall_up_factor = 2**self.num_upsamplers
|
||||||
|
|
||||||
|
# upsample size should be forwarded when sample is not a multiple of `default_overall_up_factor`
|
||||||
|
forward_upsample_size = False
|
||||||
|
upsample_size = None
|
||||||
|
|
||||||
|
if any(s % default_overall_up_factor != 0 for s in sample.shape[-2:]):
|
||||||
|
logger.info("Forward upsample size to force interpolation output size.")
|
||||||
|
forward_upsample_size = True
|
||||||
|
|
||||||
|
# prepare attention_mask
|
||||||
|
if attention_mask is not None:
|
||||||
|
attention_mask = (1 - attention_mask.to(sample.dtype)) * -10000.0
|
||||||
|
attention_mask = attention_mask.unsqueeze(1)
|
||||||
|
|
||||||
|
# center input if necessary
|
||||||
|
if self.config.center_input_sample:
|
||||||
|
sample = 2 * sample - 1.0
|
||||||
|
|
||||||
|
# time
|
||||||
|
timesteps = timestep
|
||||||
|
if not torch.is_tensor(timesteps):
|
||||||
|
# This would be a good case for the `match` statement (Python 3.10+)
|
||||||
|
is_mps = sample.device.type == "mps"
|
||||||
|
if isinstance(timestep, float):
|
||||||
|
dtype = torch.float32 if is_mps else torch.float64
|
||||||
|
else:
|
||||||
|
dtype = torch.int32 if is_mps else torch.int64
|
||||||
|
timesteps = torch.tensor([timesteps], dtype=dtype, device=sample.device)
|
||||||
|
elif len(timesteps.shape) == 0:
|
||||||
|
timesteps = timesteps[None].to(sample.device)
|
||||||
|
|
||||||
|
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||||
|
timesteps = timesteps.expand(sample.shape[0])
|
||||||
|
|
||||||
|
t_emb = self.time_proj(timesteps)
|
||||||
|
|
||||||
|
# timesteps does not contain any weights and will always return f32 tensors
|
||||||
|
# but time_embedding might actually be running in fp16. so we need to cast here.
|
||||||
|
# there might be better ways to encapsulate this.
|
||||||
|
t_emb = t_emb.to(dtype=self.dtype)
|
||||||
|
emb = self.time_embedding(t_emb)
|
||||||
|
|
||||||
|
if self.class_embedding is not None:
|
||||||
|
if class_labels is None:
|
||||||
|
raise ValueError("class_labels should be provided when num_class_embeds > 0")
|
||||||
|
|
||||||
|
if self.config.class_embed_type == "timestep":
|
||||||
|
class_labels = self.time_proj(class_labels)
|
||||||
|
|
||||||
|
class_emb = self.class_embedding(class_labels).to(dtype=self.dtype)
|
||||||
|
emb = emb + class_emb
|
||||||
|
|
||||||
|
# pre-process
|
||||||
|
sample = self.conv_in(sample)
|
||||||
|
|
||||||
|
# down
|
||||||
|
down_block_res_samples = (sample,)
|
||||||
|
for downsample_block in self.down_blocks:
|
||||||
|
if hasattr(downsample_block, "has_cross_attention") and downsample_block.has_cross_attention:
|
||||||
|
sample, res_samples = downsample_block(
|
||||||
|
hidden_states=sample,
|
||||||
|
temb=emb,
|
||||||
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
|
attention_mask=attention_mask,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
sample, res_samples = downsample_block(hidden_states=sample, temb=emb, encoder_hidden_states=encoder_hidden_states)
|
||||||
|
|
||||||
|
down_block_res_samples += res_samples
|
||||||
|
|
||||||
|
# support controlnet
|
||||||
|
down_block_res_samples = list(down_block_res_samples)
|
||||||
|
if down_block_additional_residuals is not None:
|
||||||
|
for i, down_block_additional_residual in enumerate(down_block_additional_residuals):
|
||||||
|
if down_block_additional_residual.dim() == 4: # boardcast
|
||||||
|
down_block_additional_residual = down_block_additional_residual.unsqueeze(2)
|
||||||
|
down_block_res_samples[i] = down_block_res_samples[i] + down_block_additional_residual
|
||||||
|
|
||||||
|
# mid
|
||||||
|
sample = self.mid_block(
|
||||||
|
sample, emb, encoder_hidden_states=encoder_hidden_states, attention_mask=attention_mask
|
||||||
|
)
|
||||||
|
|
||||||
|
# support controlnet
|
||||||
|
if mid_block_additional_residual is not None:
|
||||||
|
if mid_block_additional_residual.dim() == 4: # boardcast
|
||||||
|
mid_block_additional_residual = mid_block_additional_residual.unsqueeze(2)
|
||||||
|
sample = sample + mid_block_additional_residual
|
||||||
|
|
||||||
|
# up
|
||||||
|
for i, upsample_block in enumerate(self.up_blocks):
|
||||||
|
is_final_block = i == len(self.up_blocks) - 1
|
||||||
|
|
||||||
|
res_samples = down_block_res_samples[-len(upsample_block.resnets) :]
|
||||||
|
down_block_res_samples = down_block_res_samples[: -len(upsample_block.resnets)]
|
||||||
|
|
||||||
|
# if we have not reached the final block and need to forward the
|
||||||
|
# upsample size, we do it here
|
||||||
|
if not is_final_block and forward_upsample_size:
|
||||||
|
upsample_size = down_block_res_samples[-1].shape[2:]
|
||||||
|
|
||||||
|
if hasattr(upsample_block, "has_cross_attention") and upsample_block.has_cross_attention:
|
||||||
|
sample = upsample_block(
|
||||||
|
hidden_states=sample,
|
||||||
|
temb=emb,
|
||||||
|
res_hidden_states_tuple=res_samples,
|
||||||
|
encoder_hidden_states=encoder_hidden_states,
|
||||||
|
upsample_size=upsample_size,
|
||||||
|
attention_mask=attention_mask,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
sample = upsample_block(
|
||||||
|
hidden_states=sample, temb=emb, res_hidden_states_tuple=res_samples, upsample_size=upsample_size, encoder_hidden_states=encoder_hidden_states,
|
||||||
|
)
|
||||||
|
|
||||||
|
# post-process
|
||||||
|
sample = self.conv_norm_out(sample)
|
||||||
|
sample = self.conv_act(sample)
|
||||||
|
sample = self.conv_out(sample)
|
||||||
|
|
||||||
|
if not return_dict:
|
||||||
|
return (sample,)
|
||||||
|
|
||||||
|
return UNet3DConditionOutput(sample=sample)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_pretrained_2d(cls, pretrained_model_path, subfolder=None, unet_additional_kwargs=None):
|
||||||
|
if subfolder is not None:
|
||||||
|
pretrained_model_path = os.path.join(pretrained_model_path, subfolder)
|
||||||
|
print(f"loaded 3D unet's pretrained weights from {pretrained_model_path} ...")
|
||||||
|
|
||||||
|
config_file = os.path.join(pretrained_model_path, 'config.json')
|
||||||
|
if not os.path.isfile(config_file):
|
||||||
|
raise RuntimeError(f"{config_file} does not exist")
|
||||||
|
with open(config_file, "r") as f:
|
||||||
|
config = json.load(f)
|
||||||
|
config["_class_name"] = cls.__name__
|
||||||
|
config["down_block_types"] = [
|
||||||
|
"CrossAttnDownBlock3D",
|
||||||
|
"CrossAttnDownBlock3D",
|
||||||
|
"CrossAttnDownBlock3D",
|
||||||
|
"DownBlock3D"
|
||||||
|
]
|
||||||
|
config["up_block_types"] = [
|
||||||
|
"UpBlock3D",
|
||||||
|
"CrossAttnUpBlock3D",
|
||||||
|
"CrossAttnUpBlock3D",
|
||||||
|
"CrossAttnUpBlock3D"
|
||||||
|
]
|
||||||
|
|
||||||
|
from diffusers.utils import WEIGHTS_NAME, SAFETENSORS_WEIGHTS_NAME
|
||||||
|
model = cls.from_config(config, **unet_additional_kwargs)
|
||||||
|
|
||||||
|
model_file = os.path.join(pretrained_model_path, WEIGHTS_NAME)
|
||||||
|
model_file_safe = os.path.join(pretrained_model_path, SAFETENSORS_WEIGHTS_NAME)
|
||||||
|
|
||||||
|
if os.path.isfile(model_file_safe):
|
||||||
|
model_file = model_file_safe
|
||||||
|
|
||||||
|
if not os.path.isfile(model_file):
|
||||||
|
raise RuntimeError(f"{model_file} does not exist")
|
||||||
|
|
||||||
|
if SAFETENSORS_WEIGHTS_NAME in model_file:
|
||||||
|
from safetensors.torch import load_file
|
||||||
|
state_dict = load_file(model_file)
|
||||||
|
else:
|
||||||
|
state_dict = torch.load(model_file, map_location="cpu")
|
||||||
|
|
||||||
|
m, u = model.load_state_dict(state_dict, strict=False)
|
||||||
|
print(f"### missing keys: {len(m)}; \n### unexpected keys: {len(u)};")
|
||||||
|
|
||||||
|
params = [p.numel() if "motion_modules." in n else 0 for n, p in model.named_parameters()]
|
||||||
|
print(f"### Motion Module Parameters: {sum(params) / 1e6} M")
|
||||||
|
|
||||||
|
return model
|
||||||
@@ -0,0 +1,764 @@
|
|||||||
|
# Adapted from https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/unet_2d_blocks.py
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
from .attention import Transformer3DModel
|
||||||
|
from .resnet import Downsample3D, ResnetBlock3D, Upsample3D
|
||||||
|
from .motion_module import get_motion_module
|
||||||
|
|
||||||
|
import pdb
|
||||||
|
|
||||||
|
def checkpoint_no_reentrant(*args, **kwargs):
|
||||||
|
kwargs['use_reentrant'] = False
|
||||||
|
return torch.utils.checkpoint.checkpoint(*args, **kwargs)
|
||||||
|
|
||||||
|
def get_down_block(
|
||||||
|
down_block_type,
|
||||||
|
num_layers,
|
||||||
|
in_channels,
|
||||||
|
out_channels,
|
||||||
|
temb_channels,
|
||||||
|
add_downsample,
|
||||||
|
resnet_eps,
|
||||||
|
resnet_act_fn,
|
||||||
|
attn_num_head_channels,
|
||||||
|
resnet_groups=None,
|
||||||
|
cross_attention_dim=None,
|
||||||
|
downsample_padding=None,
|
||||||
|
dual_cross_attention=False,
|
||||||
|
use_linear_projection=False,
|
||||||
|
only_cross_attention=False,
|
||||||
|
upcast_attention=False,
|
||||||
|
resnet_time_scale_shift="default",
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=False,
|
||||||
|
unet_use_temporal_attention=False,
|
||||||
|
use_inflated_groupnorm=False,
|
||||||
|
|
||||||
|
use_motion_module=None,
|
||||||
|
|
||||||
|
motion_module_type=None,
|
||||||
|
motion_module_kwargs=None,
|
||||||
|
):
|
||||||
|
down_block_type = down_block_type[7:] if down_block_type.startswith("UNetRes") else down_block_type
|
||||||
|
if down_block_type == "DownBlock3D":
|
||||||
|
return DownBlock3D(
|
||||||
|
num_layers=num_layers,
|
||||||
|
in_channels=in_channels,
|
||||||
|
out_channels=out_channels,
|
||||||
|
temb_channels=temb_channels,
|
||||||
|
add_downsample=add_downsample,
|
||||||
|
resnet_eps=resnet_eps,
|
||||||
|
resnet_act_fn=resnet_act_fn,
|
||||||
|
resnet_groups=resnet_groups,
|
||||||
|
downsample_padding=downsample_padding,
|
||||||
|
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||||
|
|
||||||
|
use_inflated_groupnorm=use_inflated_groupnorm,
|
||||||
|
|
||||||
|
use_motion_module=use_motion_module,
|
||||||
|
motion_module_type=motion_module_type,
|
||||||
|
motion_module_kwargs=motion_module_kwargs,
|
||||||
|
)
|
||||||
|
elif down_block_type == "CrossAttnDownBlock3D":
|
||||||
|
if cross_attention_dim is None:
|
||||||
|
raise ValueError("cross_attention_dim must be specified for CrossAttnDownBlock3D")
|
||||||
|
return CrossAttnDownBlock3D(
|
||||||
|
num_layers=num_layers,
|
||||||
|
in_channels=in_channels,
|
||||||
|
out_channels=out_channels,
|
||||||
|
temb_channels=temb_channels,
|
||||||
|
add_downsample=add_downsample,
|
||||||
|
resnet_eps=resnet_eps,
|
||||||
|
resnet_act_fn=resnet_act_fn,
|
||||||
|
resnet_groups=resnet_groups,
|
||||||
|
downsample_padding=downsample_padding,
|
||||||
|
cross_attention_dim=cross_attention_dim,
|
||||||
|
attn_num_head_channels=attn_num_head_channels,
|
||||||
|
dual_cross_attention=dual_cross_attention,
|
||||||
|
use_linear_projection=use_linear_projection,
|
||||||
|
only_cross_attention=only_cross_attention,
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=unet_use_cross_frame_attention,
|
||||||
|
unet_use_temporal_attention=unet_use_temporal_attention,
|
||||||
|
use_inflated_groupnorm=use_inflated_groupnorm,
|
||||||
|
|
||||||
|
use_motion_module=use_motion_module,
|
||||||
|
motion_module_type=motion_module_type,
|
||||||
|
motion_module_kwargs=motion_module_kwargs,
|
||||||
|
)
|
||||||
|
raise ValueError(f"{down_block_type} does not exist.")
|
||||||
|
|
||||||
|
|
||||||
|
def get_up_block(
|
||||||
|
up_block_type,
|
||||||
|
num_layers,
|
||||||
|
in_channels,
|
||||||
|
out_channels,
|
||||||
|
prev_output_channel,
|
||||||
|
temb_channels,
|
||||||
|
add_upsample,
|
||||||
|
resnet_eps,
|
||||||
|
resnet_act_fn,
|
||||||
|
attn_num_head_channels,
|
||||||
|
resnet_groups=None,
|
||||||
|
cross_attention_dim=None,
|
||||||
|
dual_cross_attention=False,
|
||||||
|
use_linear_projection=False,
|
||||||
|
only_cross_attention=False,
|
||||||
|
upcast_attention=False,
|
||||||
|
resnet_time_scale_shift="default",
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=False,
|
||||||
|
unet_use_temporal_attention=False,
|
||||||
|
use_inflated_groupnorm=False,
|
||||||
|
|
||||||
|
use_motion_module=None,
|
||||||
|
motion_module_type=None,
|
||||||
|
motion_module_kwargs=None,
|
||||||
|
):
|
||||||
|
up_block_type = up_block_type[7:] if up_block_type.startswith("UNetRes") else up_block_type
|
||||||
|
if up_block_type == "UpBlock3D":
|
||||||
|
return UpBlock3D(
|
||||||
|
num_layers=num_layers,
|
||||||
|
in_channels=in_channels,
|
||||||
|
out_channels=out_channels,
|
||||||
|
prev_output_channel=prev_output_channel,
|
||||||
|
temb_channels=temb_channels,
|
||||||
|
add_upsample=add_upsample,
|
||||||
|
resnet_eps=resnet_eps,
|
||||||
|
resnet_act_fn=resnet_act_fn,
|
||||||
|
resnet_groups=resnet_groups,
|
||||||
|
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||||
|
|
||||||
|
use_inflated_groupnorm=use_inflated_groupnorm,
|
||||||
|
|
||||||
|
use_motion_module=use_motion_module,
|
||||||
|
motion_module_type=motion_module_type,
|
||||||
|
motion_module_kwargs=motion_module_kwargs,
|
||||||
|
)
|
||||||
|
elif up_block_type == "CrossAttnUpBlock3D":
|
||||||
|
if cross_attention_dim is None:
|
||||||
|
raise ValueError("cross_attention_dim must be specified for CrossAttnUpBlock3D")
|
||||||
|
return CrossAttnUpBlock3D(
|
||||||
|
num_layers=num_layers,
|
||||||
|
in_channels=in_channels,
|
||||||
|
out_channels=out_channels,
|
||||||
|
prev_output_channel=prev_output_channel,
|
||||||
|
temb_channels=temb_channels,
|
||||||
|
add_upsample=add_upsample,
|
||||||
|
resnet_eps=resnet_eps,
|
||||||
|
resnet_act_fn=resnet_act_fn,
|
||||||
|
resnet_groups=resnet_groups,
|
||||||
|
cross_attention_dim=cross_attention_dim,
|
||||||
|
attn_num_head_channels=attn_num_head_channels,
|
||||||
|
dual_cross_attention=dual_cross_attention,
|
||||||
|
use_linear_projection=use_linear_projection,
|
||||||
|
only_cross_attention=only_cross_attention,
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
resnet_time_scale_shift=resnet_time_scale_shift,
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=unet_use_cross_frame_attention,
|
||||||
|
unet_use_temporal_attention=unet_use_temporal_attention,
|
||||||
|
use_inflated_groupnorm=use_inflated_groupnorm,
|
||||||
|
|
||||||
|
use_motion_module=use_motion_module,
|
||||||
|
motion_module_type=motion_module_type,
|
||||||
|
motion_module_kwargs=motion_module_kwargs,
|
||||||
|
)
|
||||||
|
raise ValueError(f"{up_block_type} does not exist.")
|
||||||
|
|
||||||
|
|
||||||
|
class UNetMidBlock3DCrossAttn(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels: int,
|
||||||
|
temb_channels: int,
|
||||||
|
dropout: float = 0.0,
|
||||||
|
num_layers: int = 1,
|
||||||
|
resnet_eps: float = 1e-6,
|
||||||
|
resnet_time_scale_shift: str = "default",
|
||||||
|
resnet_act_fn: str = "swish",
|
||||||
|
resnet_groups: int = 32,
|
||||||
|
resnet_pre_norm: bool = True,
|
||||||
|
attn_num_head_channels=1,
|
||||||
|
output_scale_factor=1.0,
|
||||||
|
cross_attention_dim=1280,
|
||||||
|
dual_cross_attention=False,
|
||||||
|
use_linear_projection=False,
|
||||||
|
upcast_attention=False,
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=False,
|
||||||
|
unet_use_temporal_attention=False,
|
||||||
|
use_inflated_groupnorm=False,
|
||||||
|
|
||||||
|
use_motion_module=None,
|
||||||
|
|
||||||
|
motion_module_type=None,
|
||||||
|
motion_module_kwargs=None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.has_cross_attention = True
|
||||||
|
self.attn_num_head_channels = attn_num_head_channels
|
||||||
|
resnet_groups = resnet_groups if resnet_groups is not None else min(in_channels // 4, 32)
|
||||||
|
|
||||||
|
# there is always at least one resnet
|
||||||
|
resnets = [
|
||||||
|
ResnetBlock3D(
|
||||||
|
in_channels=in_channels,
|
||||||
|
out_channels=in_channels,
|
||||||
|
temb_channels=temb_channels,
|
||||||
|
eps=resnet_eps,
|
||||||
|
groups=resnet_groups,
|
||||||
|
dropout=dropout,
|
||||||
|
time_embedding_norm=resnet_time_scale_shift,
|
||||||
|
non_linearity=resnet_act_fn,
|
||||||
|
output_scale_factor=output_scale_factor,
|
||||||
|
pre_norm=resnet_pre_norm,
|
||||||
|
|
||||||
|
use_inflated_groupnorm=use_inflated_groupnorm,
|
||||||
|
)
|
||||||
|
]
|
||||||
|
attentions = []
|
||||||
|
motion_modules = []
|
||||||
|
|
||||||
|
for _ in range(num_layers):
|
||||||
|
if dual_cross_attention:
|
||||||
|
raise NotImplementedError
|
||||||
|
attentions.append(
|
||||||
|
Transformer3DModel(
|
||||||
|
attn_num_head_channels,
|
||||||
|
in_channels // attn_num_head_channels,
|
||||||
|
in_channels=in_channels,
|
||||||
|
num_layers=1,
|
||||||
|
cross_attention_dim=cross_attention_dim,
|
||||||
|
norm_num_groups=resnet_groups,
|
||||||
|
use_linear_projection=use_linear_projection,
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=unet_use_cross_frame_attention,
|
||||||
|
unet_use_temporal_attention=unet_use_temporal_attention,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
motion_modules.append(
|
||||||
|
get_motion_module(
|
||||||
|
in_channels=in_channels,
|
||||||
|
motion_module_type=motion_module_type,
|
||||||
|
motion_module_kwargs=motion_module_kwargs,
|
||||||
|
) if use_motion_module else None
|
||||||
|
)
|
||||||
|
resnets.append(
|
||||||
|
ResnetBlock3D(
|
||||||
|
in_channels=in_channels,
|
||||||
|
out_channels=in_channels,
|
||||||
|
temb_channels=temb_channels,
|
||||||
|
eps=resnet_eps,
|
||||||
|
groups=resnet_groups,
|
||||||
|
dropout=dropout,
|
||||||
|
time_embedding_norm=resnet_time_scale_shift,
|
||||||
|
non_linearity=resnet_act_fn,
|
||||||
|
output_scale_factor=output_scale_factor,
|
||||||
|
pre_norm=resnet_pre_norm,
|
||||||
|
|
||||||
|
use_inflated_groupnorm=use_inflated_groupnorm,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.attentions = nn.ModuleList(attentions)
|
||||||
|
self.resnets = nn.ModuleList(resnets)
|
||||||
|
self.motion_modules = nn.ModuleList(motion_modules)
|
||||||
|
|
||||||
|
def forward(self, hidden_states, temb=None, encoder_hidden_states=None, attention_mask=None):
|
||||||
|
hidden_states = self.resnets[0](hidden_states, temb)
|
||||||
|
for attn, resnet, motion_module in zip(self.attentions, self.resnets[1:], self.motion_modules):
|
||||||
|
hidden_states = attn(hidden_states, encoder_hidden_states=encoder_hidden_states).sample
|
||||||
|
hidden_states = motion_module(hidden_states, temb, encoder_hidden_states=encoder_hidden_states) if motion_module is not None else hidden_states
|
||||||
|
hidden_states = resnet(hidden_states, temb)
|
||||||
|
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
|
class CrossAttnDownBlock3D(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels: int,
|
||||||
|
out_channels: int,
|
||||||
|
temb_channels: int,
|
||||||
|
dropout: float = 0.0,
|
||||||
|
num_layers: int = 1,
|
||||||
|
resnet_eps: float = 1e-6,
|
||||||
|
resnet_time_scale_shift: str = "default",
|
||||||
|
resnet_act_fn: str = "swish",
|
||||||
|
resnet_groups: int = 32,
|
||||||
|
resnet_pre_norm: bool = True,
|
||||||
|
attn_num_head_channels=1,
|
||||||
|
cross_attention_dim=1280,
|
||||||
|
output_scale_factor=1.0,
|
||||||
|
downsample_padding=1,
|
||||||
|
add_downsample=True,
|
||||||
|
dual_cross_attention=False,
|
||||||
|
use_linear_projection=False,
|
||||||
|
only_cross_attention=False,
|
||||||
|
upcast_attention=False,
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=False,
|
||||||
|
unet_use_temporal_attention=False,
|
||||||
|
use_inflated_groupnorm=False,
|
||||||
|
|
||||||
|
use_motion_module=None,
|
||||||
|
|
||||||
|
motion_module_type=None,
|
||||||
|
motion_module_kwargs=None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
resnets = []
|
||||||
|
attentions = []
|
||||||
|
motion_modules = []
|
||||||
|
|
||||||
|
self.has_cross_attention = True
|
||||||
|
self.attn_num_head_channels = attn_num_head_channels
|
||||||
|
|
||||||
|
for i in range(num_layers):
|
||||||
|
in_channels = in_channels if i == 0 else out_channels
|
||||||
|
resnets.append(
|
||||||
|
ResnetBlock3D(
|
||||||
|
in_channels=in_channels,
|
||||||
|
out_channels=out_channels,
|
||||||
|
temb_channels=temb_channels,
|
||||||
|
eps=resnet_eps,
|
||||||
|
groups=resnet_groups,
|
||||||
|
dropout=dropout,
|
||||||
|
time_embedding_norm=resnet_time_scale_shift,
|
||||||
|
non_linearity=resnet_act_fn,
|
||||||
|
output_scale_factor=output_scale_factor,
|
||||||
|
pre_norm=resnet_pre_norm,
|
||||||
|
|
||||||
|
use_inflated_groupnorm=use_inflated_groupnorm,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if dual_cross_attention:
|
||||||
|
raise NotImplementedError
|
||||||
|
attentions.append(
|
||||||
|
Transformer3DModel(
|
||||||
|
attn_num_head_channels,
|
||||||
|
out_channels // attn_num_head_channels,
|
||||||
|
in_channels=out_channels,
|
||||||
|
num_layers=1,
|
||||||
|
cross_attention_dim=cross_attention_dim,
|
||||||
|
norm_num_groups=resnet_groups,
|
||||||
|
use_linear_projection=use_linear_projection,
|
||||||
|
only_cross_attention=only_cross_attention,
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=unet_use_cross_frame_attention,
|
||||||
|
unet_use_temporal_attention=unet_use_temporal_attention,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
motion_modules.append(
|
||||||
|
get_motion_module(
|
||||||
|
in_channels=out_channels,
|
||||||
|
motion_module_type=motion_module_type,
|
||||||
|
motion_module_kwargs=motion_module_kwargs,
|
||||||
|
) if use_motion_module else None
|
||||||
|
)
|
||||||
|
|
||||||
|
self.attentions = nn.ModuleList(attentions)
|
||||||
|
self.resnets = nn.ModuleList(resnets)
|
||||||
|
self.motion_modules = nn.ModuleList(motion_modules)
|
||||||
|
|
||||||
|
if add_downsample:
|
||||||
|
self.downsamplers = nn.ModuleList(
|
||||||
|
[
|
||||||
|
Downsample3D(
|
||||||
|
out_channels, use_conv=True, out_channels=out_channels, padding=downsample_padding, name="op"
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.downsamplers = None
|
||||||
|
|
||||||
|
self.gradient_checkpointing = False
|
||||||
|
|
||||||
|
def forward(self, hidden_states, temb=None, encoder_hidden_states=None, attention_mask=None):
|
||||||
|
output_states = ()
|
||||||
|
|
||||||
|
for resnet, attn, motion_module in zip(self.resnets, self.attentions, self.motion_modules):
|
||||||
|
if self.training and self.gradient_checkpointing:
|
||||||
|
|
||||||
|
def create_custom_forward(module, return_dict=None):
|
||||||
|
def custom_forward(*inputs):
|
||||||
|
if return_dict is not None:
|
||||||
|
return module(*inputs, return_dict=return_dict)
|
||||||
|
else:
|
||||||
|
return module(*inputs)
|
||||||
|
|
||||||
|
return custom_forward
|
||||||
|
|
||||||
|
hidden_states = checkpoint_no_reentrant(create_custom_forward(resnet), hidden_states, temb)
|
||||||
|
hidden_states = checkpoint_no_reentrant(
|
||||||
|
create_custom_forward(attn, return_dict=False),
|
||||||
|
hidden_states,
|
||||||
|
encoder_hidden_states,
|
||||||
|
)[0]
|
||||||
|
if motion_module is not None:
|
||||||
|
hidden_states = checkpoint_no_reentrant(create_custom_forward(motion_module), hidden_states.requires_grad_(), temb, encoder_hidden_states)
|
||||||
|
|
||||||
|
else:
|
||||||
|
hidden_states = resnet(hidden_states, temb)
|
||||||
|
hidden_states = attn(hidden_states, encoder_hidden_states=encoder_hidden_states).sample
|
||||||
|
|
||||||
|
# add motion module
|
||||||
|
hidden_states = motion_module(hidden_states, temb, encoder_hidden_states=encoder_hidden_states) if motion_module is not None else hidden_states
|
||||||
|
|
||||||
|
output_states += (hidden_states,)
|
||||||
|
|
||||||
|
if self.downsamplers is not None:
|
||||||
|
for downsampler in self.downsamplers:
|
||||||
|
hidden_states = downsampler(hidden_states)
|
||||||
|
|
||||||
|
output_states += (hidden_states,)
|
||||||
|
|
||||||
|
return hidden_states, output_states
|
||||||
|
|
||||||
|
|
||||||
|
class DownBlock3D(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels: int,
|
||||||
|
out_channels: int,
|
||||||
|
temb_channels: int,
|
||||||
|
dropout: float = 0.0,
|
||||||
|
num_layers: int = 1,
|
||||||
|
resnet_eps: float = 1e-6,
|
||||||
|
resnet_time_scale_shift: str = "default",
|
||||||
|
resnet_act_fn: str = "swish",
|
||||||
|
resnet_groups: int = 32,
|
||||||
|
resnet_pre_norm: bool = True,
|
||||||
|
output_scale_factor=1.0,
|
||||||
|
add_downsample=True,
|
||||||
|
downsample_padding=1,
|
||||||
|
|
||||||
|
use_inflated_groupnorm=False,
|
||||||
|
|
||||||
|
use_motion_module=None,
|
||||||
|
motion_module_type=None,
|
||||||
|
motion_module_kwargs=None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
resnets = []
|
||||||
|
motion_modules = []
|
||||||
|
|
||||||
|
for i in range(num_layers):
|
||||||
|
in_channels = in_channels if i == 0 else out_channels
|
||||||
|
resnets.append(
|
||||||
|
ResnetBlock3D(
|
||||||
|
in_channels=in_channels,
|
||||||
|
out_channels=out_channels,
|
||||||
|
temb_channels=temb_channels,
|
||||||
|
eps=resnet_eps,
|
||||||
|
groups=resnet_groups,
|
||||||
|
dropout=dropout,
|
||||||
|
time_embedding_norm=resnet_time_scale_shift,
|
||||||
|
non_linearity=resnet_act_fn,
|
||||||
|
output_scale_factor=output_scale_factor,
|
||||||
|
pre_norm=resnet_pre_norm,
|
||||||
|
|
||||||
|
use_inflated_groupnorm=use_inflated_groupnorm,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
motion_modules.append(
|
||||||
|
get_motion_module(
|
||||||
|
in_channels=out_channels,
|
||||||
|
motion_module_type=motion_module_type,
|
||||||
|
motion_module_kwargs=motion_module_kwargs,
|
||||||
|
) if use_motion_module else None
|
||||||
|
)
|
||||||
|
|
||||||
|
self.resnets = nn.ModuleList(resnets)
|
||||||
|
self.motion_modules = nn.ModuleList(motion_modules)
|
||||||
|
|
||||||
|
if add_downsample:
|
||||||
|
self.downsamplers = nn.ModuleList(
|
||||||
|
[
|
||||||
|
Downsample3D(
|
||||||
|
out_channels, use_conv=True, out_channels=out_channels, padding=downsample_padding, name="op"
|
||||||
|
)
|
||||||
|
]
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.downsamplers = None
|
||||||
|
|
||||||
|
self.gradient_checkpointing = False
|
||||||
|
|
||||||
|
def forward(self, hidden_states, temb=None, encoder_hidden_states=None):
|
||||||
|
output_states = ()
|
||||||
|
|
||||||
|
for resnet, motion_module in zip(self.resnets, self.motion_modules):
|
||||||
|
if self.training and self.gradient_checkpointing:
|
||||||
|
def create_custom_forward(module):
|
||||||
|
def custom_forward(*inputs):
|
||||||
|
return module(*inputs)
|
||||||
|
|
||||||
|
return custom_forward
|
||||||
|
|
||||||
|
hidden_states = checkpoint_no_reentrant(create_custom_forward(resnet), hidden_states, temb)
|
||||||
|
if motion_module is not None:
|
||||||
|
hidden_states = checkpoint_no_reentrant(create_custom_forward(motion_module), hidden_states.requires_grad_(), temb, encoder_hidden_states)
|
||||||
|
else:
|
||||||
|
hidden_states = resnet(hidden_states, temb)
|
||||||
|
|
||||||
|
# add motion module
|
||||||
|
hidden_states = motion_module(hidden_states, temb, encoder_hidden_states=encoder_hidden_states) if motion_module is not None else hidden_states
|
||||||
|
|
||||||
|
output_states += (hidden_states,)
|
||||||
|
|
||||||
|
if self.downsamplers is not None:
|
||||||
|
for downsampler in self.downsamplers:
|
||||||
|
hidden_states = downsampler(hidden_states)
|
||||||
|
|
||||||
|
output_states += (hidden_states,)
|
||||||
|
|
||||||
|
return hidden_states, output_states
|
||||||
|
|
||||||
|
|
||||||
|
class CrossAttnUpBlock3D(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels: int,
|
||||||
|
out_channels: int,
|
||||||
|
prev_output_channel: int,
|
||||||
|
temb_channels: int,
|
||||||
|
dropout: float = 0.0,
|
||||||
|
num_layers: int = 1,
|
||||||
|
resnet_eps: float = 1e-6,
|
||||||
|
resnet_time_scale_shift: str = "default",
|
||||||
|
resnet_act_fn: str = "swish",
|
||||||
|
resnet_groups: int = 32,
|
||||||
|
resnet_pre_norm: bool = True,
|
||||||
|
attn_num_head_channels=1,
|
||||||
|
cross_attention_dim=1280,
|
||||||
|
output_scale_factor=1.0,
|
||||||
|
add_upsample=True,
|
||||||
|
dual_cross_attention=False,
|
||||||
|
use_linear_projection=False,
|
||||||
|
only_cross_attention=False,
|
||||||
|
upcast_attention=False,
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=False,
|
||||||
|
unet_use_temporal_attention=False,
|
||||||
|
use_inflated_groupnorm=False,
|
||||||
|
|
||||||
|
use_motion_module=None,
|
||||||
|
|
||||||
|
motion_module_type=None,
|
||||||
|
motion_module_kwargs=None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
resnets = []
|
||||||
|
attentions = []
|
||||||
|
motion_modules = []
|
||||||
|
|
||||||
|
self.has_cross_attention = True
|
||||||
|
self.attn_num_head_channels = attn_num_head_channels
|
||||||
|
|
||||||
|
for i in range(num_layers):
|
||||||
|
res_skip_channels = in_channels if (i == num_layers - 1) else out_channels
|
||||||
|
resnet_in_channels = prev_output_channel if i == 0 else out_channels
|
||||||
|
|
||||||
|
resnets.append(
|
||||||
|
ResnetBlock3D(
|
||||||
|
in_channels=resnet_in_channels + res_skip_channels,
|
||||||
|
out_channels=out_channels,
|
||||||
|
temb_channels=temb_channels,
|
||||||
|
eps=resnet_eps,
|
||||||
|
groups=resnet_groups,
|
||||||
|
dropout=dropout,
|
||||||
|
time_embedding_norm=resnet_time_scale_shift,
|
||||||
|
non_linearity=resnet_act_fn,
|
||||||
|
output_scale_factor=output_scale_factor,
|
||||||
|
pre_norm=resnet_pre_norm,
|
||||||
|
|
||||||
|
use_inflated_groupnorm=use_inflated_groupnorm,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if dual_cross_attention:
|
||||||
|
raise NotImplementedError
|
||||||
|
attentions.append(
|
||||||
|
Transformer3DModel(
|
||||||
|
attn_num_head_channels,
|
||||||
|
out_channels // attn_num_head_channels,
|
||||||
|
in_channels=out_channels,
|
||||||
|
num_layers=1,
|
||||||
|
cross_attention_dim=cross_attention_dim,
|
||||||
|
norm_num_groups=resnet_groups,
|
||||||
|
use_linear_projection=use_linear_projection,
|
||||||
|
only_cross_attention=only_cross_attention,
|
||||||
|
upcast_attention=upcast_attention,
|
||||||
|
|
||||||
|
unet_use_cross_frame_attention=unet_use_cross_frame_attention,
|
||||||
|
unet_use_temporal_attention=unet_use_temporal_attention,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
motion_modules.append(
|
||||||
|
get_motion_module(
|
||||||
|
in_channels=out_channels,
|
||||||
|
motion_module_type=motion_module_type,
|
||||||
|
motion_module_kwargs=motion_module_kwargs,
|
||||||
|
) if use_motion_module else None
|
||||||
|
)
|
||||||
|
|
||||||
|
self.attentions = nn.ModuleList(attentions)
|
||||||
|
self.resnets = nn.ModuleList(resnets)
|
||||||
|
self.motion_modules = nn.ModuleList(motion_modules)
|
||||||
|
|
||||||
|
if add_upsample:
|
||||||
|
self.upsamplers = nn.ModuleList([Upsample3D(out_channels, use_conv=True, out_channels=out_channels)])
|
||||||
|
else:
|
||||||
|
self.upsamplers = None
|
||||||
|
|
||||||
|
self.gradient_checkpointing = False
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
hidden_states,
|
||||||
|
res_hidden_states_tuple,
|
||||||
|
temb=None,
|
||||||
|
encoder_hidden_states=None,
|
||||||
|
upsample_size=None,
|
||||||
|
attention_mask=None,
|
||||||
|
):
|
||||||
|
for resnet, attn, motion_module in zip(self.resnets, self.attentions, self.motion_modules):
|
||||||
|
# pop res hidden states
|
||||||
|
res_hidden_states = res_hidden_states_tuple[-1]
|
||||||
|
res_hidden_states_tuple = res_hidden_states_tuple[:-1]
|
||||||
|
hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1)
|
||||||
|
|
||||||
|
if self.training and self.gradient_checkpointing:
|
||||||
|
|
||||||
|
def create_custom_forward(module, return_dict=None):
|
||||||
|
def custom_forward(*inputs):
|
||||||
|
if return_dict is not None:
|
||||||
|
return module(*inputs, return_dict=return_dict)
|
||||||
|
else:
|
||||||
|
return module(*inputs)
|
||||||
|
|
||||||
|
return custom_forward
|
||||||
|
|
||||||
|
hidden_states = checkpoint_no_reentrant(create_custom_forward(resnet), hidden_states, temb)
|
||||||
|
hidden_states = checkpoint_no_reentrant(
|
||||||
|
create_custom_forward(attn, return_dict=False),
|
||||||
|
hidden_states,
|
||||||
|
encoder_hidden_states,
|
||||||
|
)[0]
|
||||||
|
if motion_module is not None:
|
||||||
|
hidden_states = checkpoint_no_reentrant(create_custom_forward(motion_module), hidden_states.requires_grad_(), temb, encoder_hidden_states)
|
||||||
|
|
||||||
|
else:
|
||||||
|
hidden_states = resnet(hidden_states, temb)
|
||||||
|
hidden_states = attn(hidden_states, encoder_hidden_states=encoder_hidden_states).sample
|
||||||
|
|
||||||
|
# add motion module
|
||||||
|
hidden_states = motion_module(hidden_states, temb, encoder_hidden_states=encoder_hidden_states) if motion_module is not None else hidden_states
|
||||||
|
|
||||||
|
if self.upsamplers is not None:
|
||||||
|
for upsampler in self.upsamplers:
|
||||||
|
hidden_states = upsampler(hidden_states, upsample_size)
|
||||||
|
|
||||||
|
return hidden_states
|
||||||
|
|
||||||
|
|
||||||
|
class UpBlock3D(nn.Module):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels: int,
|
||||||
|
prev_output_channel: int,
|
||||||
|
out_channels: int,
|
||||||
|
temb_channels: int,
|
||||||
|
dropout: float = 0.0,
|
||||||
|
num_layers: int = 1,
|
||||||
|
resnet_eps: float = 1e-6,
|
||||||
|
resnet_time_scale_shift: str = "default",
|
||||||
|
resnet_act_fn: str = "swish",
|
||||||
|
resnet_groups: int = 32,
|
||||||
|
resnet_pre_norm: bool = True,
|
||||||
|
output_scale_factor=1.0,
|
||||||
|
add_upsample=True,
|
||||||
|
|
||||||
|
use_inflated_groupnorm=False,
|
||||||
|
|
||||||
|
use_motion_module=None,
|
||||||
|
motion_module_type=None,
|
||||||
|
motion_module_kwargs=None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
resnets = []
|
||||||
|
motion_modules = []
|
||||||
|
|
||||||
|
for i in range(num_layers):
|
||||||
|
res_skip_channels = in_channels if (i == num_layers - 1) else out_channels
|
||||||
|
resnet_in_channels = prev_output_channel if i == 0 else out_channels
|
||||||
|
|
||||||
|
resnets.append(
|
||||||
|
ResnetBlock3D(
|
||||||
|
in_channels=resnet_in_channels + res_skip_channels,
|
||||||
|
out_channels=out_channels,
|
||||||
|
temb_channels=temb_channels,
|
||||||
|
eps=resnet_eps,
|
||||||
|
groups=resnet_groups,
|
||||||
|
dropout=dropout,
|
||||||
|
time_embedding_norm=resnet_time_scale_shift,
|
||||||
|
non_linearity=resnet_act_fn,
|
||||||
|
output_scale_factor=output_scale_factor,
|
||||||
|
pre_norm=resnet_pre_norm,
|
||||||
|
|
||||||
|
use_inflated_groupnorm=use_inflated_groupnorm,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
motion_modules.append(
|
||||||
|
get_motion_module(
|
||||||
|
in_channels=out_channels,
|
||||||
|
motion_module_type=motion_module_type,
|
||||||
|
motion_module_kwargs=motion_module_kwargs,
|
||||||
|
) if use_motion_module else None
|
||||||
|
)
|
||||||
|
|
||||||
|
self.resnets = nn.ModuleList(resnets)
|
||||||
|
self.motion_modules = nn.ModuleList(motion_modules)
|
||||||
|
|
||||||
|
if add_upsample:
|
||||||
|
self.upsamplers = nn.ModuleList([Upsample3D(out_channels, use_conv=True, out_channels=out_channels)])
|
||||||
|
else:
|
||||||
|
self.upsamplers = None
|
||||||
|
|
||||||
|
self.gradient_checkpointing = False
|
||||||
|
|
||||||
|
def forward(self, hidden_states, res_hidden_states_tuple, temb=None, upsample_size=None, encoder_hidden_states=None,):
|
||||||
|
for resnet, motion_module in zip(self.resnets, self.motion_modules):
|
||||||
|
# pop res hidden states
|
||||||
|
res_hidden_states = res_hidden_states_tuple[-1]
|
||||||
|
res_hidden_states_tuple = res_hidden_states_tuple[:-1]
|
||||||
|
hidden_states = torch.cat([hidden_states, res_hidden_states], dim=1)
|
||||||
|
|
||||||
|
if self.training and self.gradient_checkpointing:
|
||||||
|
def create_custom_forward(module):
|
||||||
|
def custom_forward(*inputs):
|
||||||
|
return module(*inputs)
|
||||||
|
|
||||||
|
return custom_forward
|
||||||
|
|
||||||
|
hidden_states = checkpoint_no_reentrant(create_custom_forward(resnet), hidden_states, temb)
|
||||||
|
if motion_module is not None:
|
||||||
|
hidden_states = checkpoint_no_reentrant(create_custom_forward(motion_module), hidden_states.requires_grad_(), temb, encoder_hidden_states)
|
||||||
|
else:
|
||||||
|
hidden_states = resnet(hidden_states, temb)
|
||||||
|
hidden_states = motion_module(hidden_states, temb, encoder_hidden_states=encoder_hidden_states) if motion_module is not None else hidden_states
|
||||||
|
|
||||||
|
if self.upsamplers is not None:
|
||||||
|
for upsampler in self.upsamplers:
|
||||||
|
hidden_states = upsampler(hidden_states, upsample_size)
|
||||||
|
|
||||||
|
return hidden_states
|
||||||
@@ -0,0 +1,465 @@
|
|||||||
|
# Adapted from https://github.com/showlab/Tune-A-Video/blob/main/tuneavideo/pipelines/pipeline_tuneavideo.py
|
||||||
|
|
||||||
|
import inspect
|
||||||
|
from typing import Callable, List, Optional, Union
|
||||||
|
from dataclasses import dataclass
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import torch
|
||||||
|
from tqdm import tqdm
|
||||||
|
|
||||||
|
from diffusers.utils import is_accelerate_available
|
||||||
|
from packaging import version
|
||||||
|
from transformers import CLIPTextModel, CLIPTokenizer
|
||||||
|
|
||||||
|
from diffusers.configuration_utils import FrozenDict
|
||||||
|
from diffusers.models import AutoencoderKL
|
||||||
|
from diffusers import DiffusionPipeline
|
||||||
|
from diffusers.schedulers import (
|
||||||
|
DDIMScheduler,
|
||||||
|
DPMSolverMultistepScheduler,
|
||||||
|
EulerAncestralDiscreteScheduler,
|
||||||
|
EulerDiscreteScheduler,
|
||||||
|
LMSDiscreteScheduler,
|
||||||
|
PNDMScheduler,
|
||||||
|
)
|
||||||
|
from diffusers.utils import deprecate, logging, BaseOutput
|
||||||
|
|
||||||
|
from einops import rearrange
|
||||||
|
|
||||||
|
from ..models.unet import UNet3DConditionModel
|
||||||
|
from ..models.sparse_controlnet import SparseControlNetModel
|
||||||
|
import pdb
|
||||||
|
|
||||||
|
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class AnimationPipelineOutput(BaseOutput):
|
||||||
|
videos: Union[torch.Tensor, np.ndarray]
|
||||||
|
|
||||||
|
|
||||||
|
class AnimationPipeline(DiffusionPipeline):
|
||||||
|
_optional_components = []
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
vae: AutoencoderKL,
|
||||||
|
text_encoder: CLIPTextModel,
|
||||||
|
tokenizer: CLIPTokenizer,
|
||||||
|
unet: UNet3DConditionModel,
|
||||||
|
scheduler: Union[
|
||||||
|
DDIMScheduler,
|
||||||
|
PNDMScheduler,
|
||||||
|
LMSDiscreteScheduler,
|
||||||
|
EulerDiscreteScheduler,
|
||||||
|
EulerAncestralDiscreteScheduler,
|
||||||
|
DPMSolverMultistepScheduler,
|
||||||
|
],
|
||||||
|
controlnet: Union[SparseControlNetModel, None] = None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
if hasattr(scheduler.config, "steps_offset") and scheduler.config.steps_offset != 1:
|
||||||
|
deprecation_message = (
|
||||||
|
f"The configuration file of this scheduler: {scheduler} is outdated. `steps_offset`"
|
||||||
|
f" should be set to 1 instead of {scheduler.config.steps_offset}. Please make sure "
|
||||||
|
"to update the config accordingly as leaving `steps_offset` might led to incorrect results"
|
||||||
|
" in future versions. If you have downloaded this checkpoint from the Hugging Face Hub,"
|
||||||
|
" it would be very nice if you could open a Pull request for the `scheduler/scheduler_config.json`"
|
||||||
|
" file"
|
||||||
|
)
|
||||||
|
deprecate("steps_offset!=1", "1.0.0", deprecation_message, standard_warn=False)
|
||||||
|
new_config = dict(scheduler.config)
|
||||||
|
new_config["steps_offset"] = 1
|
||||||
|
scheduler._internal_dict = FrozenDict(new_config)
|
||||||
|
|
||||||
|
if hasattr(scheduler.config, "clip_sample") and scheduler.config.clip_sample is True:
|
||||||
|
deprecation_message = (
|
||||||
|
f"The configuration file of this scheduler: {scheduler} has not set the configuration `clip_sample`."
|
||||||
|
" `clip_sample` should be set to False in the configuration file. Please make sure to update the"
|
||||||
|
" config accordingly as not setting `clip_sample` in the config might lead to incorrect results in"
|
||||||
|
" future versions. If you have downloaded this checkpoint from the Hugging Face Hub, it would be very"
|
||||||
|
" nice if you could open a Pull request for the `scheduler/scheduler_config.json` file"
|
||||||
|
)
|
||||||
|
deprecate("clip_sample not set", "1.0.0", deprecation_message, standard_warn=False)
|
||||||
|
new_config = dict(scheduler.config)
|
||||||
|
new_config["clip_sample"] = False
|
||||||
|
scheduler._internal_dict = FrozenDict(new_config)
|
||||||
|
|
||||||
|
is_unet_version_less_0_9_0 = hasattr(unet.config, "_diffusers_version") and version.parse(
|
||||||
|
version.parse(unet.config._diffusers_version).base_version
|
||||||
|
) < version.parse("0.9.0.dev0")
|
||||||
|
is_unet_sample_size_less_64 = hasattr(unet.config, "sample_size") and unet.config.sample_size < 64
|
||||||
|
if is_unet_version_less_0_9_0 and is_unet_sample_size_less_64:
|
||||||
|
deprecation_message = (
|
||||||
|
"The configuration file of the unet has set the default `sample_size` to smaller than"
|
||||||
|
" 64 which seems highly unlikely. If your checkpoint is a fine-tuned version of any of the"
|
||||||
|
" following: \n- CompVis/stable-diffusion-v1-4 \n- CompVis/stable-diffusion-v1-3 \n-"
|
||||||
|
" CompVis/stable-diffusion-v1-2 \n- CompVis/stable-diffusion-v1-1 \n- runwayml/stable-diffusion-v1-5"
|
||||||
|
" \n- runwayml/stable-diffusion-inpainting \n you should change 'sample_size' to 64 in the"
|
||||||
|
" configuration file. Please make sure to update the config accordingly as leaving `sample_size=32`"
|
||||||
|
" in the config might lead to incorrect results in future versions. If you have downloaded this"
|
||||||
|
" checkpoint from the Hugging Face Hub, it would be very nice if you could open a Pull request for"
|
||||||
|
" the `unet/config.json` file"
|
||||||
|
)
|
||||||
|
deprecate("sample_size<64", "1.0.0", deprecation_message, standard_warn=False)
|
||||||
|
new_config = dict(unet.config)
|
||||||
|
new_config["sample_size"] = 64
|
||||||
|
unet._internal_dict = FrozenDict(new_config)
|
||||||
|
|
||||||
|
self.register_modules(
|
||||||
|
vae=vae,
|
||||||
|
text_encoder=text_encoder,
|
||||||
|
tokenizer=tokenizer,
|
||||||
|
unet=unet,
|
||||||
|
scheduler=scheduler,
|
||||||
|
controlnet=controlnet,
|
||||||
|
)
|
||||||
|
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
|
||||||
|
|
||||||
|
def enable_vae_slicing(self):
|
||||||
|
self.vae.enable_slicing()
|
||||||
|
|
||||||
|
def disable_vae_slicing(self):
|
||||||
|
self.vae.disable_slicing()
|
||||||
|
|
||||||
|
def enable_sequential_cpu_offload(self, gpu_id=0):
|
||||||
|
if is_accelerate_available():
|
||||||
|
from accelerate import cpu_offload
|
||||||
|
else:
|
||||||
|
raise ImportError("Please install accelerate via `pip install accelerate`")
|
||||||
|
|
||||||
|
device = torch.device(f"cuda:{gpu_id}")
|
||||||
|
|
||||||
|
for cpu_offloaded_model in [self.unet, self.text_encoder, self.vae]:
|
||||||
|
if cpu_offloaded_model is not None:
|
||||||
|
cpu_offload(cpu_offloaded_model, device)
|
||||||
|
|
||||||
|
|
||||||
|
@property
|
||||||
|
def _execution_device(self):
|
||||||
|
if self.device != torch.device("meta") or not hasattr(self.unet, "_hf_hook"):
|
||||||
|
return self.device
|
||||||
|
for module in self.unet.modules():
|
||||||
|
if (
|
||||||
|
hasattr(module, "_hf_hook")
|
||||||
|
and hasattr(module._hf_hook, "execution_device")
|
||||||
|
and module._hf_hook.execution_device is not None
|
||||||
|
):
|
||||||
|
return torch.device(module._hf_hook.execution_device)
|
||||||
|
return self.device
|
||||||
|
|
||||||
|
def _encode_prompt(self, prompt, device, num_videos_per_prompt, do_classifier_free_guidance, negative_prompt):
|
||||||
|
batch_size = len(prompt) if isinstance(prompt, list) else 1
|
||||||
|
|
||||||
|
text_inputs = self.tokenizer(
|
||||||
|
prompt,
|
||||||
|
padding="max_length",
|
||||||
|
max_length=self.tokenizer.model_max_length,
|
||||||
|
truncation=True,
|
||||||
|
return_tensors="pt",
|
||||||
|
)
|
||||||
|
text_input_ids = text_inputs.input_ids
|
||||||
|
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
|
||||||
|
|
||||||
|
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
|
||||||
|
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, self.tokenizer.model_max_length - 1 : -1])
|
||||||
|
logger.warning(
|
||||||
|
"The following part of your input was truncated because CLIP can only handle sequences up to"
|
||||||
|
f" {self.tokenizer.model_max_length} tokens: {removed_text}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if hasattr(self.text_encoder.config, "use_attention_mask") and self.text_encoder.config.use_attention_mask:
|
||||||
|
attention_mask = text_inputs.attention_mask.to(device)
|
||||||
|
else:
|
||||||
|
attention_mask = None
|
||||||
|
|
||||||
|
text_embeddings = self.text_encoder(
|
||||||
|
text_input_ids.to(device),
|
||||||
|
attention_mask=attention_mask,
|
||||||
|
)
|
||||||
|
text_embeddings = text_embeddings[0]
|
||||||
|
|
||||||
|
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||||
|
bs_embed, seq_len, _ = text_embeddings.shape
|
||||||
|
text_embeddings = text_embeddings.repeat(1, num_videos_per_prompt, 1)
|
||||||
|
text_embeddings = text_embeddings.view(bs_embed * num_videos_per_prompt, seq_len, -1)
|
||||||
|
|
||||||
|
# get unconditional embeddings for classifier free guidance
|
||||||
|
if do_classifier_free_guidance:
|
||||||
|
uncond_tokens: List[str]
|
||||||
|
if negative_prompt is None:
|
||||||
|
uncond_tokens = [""] * batch_size
|
||||||
|
elif type(prompt) is not type(negative_prompt):
|
||||||
|
raise TypeError(
|
||||||
|
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
|
||||||
|
f" {type(prompt)}."
|
||||||
|
)
|
||||||
|
elif isinstance(negative_prompt, str):
|
||||||
|
uncond_tokens = [negative_prompt]
|
||||||
|
elif batch_size != len(negative_prompt):
|
||||||
|
raise ValueError(
|
||||||
|
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
|
||||||
|
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
|
||||||
|
" the batch size of `prompt`."
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
uncond_tokens = negative_prompt
|
||||||
|
|
||||||
|
max_length = text_input_ids.shape[-1]
|
||||||
|
uncond_input = self.tokenizer(
|
||||||
|
uncond_tokens,
|
||||||
|
padding="max_length",
|
||||||
|
max_length=max_length,
|
||||||
|
truncation=True,
|
||||||
|
return_tensors="pt",
|
||||||
|
)
|
||||||
|
|
||||||
|
if hasattr(self.text_encoder.config, "use_attention_mask") and self.text_encoder.config.use_attention_mask:
|
||||||
|
attention_mask = uncond_input.attention_mask.to(device)
|
||||||
|
else:
|
||||||
|
attention_mask = None
|
||||||
|
|
||||||
|
uncond_embeddings = self.text_encoder(
|
||||||
|
uncond_input.input_ids.to(device),
|
||||||
|
attention_mask=attention_mask,
|
||||||
|
)
|
||||||
|
uncond_embeddings = uncond_embeddings[0]
|
||||||
|
|
||||||
|
# duplicate unconditional embeddings for each generation per prompt, using mps friendly method
|
||||||
|
seq_len = uncond_embeddings.shape[1]
|
||||||
|
uncond_embeddings = uncond_embeddings.repeat(1, num_videos_per_prompt, 1)
|
||||||
|
uncond_embeddings = uncond_embeddings.view(batch_size * num_videos_per_prompt, seq_len, -1)
|
||||||
|
|
||||||
|
# For classifier free guidance, we need to do two forward passes.
|
||||||
|
# Here we concatenate the unconditional and text embeddings into a single batch
|
||||||
|
# to avoid doing two forward passes
|
||||||
|
text_embeddings = torch.cat([uncond_embeddings, text_embeddings])
|
||||||
|
|
||||||
|
return text_embeddings
|
||||||
|
|
||||||
|
def decode_latents(self, latents):
|
||||||
|
video_length = latents.shape[2]
|
||||||
|
latents = 1 / 0.18215 * latents
|
||||||
|
latents = rearrange(latents, "b c f h w -> (b f) c h w")
|
||||||
|
# video = self.vae.decode(latents).sample
|
||||||
|
video = []
|
||||||
|
for frame_idx in tqdm(range(latents.shape[0])):
|
||||||
|
video.append(self.vae.decode(latents[frame_idx:frame_idx+1]).sample)
|
||||||
|
video = torch.cat(video)
|
||||||
|
video = rearrange(video, "(b f) c h w -> b c f h w", f=video_length)
|
||||||
|
video = (video / 2 + 0.5).clamp(0, 1)
|
||||||
|
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloa16
|
||||||
|
video = video.cpu().float().numpy()
|
||||||
|
return video
|
||||||
|
|
||||||
|
def prepare_extra_step_kwargs(self, generator, eta):
|
||||||
|
# prepare extra kwargs for the scheduler step, since not all schedulers have the same signature
|
||||||
|
# eta (η) is only used with the DDIMScheduler, it will be ignored for other schedulers.
|
||||||
|
# eta corresponds to η in DDIM paper: https://arxiv.org/abs/2010.02502
|
||||||
|
# and should be between [0, 1]
|
||||||
|
|
||||||
|
accepts_eta = "eta" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||||
|
extra_step_kwargs = {}
|
||||||
|
if accepts_eta:
|
||||||
|
extra_step_kwargs["eta"] = eta
|
||||||
|
|
||||||
|
# check if the scheduler accepts generator
|
||||||
|
accepts_generator = "generator" in set(inspect.signature(self.scheduler.step).parameters.keys())
|
||||||
|
if accepts_generator:
|
||||||
|
extra_step_kwargs["generator"] = generator
|
||||||
|
return extra_step_kwargs
|
||||||
|
|
||||||
|
def check_inputs(self, prompt, height, width, callback_steps):
|
||||||
|
if not isinstance(prompt, str) and not isinstance(prompt, list):
|
||||||
|
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
|
||||||
|
|
||||||
|
if height % 8 != 0 or width % 8 != 0:
|
||||||
|
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
|
||||||
|
|
||||||
|
if (callback_steps is None) or (
|
||||||
|
callback_steps is not None and (not isinstance(callback_steps, int) or callback_steps <= 0)
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
f"`callback_steps` has to be a positive integer but is {callback_steps} of type"
|
||||||
|
f" {type(callback_steps)}."
|
||||||
|
)
|
||||||
|
|
||||||
|
def prepare_latents(self, batch_size, num_channels_latents, video_length, height, width, dtype, device, generator, latents=None):
|
||||||
|
shape = (batch_size, num_channels_latents, video_length, height // self.vae_scale_factor, width // self.vae_scale_factor)
|
||||||
|
if isinstance(generator, list) and len(generator) != batch_size:
|
||||||
|
raise ValueError(
|
||||||
|
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||||
|
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||||
|
)
|
||||||
|
if latents is None:
|
||||||
|
rand_device = "cpu" if device.type == "mps" else device
|
||||||
|
|
||||||
|
if isinstance(generator, list):
|
||||||
|
shape = shape
|
||||||
|
# shape = (1,) + shape[1:]
|
||||||
|
latents = [
|
||||||
|
torch.randn(shape, generator=generator[i], device=rand_device, dtype=dtype)
|
||||||
|
for i in range(batch_size)
|
||||||
|
]
|
||||||
|
latents = torch.cat(latents, dim=0).to(device)
|
||||||
|
else:
|
||||||
|
latents = torch.randn(shape, generator=generator, device=rand_device, dtype=dtype).to(device)
|
||||||
|
else:
|
||||||
|
if latents.shape != shape:
|
||||||
|
raise ValueError(f"Unexpected latents shape, got {latents.shape}, expected {shape}")
|
||||||
|
latents = latents.to(device)
|
||||||
|
|
||||||
|
# scale the initial noise by the standard deviation required by the scheduler
|
||||||
|
latents = latents * self.scheduler.init_noise_sigma
|
||||||
|
return latents
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def __call__(
|
||||||
|
self,
|
||||||
|
prompt: Union[str, List[str]],
|
||||||
|
video_length: Optional[int],
|
||||||
|
height: Optional[int] = None,
|
||||||
|
width: Optional[int] = None,
|
||||||
|
num_inference_steps: int = 50,
|
||||||
|
guidance_scale: float = 7.5,
|
||||||
|
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||||
|
num_videos_per_prompt: Optional[int] = 1,
|
||||||
|
eta: float = 0.0,
|
||||||
|
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||||
|
latents: Optional[torch.FloatTensor] = None,
|
||||||
|
output_type: Optional[str] = "tensor",
|
||||||
|
return_dict: bool = True,
|
||||||
|
callback: Optional[Callable[[int, int, torch.FloatTensor], None]] = None,
|
||||||
|
callback_steps: Optional[int] = 1,
|
||||||
|
|
||||||
|
# support controlnet
|
||||||
|
controlnet_images: torch.FloatTensor = None,
|
||||||
|
controlnet_image_index: list = [0],
|
||||||
|
controlnet_conditioning_scale: Union[float, List[float]] = 1.0,
|
||||||
|
|
||||||
|
**kwargs,
|
||||||
|
):
|
||||||
|
# Default height and width to unet
|
||||||
|
height = height or self.unet.config.sample_size * self.vae_scale_factor
|
||||||
|
width = width or self.unet.config.sample_size * self.vae_scale_factor
|
||||||
|
|
||||||
|
# Check inputs. Raise error if not correct
|
||||||
|
self.check_inputs(prompt, height, width, callback_steps)
|
||||||
|
|
||||||
|
# Define call parameters
|
||||||
|
# batch_size = 1 if isinstance(prompt, str) else len(prompt)
|
||||||
|
batch_size = 1
|
||||||
|
if latents is not None:
|
||||||
|
batch_size = latents.shape[0]
|
||||||
|
if isinstance(prompt, list):
|
||||||
|
batch_size = len(prompt)
|
||||||
|
|
||||||
|
device = self._execution_device
|
||||||
|
# here `guidance_scale` is defined analog to the guidance weight `w` of equation (2)
|
||||||
|
# of the Imagen paper: https://arxiv.org/pdf/2205.11487.pdf . `guidance_scale = 1`
|
||||||
|
# corresponds to doing no classifier free guidance.
|
||||||
|
do_classifier_free_guidance = guidance_scale > 1.0
|
||||||
|
|
||||||
|
# Encode input prompt
|
||||||
|
prompt = prompt if isinstance(prompt, list) else [prompt] * batch_size
|
||||||
|
if negative_prompt is not None:
|
||||||
|
negative_prompt = negative_prompt if isinstance(negative_prompt, list) else [negative_prompt] * batch_size
|
||||||
|
text_embeddings = self._encode_prompt(
|
||||||
|
prompt, device, num_videos_per_prompt, do_classifier_free_guidance, negative_prompt
|
||||||
|
)
|
||||||
|
|
||||||
|
# Prepare timesteps
|
||||||
|
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||||
|
timesteps = self.scheduler.timesteps
|
||||||
|
|
||||||
|
# Prepare latent variables
|
||||||
|
num_channels_latents = self.unet.in_channels
|
||||||
|
latents = self.prepare_latents(
|
||||||
|
batch_size * num_videos_per_prompt,
|
||||||
|
num_channels_latents,
|
||||||
|
video_length,
|
||||||
|
height,
|
||||||
|
width,
|
||||||
|
text_embeddings.dtype,
|
||||||
|
device,
|
||||||
|
generator,
|
||||||
|
latents,
|
||||||
|
)
|
||||||
|
latents_dtype = latents.dtype
|
||||||
|
|
||||||
|
# Prepare extra step kwargs.
|
||||||
|
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
|
||||||
|
|
||||||
|
# Denoising loop
|
||||||
|
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||||
|
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||||
|
for i, t in enumerate(timesteps):
|
||||||
|
# expand the latents if we are doing classifier free guidance
|
||||||
|
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||||
|
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||||
|
|
||||||
|
down_block_additional_residuals = mid_block_additional_residual = None
|
||||||
|
if (getattr(self, "controlnet", None) != None) and (controlnet_images != None):
|
||||||
|
assert controlnet_images.dim() == 5
|
||||||
|
|
||||||
|
controlnet_noisy_latents = latent_model_input
|
||||||
|
controlnet_prompt_embeds = text_embeddings
|
||||||
|
|
||||||
|
controlnet_images = controlnet_images.to(latents.device)
|
||||||
|
|
||||||
|
controlnet_cond_shape = list(controlnet_images.shape)
|
||||||
|
controlnet_cond_shape[2] = video_length
|
||||||
|
controlnet_cond = torch.zeros(controlnet_cond_shape).to(latents.device)
|
||||||
|
|
||||||
|
controlnet_conditioning_mask_shape = list(controlnet_cond.shape)
|
||||||
|
controlnet_conditioning_mask_shape[1] = 1
|
||||||
|
controlnet_conditioning_mask = torch.zeros(controlnet_conditioning_mask_shape).to(latents.device)
|
||||||
|
|
||||||
|
assert controlnet_images.shape[2] >= len(controlnet_image_index)
|
||||||
|
controlnet_cond[:,:,controlnet_image_index] = controlnet_images[:,:,:len(controlnet_image_index)]
|
||||||
|
controlnet_conditioning_mask[:,:,controlnet_image_index] = 1
|
||||||
|
|
||||||
|
down_block_additional_residuals, mid_block_additional_residual = self.controlnet(
|
||||||
|
controlnet_noisy_latents, t,
|
||||||
|
encoder_hidden_states=controlnet_prompt_embeds,
|
||||||
|
controlnet_cond=controlnet_cond,
|
||||||
|
conditioning_mask=controlnet_conditioning_mask,
|
||||||
|
conditioning_scale=controlnet_conditioning_scale,
|
||||||
|
guess_mode=False, return_dict=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
# predict the noise residual
|
||||||
|
noise_pred = self.unet(
|
||||||
|
latent_model_input, t,
|
||||||
|
encoder_hidden_states=text_embeddings,
|
||||||
|
down_block_additional_residuals = down_block_additional_residuals,
|
||||||
|
mid_block_additional_residual = mid_block_additional_residual,
|
||||||
|
).sample.to(dtype=latents_dtype)
|
||||||
|
|
||||||
|
# perform guidance
|
||||||
|
if do_classifier_free_guidance:
|
||||||
|
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||||
|
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||||
|
|
||||||
|
# compute the previous noisy sample x_t -> x_t-1
|
||||||
|
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs).prev_sample
|
||||||
|
|
||||||
|
# call the callback, if provided
|
||||||
|
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||||
|
progress_bar.update()
|
||||||
|
if callback is not None and i % callback_steps == 0:
|
||||||
|
callback(i, t, latents)
|
||||||
|
|
||||||
|
# Post-processing
|
||||||
|
video = self.decode_latents(latents)
|
||||||
|
|
||||||
|
# Convert to tensor
|
||||||
|
if output_type == "tensor":
|
||||||
|
video = torch.from_numpy(video)
|
||||||
|
|
||||||
|
if not return_dict:
|
||||||
|
return video
|
||||||
|
|
||||||
|
return AnimationPipelineOutput(videos=video)
|
||||||
@@ -0,0 +1,240 @@
|
|||||||
|
# Script for converting a HF Diffusers saved pipeline to a Stable Diffusion checkpoint.
|
||||||
|
# *Only* converts the UNet, VAE, and Text Encoder.
|
||||||
|
# Does not convert optimizer state or any other thing.
|
||||||
|
# Written by jachiam
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import os.path as osp
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
|
||||||
|
# =================#
|
||||||
|
# UNet Conversion #
|
||||||
|
# =================#
|
||||||
|
|
||||||
|
unet_conversion_map = [
|
||||||
|
# (stable-diffusion, HF Diffusers)
|
||||||
|
("time_embed.0.weight", "time_embedding.linear_1.weight"),
|
||||||
|
("time_embed.0.bias", "time_embedding.linear_1.bias"),
|
||||||
|
("time_embed.2.weight", "time_embedding.linear_2.weight"),
|
||||||
|
("time_embed.2.bias", "time_embedding.linear_2.bias"),
|
||||||
|
("input_blocks.0.0.weight", "conv_in.weight"),
|
||||||
|
("input_blocks.0.0.bias", "conv_in.bias"),
|
||||||
|
("out.0.weight", "conv_norm_out.weight"),
|
||||||
|
("out.0.bias", "conv_norm_out.bias"),
|
||||||
|
("out.2.weight", "conv_out.weight"),
|
||||||
|
("out.2.bias", "conv_out.bias"),
|
||||||
|
]
|
||||||
|
|
||||||
|
unet_conversion_map_resnet = [
|
||||||
|
# (stable-diffusion, HF Diffusers)
|
||||||
|
("in_layers.0", "norm1"),
|
||||||
|
("in_layers.2", "conv1"),
|
||||||
|
("out_layers.0", "norm2"),
|
||||||
|
("out_layers.3", "conv2"),
|
||||||
|
("emb_layers.1", "time_emb_proj"),
|
||||||
|
("skip_connection", "conv_shortcut"),
|
||||||
|
]
|
||||||
|
|
||||||
|
unet_conversion_map_layer = []
|
||||||
|
# hardcoded number of downblocks and resnets/attentions...
|
||||||
|
# would need smarter logic for other networks.
|
||||||
|
for i in range(4):
|
||||||
|
# loop over downblocks/upblocks
|
||||||
|
|
||||||
|
for j in range(2):
|
||||||
|
# loop over resnets/attentions for downblocks
|
||||||
|
hf_down_res_prefix = f"down_blocks.{i}.resnets.{j}."
|
||||||
|
sd_down_res_prefix = f"input_blocks.{3*i + j + 1}.0."
|
||||||
|
unet_conversion_map_layer.append((sd_down_res_prefix, hf_down_res_prefix))
|
||||||
|
|
||||||
|
if i < 3:
|
||||||
|
# no attention layers in down_blocks.3
|
||||||
|
hf_down_atn_prefix = f"down_blocks.{i}.attentions.{j}."
|
||||||
|
sd_down_atn_prefix = f"input_blocks.{3*i + j + 1}.1."
|
||||||
|
unet_conversion_map_layer.append((sd_down_atn_prefix, hf_down_atn_prefix))
|
||||||
|
|
||||||
|
for j in range(3):
|
||||||
|
# loop over resnets/attentions for upblocks
|
||||||
|
hf_up_res_prefix = f"up_blocks.{i}.resnets.{j}."
|
||||||
|
sd_up_res_prefix = f"output_blocks.{3*i + j}.0."
|
||||||
|
unet_conversion_map_layer.append((sd_up_res_prefix, hf_up_res_prefix))
|
||||||
|
|
||||||
|
if i > 0:
|
||||||
|
# no attention layers in up_blocks.0
|
||||||
|
hf_up_atn_prefix = f"up_blocks.{i}.attentions.{j}."
|
||||||
|
sd_up_atn_prefix = f"output_blocks.{3*i + j}.1."
|
||||||
|
unet_conversion_map_layer.append((sd_up_atn_prefix, hf_up_atn_prefix))
|
||||||
|
|
||||||
|
if i < 3:
|
||||||
|
# no downsample in down_blocks.3
|
||||||
|
hf_downsample_prefix = f"down_blocks.{i}.downsamplers.0.conv."
|
||||||
|
sd_downsample_prefix = f"input_blocks.{3*(i+1)}.0.op."
|
||||||
|
unet_conversion_map_layer.append((sd_downsample_prefix, hf_downsample_prefix))
|
||||||
|
|
||||||
|
# no upsample in up_blocks.3
|
||||||
|
hf_upsample_prefix = f"up_blocks.{i}.upsamplers.0."
|
||||||
|
sd_upsample_prefix = f"output_blocks.{3*i + 2}.{1 if i == 0 else 2}."
|
||||||
|
unet_conversion_map_layer.append((sd_upsample_prefix, hf_upsample_prefix))
|
||||||
|
|
||||||
|
hf_mid_atn_prefix = "mid_block.attentions.0."
|
||||||
|
sd_mid_atn_prefix = "middle_block.1."
|
||||||
|
unet_conversion_map_layer.append((sd_mid_atn_prefix, hf_mid_atn_prefix))
|
||||||
|
|
||||||
|
for j in range(2):
|
||||||
|
hf_mid_res_prefix = f"mid_block.resnets.{j}."
|
||||||
|
sd_mid_res_prefix = f"middle_block.{2*j}."
|
||||||
|
unet_conversion_map_layer.append((sd_mid_res_prefix, hf_mid_res_prefix))
|
||||||
|
|
||||||
|
|
||||||
|
def convert_unet_state_dict(unet_state_dict):
|
||||||
|
# buyer beware: this is a *brittle* function,
|
||||||
|
# and correct output requires that all of these pieces interact in
|
||||||
|
# the exact order in which I have arranged them.
|
||||||
|
mapping = {k: k for k in unet_state_dict.keys()}
|
||||||
|
for sd_name, hf_name in unet_conversion_map:
|
||||||
|
mapping[hf_name] = sd_name
|
||||||
|
for k, v in mapping.items():
|
||||||
|
if "resnets" in k:
|
||||||
|
for sd_part, hf_part in unet_conversion_map_resnet:
|
||||||
|
v = v.replace(hf_part, sd_part)
|
||||||
|
mapping[k] = v
|
||||||
|
for k, v in mapping.items():
|
||||||
|
for sd_part, hf_part in unet_conversion_map_layer:
|
||||||
|
v = v.replace(hf_part, sd_part)
|
||||||
|
mapping[k] = v
|
||||||
|
new_state_dict = {v: unet_state_dict[k] for k, v in mapping.items() if k in unet_state_dict}
|
||||||
|
return prepend_unet_key(new_state_dict)
|
||||||
|
|
||||||
|
def prepend_unet_key(unet_state_dict):
|
||||||
|
return {"model.diffusion_model." + k: v for k, v in unet_state_dict.items()}
|
||||||
|
|
||||||
|
def prepend_text_encoder_key(text_encoder_state_dict):
|
||||||
|
return {"cond_stage_model.transformer." + k: v for k, v in text_encoder_state_dict.items()}
|
||||||
|
|
||||||
|
# ================#
|
||||||
|
# VAE Conversion #
|
||||||
|
# ================#
|
||||||
|
|
||||||
|
vae_conversion_map = [
|
||||||
|
# (stable-diffusion, HF Diffusers)
|
||||||
|
("nin_shortcut", "conv_shortcut"),
|
||||||
|
("norm_out", "conv_norm_out"),
|
||||||
|
("mid.attn_1.", "mid_block.attentions.0."),
|
||||||
|
]
|
||||||
|
|
||||||
|
for i in range(4):
|
||||||
|
# down_blocks have two resnets
|
||||||
|
for j in range(2):
|
||||||
|
hf_down_prefix = f"encoder.down_blocks.{i}.resnets.{j}."
|
||||||
|
sd_down_prefix = f"encoder.down.{i}.block.{j}."
|
||||||
|
vae_conversion_map.append((sd_down_prefix, hf_down_prefix))
|
||||||
|
|
||||||
|
if i < 3:
|
||||||
|
hf_downsample_prefix = f"down_blocks.{i}.downsamplers.0."
|
||||||
|
sd_downsample_prefix = f"down.{i}.downsample."
|
||||||
|
vae_conversion_map.append((sd_downsample_prefix, hf_downsample_prefix))
|
||||||
|
|
||||||
|
hf_upsample_prefix = f"up_blocks.{i}.upsamplers.0."
|
||||||
|
sd_upsample_prefix = f"up.{3-i}.upsample."
|
||||||
|
vae_conversion_map.append((sd_upsample_prefix, hf_upsample_prefix))
|
||||||
|
|
||||||
|
# up_blocks have three resnets
|
||||||
|
# also, up blocks in hf are numbered in reverse from sd
|
||||||
|
for j in range(3):
|
||||||
|
hf_up_prefix = f"decoder.up_blocks.{i}.resnets.{j}."
|
||||||
|
sd_up_prefix = f"decoder.up.{3-i}.block.{j}."
|
||||||
|
vae_conversion_map.append((sd_up_prefix, hf_up_prefix))
|
||||||
|
|
||||||
|
# this part accounts for mid blocks in both the encoder and the decoder
|
||||||
|
for i in range(2):
|
||||||
|
hf_mid_res_prefix = f"mid_block.resnets.{i}."
|
||||||
|
sd_mid_res_prefix = f"mid.block_{i+1}."
|
||||||
|
vae_conversion_map.append((sd_mid_res_prefix, hf_mid_res_prefix))
|
||||||
|
|
||||||
|
|
||||||
|
vae_conversion_map_attn = [
|
||||||
|
# (stable-diffusion, HF Diffusers)
|
||||||
|
("norm.", "group_norm."),
|
||||||
|
("q.", "query."),
|
||||||
|
("k.", "key."),
|
||||||
|
("v.", "value."),
|
||||||
|
("proj_out.", "proj_attn."),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def reshape_weight_for_sd(w):
|
||||||
|
# convert HF linear weights to SD conv2d weights
|
||||||
|
return w.reshape(*w.shape, 1, 1)
|
||||||
|
|
||||||
|
|
||||||
|
def convert_vae_state_dict(vae_state_dict):
|
||||||
|
mapping = {k: k for k in vae_state_dict.keys()}
|
||||||
|
for k, v in mapping.items():
|
||||||
|
for sd_part, hf_part in vae_conversion_map:
|
||||||
|
v = v.replace(hf_part, sd_part)
|
||||||
|
mapping[k] = v
|
||||||
|
for k, v in mapping.items():
|
||||||
|
if "attentions" in k:
|
||||||
|
for sd_part, hf_part in vae_conversion_map_attn:
|
||||||
|
v = v.replace(hf_part, sd_part)
|
||||||
|
mapping[k] = v
|
||||||
|
new_state_dict = {v: vae_state_dict[k] for k, v in mapping.items()}
|
||||||
|
weights_to_convert = ["q", "k", "v", "proj_out"]
|
||||||
|
for k, v in new_state_dict.items():
|
||||||
|
for weight_name in weights_to_convert:
|
||||||
|
if f"mid.attn_1.{weight_name}.weight" in k:
|
||||||
|
print(f"Reshaping {k} for SD format")
|
||||||
|
new_state_dict[k] = reshape_weight_for_sd(v)
|
||||||
|
return new_state_dict
|
||||||
|
|
||||||
|
|
||||||
|
# =========================#
|
||||||
|
# Text Encoder Conversion #
|
||||||
|
# =========================#
|
||||||
|
# pretty much a no-op
|
||||||
|
|
||||||
|
|
||||||
|
def convert_text_enc_state_dict(text_enc_dict):
|
||||||
|
return prepend_text_encoder_key(text_enc_dict)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
|
||||||
|
parser.add_argument("--model_path", default=None, type=str, required=True, help="Path to the model to convert.")
|
||||||
|
parser.add_argument("--checkpoint_path", default=None, type=str, required=True, help="Path to the output model.")
|
||||||
|
parser.add_argument("--half", action="store_true", help="Save weights in half precision.")
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
assert args.model_path is not None, "Must provide a model path!"
|
||||||
|
|
||||||
|
assert args.checkpoint_path is not None, "Must provide a checkpoint path!"
|
||||||
|
|
||||||
|
unet_path = osp.join(args.model_path, "unet", "diffusion_pytorch_model.bin")
|
||||||
|
vae_path = osp.join(args.model_path, "vae", "diffusion_pytorch_model.bin")
|
||||||
|
text_enc_path = osp.join(args.model_path, "text_encoder", "pytorch_model.bin")
|
||||||
|
|
||||||
|
# Convert the UNet model
|
||||||
|
unet_state_dict = torch.load(unet_path, map_location='cpu')
|
||||||
|
unet_state_dict = convert_unet_state_dict(unet_state_dict)
|
||||||
|
unet_state_dict = {"model.diffusion_model." + k: v for k, v in unet_state_dict.items()}
|
||||||
|
|
||||||
|
# Convert the VAE model
|
||||||
|
vae_state_dict = torch.load(vae_path, map_location='cpu')
|
||||||
|
vae_state_dict = convert_vae_state_dict(vae_state_dict)
|
||||||
|
vae_state_dict = {"first_stage_model." + k: v for k, v in vae_state_dict.items()}
|
||||||
|
|
||||||
|
# Convert the text encoder model
|
||||||
|
text_enc_dict = torch.load(text_enc_path, map_location='cpu')
|
||||||
|
text_enc_dict = convert_text_enc_state_dict(text_enc_dict)
|
||||||
|
text_enc_dict = {"cond_stage_model.transformer." + k: v for k, v in text_enc_dict.items()}
|
||||||
|
|
||||||
|
# Put together new checkpoint
|
||||||
|
state_dict = {**unet_state_dict, **vae_state_dict, **text_enc_dict}
|
||||||
|
if args.half:
|
||||||
|
state_dict = {k:v.half() for k,v in state_dict.items()}
|
||||||
|
state_dict = {"state_dict": state_dict}
|
||||||
|
torch.save(state_dict, args.checkpoint_path)
|
||||||
@@ -0,0 +1,383 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
import os
|
||||||
|
import loralib as loralb
|
||||||
|
from loralib import LoRALayer
|
||||||
|
import math
|
||||||
|
import json
|
||||||
|
|
||||||
|
from torch.utils.data import ConcatDataset
|
||||||
|
from transformers import CLIPTokenizer
|
||||||
|
|
||||||
|
try:
|
||||||
|
from safetensors.torch import save_file, load_file
|
||||||
|
except:
|
||||||
|
print("Safetensors is not installed. Saving while using use_safetensors will fail.")
|
||||||
|
|
||||||
|
UNET_REPLACE = ["Transformer2DModel", "ResnetBlock2D"]
|
||||||
|
TEXT_ENCODER_REPLACE = ["CLIPAttention", "CLIPTextEmbeddings"]
|
||||||
|
|
||||||
|
UNET_ATTENTION_REPLACE = ["CrossAttention"]
|
||||||
|
TEXT_ENCODER_ATTENTION_REPLACE = ["CLIPAttention", "CLIPTextEmbeddings"]
|
||||||
|
|
||||||
|
"""
|
||||||
|
Copied from: https://github.com/cloneofsimo/lora/blob/bdd51b04c49fa90a88919a19850ec3b4cf3c5ecd/lora_diffusion/lora.py#L189
|
||||||
|
"""
|
||||||
|
def find_modules(
|
||||||
|
model,
|
||||||
|
ancestor_class= None,
|
||||||
|
search_class = [torch.nn.Linear],
|
||||||
|
exclude_children_of = [loralb.Linear, loralb.Conv2d, loralb.Embedding],
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Find all modules of a certain class (or union of classes) that are direct or
|
||||||
|
indirect descendants of other modules of a certain class (or union of classes).
|
||||||
|
|
||||||
|
Returns all matching modules, along with the parent of those moduless and the
|
||||||
|
names they are referenced by.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Get the targets we should replace all linears under
|
||||||
|
if ancestor_class is not None:
|
||||||
|
ancestors = (
|
||||||
|
module
|
||||||
|
for module in model.modules()
|
||||||
|
if module.__class__.__name__ in ancestor_class
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# this, incase you want to naively iterate over all modules.
|
||||||
|
ancestors = [module for module in model.modules()]
|
||||||
|
|
||||||
|
# For each target find every linear_class module that isn't a child of a LoraInjectedLinear
|
||||||
|
for ancestor in ancestors:
|
||||||
|
for fullname, module in ancestor.named_modules():
|
||||||
|
if any([isinstance(module, _class) for _class in search_class]):
|
||||||
|
# Find the direct parent if this is a descendant, not a child, of target
|
||||||
|
*path, name = fullname.split(".")
|
||||||
|
parent = ancestor
|
||||||
|
while path:
|
||||||
|
parent = parent.get_submodule(path.pop(0))
|
||||||
|
# Skip this linear if it's a child of a LoraInjectedLinear
|
||||||
|
if exclude_children_of and any(
|
||||||
|
[isinstance(parent, _class) for _class in exclude_children_of]
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
# Otherwise, yield it
|
||||||
|
yield parent, name, module
|
||||||
|
|
||||||
|
class Conv2d(nn.Conv2d, LoRALayer):
|
||||||
|
# LoRA implemented in a dense layer
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels: int,
|
||||||
|
out_channels: int,
|
||||||
|
kernel_size: int,
|
||||||
|
r: int = 0,
|
||||||
|
lora_alpha: int = 1,
|
||||||
|
lora_dropout: float = 0.,
|
||||||
|
merge_weights: bool = True,
|
||||||
|
**kwargs
|
||||||
|
):
|
||||||
|
nn.Conv2d.__init__(self, in_channels, out_channels, kernel_size, **kwargs)
|
||||||
|
LoRALayer.__init__(self, r=r, lora_alpha=lora_alpha, lora_dropout=lora_dropout,
|
||||||
|
merge_weights=merge_weights)
|
||||||
|
assert type(kernel_size) is int
|
||||||
|
# Actual trainable parameters
|
||||||
|
if r > 0:
|
||||||
|
self.lora_A = nn.Parameter(
|
||||||
|
self.weight.new_zeros((r*kernel_size, in_channels*kernel_size))
|
||||||
|
)
|
||||||
|
self.lora_B = nn.Parameter(
|
||||||
|
self.weight.new_zeros((out_channels*kernel_size, r*kernel_size))
|
||||||
|
)
|
||||||
|
self.scaling = self.lora_alpha / self.r
|
||||||
|
# Freezing the pre-trained weight matrix
|
||||||
|
self.weight.requires_grad = False
|
||||||
|
self.reset_parameters()
|
||||||
|
|
||||||
|
def reset_parameters(self):
|
||||||
|
nn.Conv2d.reset_parameters(self)
|
||||||
|
if hasattr(self, 'lora_A'):
|
||||||
|
# initialize A the same way as the default for nn.Linear and B to zero
|
||||||
|
nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))
|
||||||
|
nn.init.zeros_(self.lora_B)
|
||||||
|
|
||||||
|
def train(self, mode: bool = True):
|
||||||
|
nn.Conv2d.train(self, mode)
|
||||||
|
if mode:
|
||||||
|
if self.merge_weights and self.merged:
|
||||||
|
# Make sure that the weights are not merged
|
||||||
|
self.weight.data -= (self.lora_B @ self.lora_A).view(self.weight.shape) * self.scaling
|
||||||
|
self.merged = False
|
||||||
|
else:
|
||||||
|
if self.merge_weights and not self.merged:
|
||||||
|
# Merge the weights and mark it
|
||||||
|
self.weight.data += (self.lora_B @ self.lora_A).view(self.weight.shape) * self.scaling
|
||||||
|
self.merged = True
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor):
|
||||||
|
if self.r > 0 and not self.merged:
|
||||||
|
return F.conv2d(
|
||||||
|
x,
|
||||||
|
self.weight + (self.lora_B @ self.lora_A).view(self.weight.shape) * self.scaling,
|
||||||
|
self.bias, self.stride, self.padding, self.dilation, self.groups
|
||||||
|
)
|
||||||
|
return nn.Conv2d.forward(self, x)
|
||||||
|
|
||||||
|
class Conv3d(nn.Conv3d, LoRALayer):
|
||||||
|
# LoRA implemented in a dense layer
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
in_channels: int,
|
||||||
|
out_channels: int,
|
||||||
|
kernel_size: int,
|
||||||
|
r: int = 0,
|
||||||
|
lora_alpha: int = 1,
|
||||||
|
lora_dropout: float = 0.,
|
||||||
|
merge_weights: bool = True,
|
||||||
|
**kwargs
|
||||||
|
):
|
||||||
|
nn.Conv3d.__init__(self, in_channels, out_channels, (kernel_size, 1, 1), **kwargs)
|
||||||
|
LoRALayer.__init__(self, r=r, lora_alpha=lora_alpha, lora_dropout=lora_dropout,
|
||||||
|
merge_weights=merge_weights)
|
||||||
|
assert type(kernel_size) is int
|
||||||
|
# Actual trainable parameters
|
||||||
|
|
||||||
|
# Get view transform shape
|
||||||
|
i, o, k = self.weight.shape[:3]
|
||||||
|
self.view_shape = (i, o, k, kernel_size, 1)
|
||||||
|
|
||||||
|
if r > 0:
|
||||||
|
self.lora_A = nn.Parameter(
|
||||||
|
self.weight.new_zeros((r*kernel_size, in_channels*kernel_size))
|
||||||
|
)
|
||||||
|
self.lora_B = nn.Parameter(
|
||||||
|
self.weight.new_zeros((out_channels*kernel_size, r*kernel_size))
|
||||||
|
)
|
||||||
|
self.scaling = self.lora_alpha / self.r
|
||||||
|
# Freezing the pre-trained weight matrix
|
||||||
|
self.weight.requires_grad = False
|
||||||
|
self.reset_parameters()
|
||||||
|
|
||||||
|
def reset_parameters(self):
|
||||||
|
nn.Conv3d.reset_parameters(self)
|
||||||
|
if hasattr(self, 'lora_A'):
|
||||||
|
# initialize A the same way as the default for nn.Linear and B to zero
|
||||||
|
nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))
|
||||||
|
nn.init.zeros_(self.lora_B)
|
||||||
|
|
||||||
|
def train(self, mode: bool = True):
|
||||||
|
nn.Conv3d.train(self, mode)
|
||||||
|
if mode:
|
||||||
|
if self.merge_weights and self.merged:
|
||||||
|
# Make sure that the weights are not merged
|
||||||
|
self.weight.data -= (self.lora_B @ self.lora_A).view(self.weight.shape) * self.scaling
|
||||||
|
self.merged = False
|
||||||
|
else:
|
||||||
|
if self.merge_weights and not self.merged:
|
||||||
|
# Merge the weights and mark it
|
||||||
|
self.weight.data += (self.lora_B @ self.lora_A).view(self.weight.shape) * self.scaling
|
||||||
|
self.merged = True
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor):
|
||||||
|
if self.r > 0 and not self.merged:
|
||||||
|
return F.conv3d(
|
||||||
|
x,
|
||||||
|
self.weight + torch.mean((self.lora_B @ self.lora_A).view(self.view_shape), dim=-2, keepdim=True) * \
|
||||||
|
self.scaling, self.bias, self.stride, self.padding, self.dilation, self.groups
|
||||||
|
)
|
||||||
|
return nn.Conv3d.forward(self, x)
|
||||||
|
|
||||||
|
def create_lora_linear(child_module, r, dropout=0, bias=False, scale=1):
|
||||||
|
return loralb.Linear(
|
||||||
|
child_module.in_features,
|
||||||
|
child_module.out_features,
|
||||||
|
merge_weights=False,
|
||||||
|
bias=bias,
|
||||||
|
lora_dropout=dropout,
|
||||||
|
lora_alpha=scale,
|
||||||
|
r=r
|
||||||
|
)
|
||||||
|
return lora_linear
|
||||||
|
|
||||||
|
def create_lora_conv(child_module, r, dropout=0, bias=False, rescale=False, scale=1):
|
||||||
|
return Conv2d(
|
||||||
|
child_module.in_channels,
|
||||||
|
child_module.out_channels,
|
||||||
|
kernel_size=child_module.kernel_size[0],
|
||||||
|
padding=child_module.padding,
|
||||||
|
stride=child_module.stride,
|
||||||
|
merge_weights=False,
|
||||||
|
bias=bias,
|
||||||
|
lora_dropout=dropout,
|
||||||
|
lora_alpha=scale,
|
||||||
|
r=r,
|
||||||
|
)
|
||||||
|
return lora_conv
|
||||||
|
|
||||||
|
def create_lora_conv3d(child_module, r, dropout=0, bias=False, rescale=False, scale=1):
|
||||||
|
return Conv3d(
|
||||||
|
child_module.in_channels,
|
||||||
|
child_module.out_channels,
|
||||||
|
kernel_size=child_module.kernel_size[0],
|
||||||
|
padding=child_module.padding,
|
||||||
|
stride=child_module.stride,
|
||||||
|
merge_weights=False,
|
||||||
|
bias=bias,
|
||||||
|
lora_dropout=dropout,
|
||||||
|
lora_alpha=scale,
|
||||||
|
r=r,
|
||||||
|
)
|
||||||
|
return lora_conv
|
||||||
|
|
||||||
|
def create_lora_emb(child_module, r, scale=1):
|
||||||
|
return loralb.Embedding(
|
||||||
|
child_module.num_embeddings,
|
||||||
|
child_module.embedding_dim,
|
||||||
|
merge_weights=False,
|
||||||
|
lora_alpha=scale,
|
||||||
|
r=r
|
||||||
|
)
|
||||||
|
|
||||||
|
def activate_lora_train(model, bias):
|
||||||
|
def unfreeze():
|
||||||
|
print(model.__class__.__name__ + " LoRA set for training.")
|
||||||
|
return loralb.mark_only_lora_as_trainable(model, bias=bias)
|
||||||
|
|
||||||
|
return unfreeze
|
||||||
|
|
||||||
|
def add_lora_to(
|
||||||
|
model,
|
||||||
|
target_module=UNET_REPLACE,
|
||||||
|
search_class=[torch.nn.Linear],
|
||||||
|
r=32,
|
||||||
|
dropout=0,
|
||||||
|
scale=0,
|
||||||
|
lora_bias='none',
|
||||||
|
):
|
||||||
|
scale = scale if (scale > 0 and isinstance(scale, int)) else r
|
||||||
|
for module, name, child_module in find_modules(
|
||||||
|
model,
|
||||||
|
ancestor_class=target_module,
|
||||||
|
search_class=search_class,
|
||||||
|
exclude_children_of=[loralb.Linear, loralb.Embedding, Conv2d, Conv3d]
|
||||||
|
):
|
||||||
|
|
||||||
|
bias = hasattr(child_module, "bias")
|
||||||
|
|
||||||
|
# Check if child module of the model has bias.
|
||||||
|
if bias:
|
||||||
|
if child_module.bias is None:
|
||||||
|
bias = False
|
||||||
|
|
||||||
|
# Check if the child module of the model is type Linear or Conv2d.
|
||||||
|
if isinstance(child_module, torch.nn.Linear):
|
||||||
|
l = create_lora_linear(child_module, r, dropout, bias=bias, scale=scale)
|
||||||
|
|
||||||
|
if isinstance(child_module, torch.nn.Conv2d):
|
||||||
|
l = create_lora_conv(child_module, r, dropout, bias=bias, scale=scale)
|
||||||
|
|
||||||
|
if isinstance(child_module, torch.nn.Conv3d):
|
||||||
|
l = create_lora_conv3d(child_module, r, dropout, bias=bias, scale=scale)
|
||||||
|
|
||||||
|
if isinstance(child_module, torch.nn.Embedding):
|
||||||
|
l = create_lora_emb(child_module, r, scale=scale)
|
||||||
|
|
||||||
|
# If the model has bias and we wish to add it, use the child_modules in place
|
||||||
|
if bias:
|
||||||
|
l.bias = child_module.bias
|
||||||
|
|
||||||
|
# Assign the frozen weight of model's Linear or Conv2d to the LoRA model.
|
||||||
|
l.weight = child_module.weight
|
||||||
|
|
||||||
|
# Replace the new LoRA model with the model's Linear or Conv2d module.
|
||||||
|
module._modules[name] = l
|
||||||
|
|
||||||
|
|
||||||
|
# Unfreeze only the newly added LoRA weights, but keep the model frozen.
|
||||||
|
return activate_lora_train(model, lora_bias)
|
||||||
|
|
||||||
|
def save_lora(
|
||||||
|
unet=None,
|
||||||
|
text_encoder=None,
|
||||||
|
save_text_weights=False,
|
||||||
|
output_dir="output",
|
||||||
|
lora_filename="lora.safetensors",
|
||||||
|
lora_bias='none',
|
||||||
|
save_for_webui=True,
|
||||||
|
only_webui=False,
|
||||||
|
metadata=None,
|
||||||
|
unet_dict_converter=None,
|
||||||
|
text_dict_converter=None
|
||||||
|
):
|
||||||
|
|
||||||
|
if not only_webui:
|
||||||
|
# Create directory for the full LoRA weights.
|
||||||
|
trainable_weights_dir = f"{output_dir}/full_weights"
|
||||||
|
lora_out_file_full_weight = f"{trainable_weights_dir}/{lora_filename}"
|
||||||
|
os.makedirs(trainable_weights_dir, exist_ok=True)
|
||||||
|
|
||||||
|
ext = '.safetensors'
|
||||||
|
# Create LoRA out filename.
|
||||||
|
lora_out_file = f"{output_dir}/{lora_filename}{ext}"
|
||||||
|
|
||||||
|
if not only_webui:
|
||||||
|
save_path_full_weights = lora_out_file_full_weight + ext
|
||||||
|
|
||||||
|
save_path = lora_out_file
|
||||||
|
|
||||||
|
if not only_webui:
|
||||||
|
for i, model in enumerate([unet, text_encoder]):
|
||||||
|
if save_text_weights and i == 1:
|
||||||
|
non_webui_weights = save_path_full_weights.replace(ext, f"_text_encoder{ext}")
|
||||||
|
|
||||||
|
else:
|
||||||
|
non_webui_weights = save_path_full_weights.replace(ext, f"_unet{ext}")
|
||||||
|
|
||||||
|
# Load only the LoRAs from the state dict.
|
||||||
|
lora_dict = loralb.lora_state_dict(model, bias=lora_bias)
|
||||||
|
|
||||||
|
# Save the models as fp32. This ensures we can finetune again without having to upcast.
|
||||||
|
save_file(lora_dict, non_webui_weights)
|
||||||
|
|
||||||
|
if save_for_webui:
|
||||||
|
# Convert the keys to compvis model and webui
|
||||||
|
unet_lora_dict = loralb.lora_state_dict(unet, bias=lora_bias)
|
||||||
|
lora_dict_fp16 = unet_dict_converter(unet_lora_dict, strict_mapping=True)
|
||||||
|
|
||||||
|
if save_text_weights:
|
||||||
|
text_encoder_dict = loralb.lora_state_dict(text_encoder, bias=lora_bias)
|
||||||
|
lora_dict_text_fp16 = text_dict_converter(text_encoder_dict)
|
||||||
|
|
||||||
|
# Update the Unet dict to include text keys.
|
||||||
|
lora_dict_fp16.update(lora_dict_text_fp16)
|
||||||
|
|
||||||
|
# Cast tensors to fp16. It's assumed we won't be finetuning these.
|
||||||
|
for k, v in lora_dict_fp16.items():
|
||||||
|
lora_dict_fp16[k] = v.to(dtype=torch.float16)
|
||||||
|
|
||||||
|
save_file(
|
||||||
|
lora_dict_fp16,
|
||||||
|
save_path,
|
||||||
|
metadata=metadata
|
||||||
|
)
|
||||||
|
|
||||||
|
def load_lora(model, lora_path: str, *args, **kwargs):
|
||||||
|
try:
|
||||||
|
if os.path.exists(lora_path):
|
||||||
|
lora_dict = load_file(lora_path)
|
||||||
|
model.load_state_dict(lora_dict, strict=False)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Could not load your lora file: {e}")
|
||||||
|
|
||||||
|
def set_mode(model, train=False):
|
||||||
|
for n, m in model.named_modules():
|
||||||
|
is_lora = hasattr(m, 'merged')
|
||||||
|
if is_lora:
|
||||||
|
m.train(train)
|
||||||
|
|
||||||
|
def set_mode_group(models, train):
|
||||||
|
for model in models:
|
||||||
|
set_mode(model, train)
|
||||||
@@ -0,0 +1,34 @@
|
|||||||
|
def load_lora(model, lora_path: str):
|
||||||
|
try:
|
||||||
|
if os.path.exists(lora_path):
|
||||||
|
lora_dict = load_file(lora_path)
|
||||||
|
POSSIBLE_KEYS = ['text_model', 'model']
|
||||||
|
|
||||||
|
reorder_dict = False
|
||||||
|
|
||||||
|
for key in POSSIBLE_KEYS:
|
||||||
|
temp_key = list(lora_dict.keys())[-2]
|
||||||
|
key_check = [kr for kr in POSSIBLE_KEYS if kr in temp_key]
|
||||||
|
|
||||||
|
reorder_dict = len(key_check) > 0
|
||||||
|
|
||||||
|
if reorder_dict:
|
||||||
|
from collections import OrderedDict
|
||||||
|
fixed_lora_dict = OrderedDict()
|
||||||
|
|
||||||
|
for k, v in list(lora_dict.items()):
|
||||||
|
first_path = k.split('.')[0]
|
||||||
|
key_replace = [kr for kr in POSSIBLE_KEYS if kr == first_path]
|
||||||
|
|
||||||
|
if len(key_replace) > 0:
|
||||||
|
new_key = k.replace(f"{key_replace[0]}.", "")
|
||||||
|
fixed_lora_dict[new_key] = v
|
||||||
|
print(new_key)
|
||||||
|
|
||||||
|
lora_dict = fixed_lora_dict
|
||||||
|
|
||||||
|
model.load_state_dict(lora_dict)
|
||||||
|
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Could not load your lora file: {e}")
|
||||||
@@ -0,0 +1,187 @@
|
|||||||
|
import torch
|
||||||
|
from torch.utils.data import Dataset
|
||||||
|
from torchvision.transforms import transforms
|
||||||
|
from pathlib import Path
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
class DreamBoothDataset(Dataset):
|
||||||
|
"""
|
||||||
|
A dataset to prepare the instance and class images with the prompts for fine-tuning the model.
|
||||||
|
It pre-processes the images and the tokenizes prompts.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
instance_data_root,
|
||||||
|
instance_prompt,
|
||||||
|
tokenizer,
|
||||||
|
class_data_root=None,
|
||||||
|
class_prompt=None,
|
||||||
|
size=512,
|
||||||
|
center_crop=False,
|
||||||
|
color_jitter=False,
|
||||||
|
h_flip=False,
|
||||||
|
resize=False,
|
||||||
|
dataset_norm=False
|
||||||
|
):
|
||||||
|
self.size = size
|
||||||
|
self.center_crop = center_crop
|
||||||
|
self.color_jitter = color_jitter
|
||||||
|
self.h_flip = h_flip
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
self.resize = resize
|
||||||
|
self.dataset_norm = dataset_norm
|
||||||
|
|
||||||
|
self.instance_data_root = Path(instance_data_root)
|
||||||
|
if not self.instance_data_root.exists():
|
||||||
|
raise ValueError("Instance images root doesn't exists.")
|
||||||
|
|
||||||
|
self.instance_images_path = list(Path(instance_data_root).iterdir())
|
||||||
|
self.num_instance_images = len(self.instance_images_path)
|
||||||
|
self.instance_prompt = instance_prompt
|
||||||
|
self._length = self.num_instance_images
|
||||||
|
|
||||||
|
if class_data_root is not None:
|
||||||
|
self.class_data_root = Path(class_data_root)
|
||||||
|
self.class_data_root.mkdir(parents=True, exist_ok=True)
|
||||||
|
self.class_images_path = list(self.class_data_root.iterdir())
|
||||||
|
self.num_class_images = len(self.class_images_path)
|
||||||
|
self._length = max(self.num_class_images, self.num_instance_images)
|
||||||
|
self.class_prompt = class_prompt
|
||||||
|
else:
|
||||||
|
self.class_data_root = None
|
||||||
|
|
||||||
|
self.image_transforms = self.compose()
|
||||||
|
self.normalized_mean_std = self.get_dataset_norm(class_data_root)
|
||||||
|
|
||||||
|
|
||||||
|
def gather_norm(self, img, mean=None, std=None):
|
||||||
|
channels = img.shape[0]
|
||||||
|
|
||||||
|
if all(x is None for x in [mean, std]):
|
||||||
|
mean, std = torch.zeros(channels), torch.zeros(channels)
|
||||||
|
|
||||||
|
for i in range(channels):
|
||||||
|
mean[i] += img[i, :, :].mean()
|
||||||
|
std[i] += img[i, :, :].std()
|
||||||
|
|
||||||
|
return mean, std
|
||||||
|
|
||||||
|
def get_dataset_norm(self, class_data_root):
|
||||||
|
if self.dataset_norm:
|
||||||
|
imgs_to_process = self.instance_images_path
|
||||||
|
|
||||||
|
if class_data_root is not None:
|
||||||
|
imgs_to_process += self.class_images_path
|
||||||
|
|
||||||
|
mean = None
|
||||||
|
std = None
|
||||||
|
|
||||||
|
for img in tqdm(imgs_to_process, desc="Processing image normalization..."):
|
||||||
|
img = Image.open(img).convert("RGB")
|
||||||
|
img = self.image_transforms(img)
|
||||||
|
|
||||||
|
mean, std = self.gather_norm(img, mean, std)
|
||||||
|
|
||||||
|
mean.div_(len(imgs_to_process))
|
||||||
|
std.div_(len(imgs_to_process))
|
||||||
|
|
||||||
|
print(f"Dataset mean and std are: {mean}, {std}")
|
||||||
|
|
||||||
|
return mean, std
|
||||||
|
else:
|
||||||
|
return [0.5], [0.5]
|
||||||
|
|
||||||
|
def compose(self):
|
||||||
|
img_transforms = []
|
||||||
|
|
||||||
|
if self.resize:
|
||||||
|
img_transforms.append(
|
||||||
|
transforms.Resize(
|
||||||
|
(self.size, self.size), interpolation=transforms.InterpolationMode.BILINEAR
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if self.center_crop:
|
||||||
|
img_transforms.append(transforms.CenterCrop(size))
|
||||||
|
if self.color_jitter:
|
||||||
|
img_transforms.append(transforms.ColorJitter(0.2, 0.1))
|
||||||
|
if self.h_flip:
|
||||||
|
img_transforms.append(transforms.RandomHorizontalFlip())
|
||||||
|
|
||||||
|
return transforms.Compose([*img_transforms, transforms.ToTensor()])
|
||||||
|
|
||||||
|
def image_transform(self, img):
|
||||||
|
img_composed = self.image_transforms(img)
|
||||||
|
|
||||||
|
if not self.dataset_norm:
|
||||||
|
mean, std = self.gather_norm(img_composed)
|
||||||
|
mean, std = [0.5], [0.5]
|
||||||
|
else:
|
||||||
|
mean, std = self.normalized_mean_std
|
||||||
|
|
||||||
|
return transforms.Normalize(mean, std)(img_composed)
|
||||||
|
|
||||||
|
def open_img(self, index, folder):
|
||||||
|
img_path = folder[index % self.num_instance_images]
|
||||||
|
img = Image.open(img_path)
|
||||||
|
|
||||||
|
if not img.mode == "RGB":
|
||||||
|
img = img.convert("RGB")
|
||||||
|
|
||||||
|
return img, str(img_path).split("/")[-1]
|
||||||
|
|
||||||
|
def tokenize_prompt(self, prompt):
|
||||||
|
return self.tokenizer(
|
||||||
|
prompt,
|
||||||
|
padding="do_not_pad",
|
||||||
|
truncation=True,
|
||||||
|
max_length=self.tokenizer.model_max_length,
|
||||||
|
).input_ids
|
||||||
|
|
||||||
|
|
||||||
|
def get_train_sample(self, index, example, base_name, folder, prompt):
|
||||||
|
image, img_name = self.open_img(index, self.instance_images_path)
|
||||||
|
example[f"{base_name}_images"] = self.image_transform(image)
|
||||||
|
example[f"{base_name}_prompt_ids"] = self.tokenize_prompt(prompt)
|
||||||
|
example[f"{base_name}_prompt"] = prompt
|
||||||
|
example[f"{base_name}_img_name"] = img_name
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return self._length
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
example = {}
|
||||||
|
|
||||||
|
self.get_train_sample(
|
||||||
|
index,
|
||||||
|
example,
|
||||||
|
"instance",
|
||||||
|
self.instance_images_path,
|
||||||
|
self.instance_prompt
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.class_data_root:
|
||||||
|
self.get_train_sample(
|
||||||
|
index,
|
||||||
|
example,
|
||||||
|
"class",
|
||||||
|
self.class_images_path,
|
||||||
|
self.class_prompt
|
||||||
|
)
|
||||||
|
return example
|
||||||
|
|
||||||
|
class PromptDataset(Dataset):
|
||||||
|
"A simple dataset to prepare the prompts to generate class images on multiple GPUs."
|
||||||
|
|
||||||
|
def __init__(self, prompt, num_samples):
|
||||||
|
self.prompt = prompt
|
||||||
|
self.num_samples = num_samples
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return self.num_samples
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
example = {}
|
||||||
|
example["prompt"] = self.prompt
|
||||||
|
example["index"] = index
|
||||||
|
return example
|
||||||
@@ -0,0 +1,372 @@
|
|||||||
|
import argparse
|
||||||
|
import os
|
||||||
|
|
||||||
|
def parse_args(input_args=None):
|
||||||
|
parser = argparse.ArgumentParser(description="Simple example of a training script.")
|
||||||
|
parser.add_argument(
|
||||||
|
"--pretrained_model_name_or_path",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
required=True,
|
||||||
|
help="Path to pretrained model or model identifier from huggingface.co/models.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--pretrained_vae_name_or_path",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="Path to pretrained vae or vae identifier from huggingface.co/models.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--revision",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
required=False,
|
||||||
|
help="Revision of pretrained model identifier from huggingface.co/models.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--tokenizer_name",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="Pretrained tokenizer name or path if not the same as model_name",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--instance_data_dir",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
required=True,
|
||||||
|
help="A folder containing the training data of instance images.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--class_data_dir",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
required=False,
|
||||||
|
help="A folder containing the training data of class images.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--json_path",
|
||||||
|
type=str,
|
||||||
|
default="",
|
||||||
|
required=True,
|
||||||
|
help="A JSON file with the same args as argparse (instance_data_dir, class_data_dir, etc.)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--instance_prompt",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
required=True,
|
||||||
|
help="The prompt with identifier specifying the instance",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--preview_prompt",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="The prompt to use when generating preview images",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--class_prompt",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="The prompt to specify images in the same class as provided instance images.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--with_prior_preservation",
|
||||||
|
default=False,
|
||||||
|
action="store_true",
|
||||||
|
help="Flag to add prior preservation loss.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--prior_loss_weight",
|
||||||
|
type=float,
|
||||||
|
default=1.0,
|
||||||
|
help="The weight of prior preservation loss.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--prior_preservation_mode",
|
||||||
|
type=str,
|
||||||
|
choices=["additive", "multiply", "single_pass", "text"],
|
||||||
|
default="multiply",
|
||||||
|
help=("The prior preservation loss mode."
|
||||||
|
"Additive: loss + (prior_loss * loss_weight)"
|
||||||
|
"Multiply: loss + loss_weight * prior_loss",
|
||||||
|
"Text: The class prompt is used as the initializer"
|
||||||
|
"Single Pass:" "Compute the losses together"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--num_class_images",
|
||||||
|
type=int,
|
||||||
|
default=800,
|
||||||
|
help=(
|
||||||
|
"Minimal class images for prior preservation loss. If not have enough images, additional images will be"
|
||||||
|
" sampled with class_prompt."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--output_dir",
|
||||||
|
type=str,
|
||||||
|
default="text-inversion-model",
|
||||||
|
help="The output directory where the model predictions and checkpoints will be written.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--seed", type=int, default=None, help="A seed for reproducible training."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--resolution",
|
||||||
|
type=int,
|
||||||
|
default=512,
|
||||||
|
help=(
|
||||||
|
"The resolution for input images, all the images in the train/validation dataset will be resized to this"
|
||||||
|
" resolution"
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--center_crop",
|
||||||
|
action="store_true",
|
||||||
|
help="Whether to center crop images before resizing to resolution",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--color_jitter",
|
||||||
|
action="store_true",
|
||||||
|
help="Whether to apply color jitter to images",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--train_text_encoder",
|
||||||
|
action="store_true",
|
||||||
|
help="Whether to train the text encoder",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--clip_layers",
|
||||||
|
type=int,
|
||||||
|
default=12,
|
||||||
|
help="Amount of hidden layers to include in CLIP (Also known as CLIP Skip / Penultimate)",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--train_batch_size",
|
||||||
|
type=int,
|
||||||
|
default=1,
|
||||||
|
help="Batch size (per device) for the training dataloader.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--sample_batch_size",
|
||||||
|
type=int,
|
||||||
|
default=1,
|
||||||
|
help="Batch size (per device) for sampling images.",
|
||||||
|
)
|
||||||
|
parser.add_argument("--num_train_epochs", type=int, default=1)
|
||||||
|
parser.add_argument(
|
||||||
|
"--max_train_steps",
|
||||||
|
type=int,
|
||||||
|
default=None,
|
||||||
|
help="Total number of training steps to perform. If provided, overrides num_train_epochs.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--save_steps",
|
||||||
|
type=int,
|
||||||
|
default=500,
|
||||||
|
help="Save checkpoint every X updates steps.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--save_for_webui",
|
||||||
|
action="store_true",
|
||||||
|
default=True,
|
||||||
|
help="Save a LoRA model for usage in the AUTOMATIC1111 webui.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--preview_steps",
|
||||||
|
type=int,
|
||||||
|
default=100,
|
||||||
|
help="Save preview every X updates steps.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--gradient_accumulation_steps",
|
||||||
|
type=int,
|
||||||
|
default=1,
|
||||||
|
help="Number of updates steps to accumulate before performing a backward/update pass.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--gradient_checkpointing",
|
||||||
|
action="store_true",
|
||||||
|
help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--lora_rank",
|
||||||
|
type=int,
|
||||||
|
default=4,
|
||||||
|
help="Rank of LoRA approximation.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--lora_bias",
|
||||||
|
type=str,
|
||||||
|
default="none",
|
||||||
|
help="Whether or not to use bias when training LoRA.",
|
||||||
|
choices=["none", "lora_only", "all"]
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--learning_rate",
|
||||||
|
type=float,
|
||||||
|
default=None,
|
||||||
|
help="Initial learning rate (after the potential warmup period) to use.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--learning_rate_text",
|
||||||
|
type=float,
|
||||||
|
default=5e-6,
|
||||||
|
help="Initial learning rate for text encoder (after the potential warmup period) to use.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--dropout",
|
||||||
|
type=float,
|
||||||
|
default=0,
|
||||||
|
help="Dropout for both UNET and Text Encoder"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--scale_lr",
|
||||||
|
action="store_true",
|
||||||
|
default=False,
|
||||||
|
help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--dataset_norm",
|
||||||
|
action="store_true",
|
||||||
|
default=False,
|
||||||
|
help="Normalizes the entire dataset by calculating the mean and standard deviation of all elements.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--save_preview",
|
||||||
|
action="store_true",
|
||||||
|
default=False,
|
||||||
|
help="Save preview images during training.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--lr_scheduler",
|
||||||
|
type=str,
|
||||||
|
default="constant",
|
||||||
|
help=(
|
||||||
|
'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
|
||||||
|
' "constant", "constant_with_warmup"]'
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--lr_warmup_steps",
|
||||||
|
type=int,
|
||||||
|
default=500,
|
||||||
|
help="Number of steps for the warmup in the lr scheduler.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--use_8bit_adam",
|
||||||
|
action="store_true",
|
||||||
|
help="Whether or not to use 8-bit Adam from bitsandbytes.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--adam_beta1",
|
||||||
|
type=float,
|
||||||
|
default=0.9,
|
||||||
|
help="The beta1 parameter for the Adam optimizer.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--adam_beta2",
|
||||||
|
type=float,
|
||||||
|
default=0.999,
|
||||||
|
help="The beta2 parameter for the Adam optimizer.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--adam_weight_decay", type=float, default=1e-2, help="Weight decay to use."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--adam_epsilon",
|
||||||
|
type=float,
|
||||||
|
default=1e-08,
|
||||||
|
help="Epsilon value for the Adam optimizer",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--max_grad_norm", default=1.0, type=float, help="Max gradient norm."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--push_to_hub",
|
||||||
|
action="store_true",
|
||||||
|
help="Whether or not to push the model to the Hub.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--hub_token",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="The token to use to push to the Model Hub.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--logging_dir",
|
||||||
|
type=str,
|
||||||
|
default="logs",
|
||||||
|
help=(
|
||||||
|
"[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||||
|
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--mixed_precision",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
choices=["no", "fp16", "bf16"],
|
||||||
|
help=(
|
||||||
|
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
|
||||||
|
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
|
||||||
|
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--local_rank",
|
||||||
|
type=int,
|
||||||
|
default=-1,
|
||||||
|
help="For distributed training: local_rank. Not to be confused with LoRA rank.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--resume_unet",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help=("File path for unet lora to resume training."),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--resume_text_encoder",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help=("File path for text encoder lora to resume training."),
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--resize",
|
||||||
|
type=bool,
|
||||||
|
default=True,
|
||||||
|
required=False,
|
||||||
|
help="Should images be resized to --resolution before training?",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--only_attn", action="store_true", help="Only finetune attention layers."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--only_webui", action="store_true", help="Only save weights for webui."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--use_xformers", action="store_true", help="Whether or not to use xformers"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--lora_name",
|
||||||
|
type=str,
|
||||||
|
default="stable_lora",
|
||||||
|
help="The name of your project. Will get saved as LoRA metadata."
|
||||||
|
)
|
||||||
|
|
||||||
|
if input_args is not None:
|
||||||
|
args = parser.parse_args(input_args)
|
||||||
|
else:
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
env_local_rank = int(os.environ.get("LOCAL_RANK", -1))
|
||||||
|
if env_local_rank != -1 and env_local_rank != args.local_rank:
|
||||||
|
args.local_rank = env_local_rank
|
||||||
|
|
||||||
|
if args.with_prior_preservation:
|
||||||
|
if args.class_data_dir is None:
|
||||||
|
raise ValueError("You must specify a data directory for class images.")
|
||||||
|
if args.class_prompt is None:
|
||||||
|
raise ValueError("You must specify prompt for class images.")
|
||||||
|
|
||||||
|
return args
|
||||||
@@ -0,0 +1,132 @@
|
|||||||
|
import sys
|
||||||
|
|
||||||
|
from omegaconf import OmegaConf
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
QUALITY_TYPES = ["low", "preferred", "best"]
|
||||||
|
|
||||||
|
def create_quality_config(
|
||||||
|
width: int,
|
||||||
|
height: int,
|
||||||
|
sample_width: int = 384,
|
||||||
|
sample_height: int = 384,
|
||||||
|
use_bucketing: bool = False,
|
||||||
|
lora_rank: int = 2
|
||||||
|
):
|
||||||
|
config = dict(
|
||||||
|
width=width,
|
||||||
|
height=height,
|
||||||
|
use_bucketing=use_bucketing,
|
||||||
|
sample_size=(
|
||||||
|
sample_height if sample_height!= 0 else 256,
|
||||||
|
sample_width if sample_width != 0 else 256
|
||||||
|
),
|
||||||
|
lora_rank=lora_rank
|
||||||
|
)
|
||||||
|
|
||||||
|
return SimpleNamespace(**config)
|
||||||
|
|
||||||
|
def set_train_data(config: SimpleNamespace, quality_config: SimpleNamespace):
|
||||||
|
train_data_map = ["sample_size", "width", "height", "sample_size", "use_bucketing"]
|
||||||
|
|
||||||
|
# Set LoRA Rank fallback
|
||||||
|
setattr(config, 'lora_rank', getattr(quality_config, 'lora_rank', 8))
|
||||||
|
|
||||||
|
for train_setting in train_data_map:
|
||||||
|
if getattr(config.train_data, 'manual_sample_size', False) and \
|
||||||
|
train_setting == "sample_size":
|
||||||
|
continue
|
||||||
|
|
||||||
|
setattr(config.train_data, train_setting, getattr(quality_config, train_setting))
|
||||||
|
|
||||||
|
def set_single_video_args(config: SimpleNamespace, simple_config: SimpleNamespace):
|
||||||
|
config.dataset_types = ["single_video"]
|
||||||
|
|
||||||
|
single_data_map = [
|
||||||
|
("max_chunks", "max_chunks"),
|
||||||
|
("single_video_path", "path"),
|
||||||
|
("sample_start_idx", "start_time"),
|
||||||
|
("single_video_prompt", "training_prompt"),
|
||||||
|
]
|
||||||
|
|
||||||
|
for single_data_key, simple_config_key in single_data_map:
|
||||||
|
if simple_config_key == "max_chunks":
|
||||||
|
setattr(
|
||||||
|
config.train_data,
|
||||||
|
single_data_key,
|
||||||
|
getattr(simple_config.video, simple_config_key, sys.maxsize)
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
setattr(config.train_data, single_data_key, getattr(simple_config.video, simple_config_key))
|
||||||
|
|
||||||
|
def set_folder_of_videos_args(config: SimpleNamespace, simple_config: SimpleNamespace):
|
||||||
|
config.dataset_types = ["folder"]
|
||||||
|
|
||||||
|
folder_data_map = [
|
||||||
|
("max_chunks", "max_chunks"),
|
||||||
|
("path", "path"),
|
||||||
|
("single_video_prompt", "training_prompt"),
|
||||||
|
("fallback_prompt", "training_prompt"),
|
||||||
|
("prompts", "validation_prompt")
|
||||||
|
]
|
||||||
|
|
||||||
|
for folder_data_key, simple_config_key in folder_data_map:
|
||||||
|
if simple_config_key == "max_chunks":
|
||||||
|
setattr(
|
||||||
|
config.train_data,
|
||||||
|
folder_data_key,
|
||||||
|
getattr(simple_config.video, simple_config_key, sys.maxsize)
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
setattr(config.train_data, folder_data_key, getattr(simple_config.video, simple_config_key))
|
||||||
|
|
||||||
|
def build_quality_configs():
|
||||||
|
LowQualityConfig = create_quality_config(256, 256, 512, 512, lora_rank=32)
|
||||||
|
PreferredConfig = create_quality_config(384, 384, 384, 384, use_bucketing=True, lora_rank=64)
|
||||||
|
BestQualityConfig = create_quality_config(512, 512, 512, 512, use_bucketing=True, lora_rank=64)
|
||||||
|
|
||||||
|
quality_configs = {"low": LowQualityConfig, "preferred": PreferredConfig, "best": BestQualityConfig}
|
||||||
|
|
||||||
|
return quality_configs
|
||||||
|
|
||||||
|
def get_simple_config(config: OmegaConf):
|
||||||
|
simple_config = None
|
||||||
|
quality_configs = build_quality_configs()
|
||||||
|
|
||||||
|
try:
|
||||||
|
checkpoints_map = [
|
||||||
|
"pretrained_model_path",
|
||||||
|
"motion_module_path",
|
||||||
|
"unet_checkpoint_path",
|
||||||
|
"domain_adapter_path"
|
||||||
|
]
|
||||||
|
|
||||||
|
simple_config = config
|
||||||
|
config = OmegaConf.load(config.training_config)
|
||||||
|
config.lora_name = simple_config.save_name
|
||||||
|
|
||||||
|
for checkpoint_key in checkpoints_map:
|
||||||
|
setattr(config, checkpoint_key, getattr(simple_config, checkpoint_key))
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
raise ValueError("Could not load training config", e)
|
||||||
|
|
||||||
|
if simple_config.quality.lower() not in QUALITY_TYPES:
|
||||||
|
raise ValueError(f"Quality must be the following: {QUALITY_TYPES}")
|
||||||
|
|
||||||
|
quality_config = quality_configs.get(simple_config.quality.lower())
|
||||||
|
set_train_data(config, quality_config)
|
||||||
|
|
||||||
|
if simple_config.mode_type == "single_video":
|
||||||
|
set_single_video_args(config, simple_config)
|
||||||
|
elif simple_config.mode_type == "folder":
|
||||||
|
set_folder_of_videos_args(config, simple_config)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"{simple_config.mode_type} not imlemented. Choose 'single_video' or 'folder'")
|
||||||
|
|
||||||
|
config.validation_data.prompts[0] = simple_config.video.validation_prompt
|
||||||
|
|
||||||
|
return config
|
||||||
|
|
||||||
@@ -0,0 +1,529 @@
|
|||||||
|
# Script for converting a HF Diffusers saved pipeline to a Stable Diffusion checkpoint.
|
||||||
|
# *Only* converts the UNet, and Text Encoder.
|
||||||
|
# Does not convert optimizer state or any other thing.
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import os.path as osp
|
||||||
|
import re
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from safetensors.torch import load_file, save_file
|
||||||
|
|
||||||
|
# =================#
|
||||||
|
# UNet Conversion #
|
||||||
|
# =================#
|
||||||
|
|
||||||
|
unet_conversion_map = [
|
||||||
|
# (ModelScope, HF Diffusers)
|
||||||
|
|
||||||
|
# from Vanilla ModelScope/StableDiffusion
|
||||||
|
("time_embed.0.weight", "time_embedding.linear_1.weight"),
|
||||||
|
("time_embed.0.bias", "time_embedding.linear_1.bias"),
|
||||||
|
("time_embed.2.weight", "time_embedding.linear_2.weight"),
|
||||||
|
("time_embed.2.bias", "time_embedding.linear_2.bias"),
|
||||||
|
|
||||||
|
|
||||||
|
# from Vanilla ModelScope/StableDiffusion
|
||||||
|
("input_blocks.0.0.weight", "conv_in.weight"),
|
||||||
|
("input_blocks.0.0.bias", "conv_in.bias"),
|
||||||
|
|
||||||
|
|
||||||
|
# from Vanilla ModelScope/StableDiffusion
|
||||||
|
("out.0.weight", "conv_norm_out.weight"),
|
||||||
|
("out.0.bias", "conv_norm_out.bias"),
|
||||||
|
("out.2.weight", "conv_out.weight"),
|
||||||
|
("out.2.bias", "conv_out.bias"),
|
||||||
|
]
|
||||||
|
|
||||||
|
unet_conversion_map_resnet = [
|
||||||
|
# (ModelScope, HF Diffusers)
|
||||||
|
|
||||||
|
# SD
|
||||||
|
("in_layers.0", "norm1"),
|
||||||
|
("in_layers.2", "conv1"),
|
||||||
|
("out_layers.0", "norm2"),
|
||||||
|
("out_layers.3", "conv2"),
|
||||||
|
("emb_layers.1", "time_emb_proj"),
|
||||||
|
("skip_connection", "conv_shortcut"),
|
||||||
|
|
||||||
|
# MS
|
||||||
|
#("temopral_conv", "temp_convs"), # ROFL, they have a typo here --kabachuha
|
||||||
|
]
|
||||||
|
|
||||||
|
unet_conversion_map_layer = []
|
||||||
|
|
||||||
|
# Convert input TemporalTransformer
|
||||||
|
unet_conversion_map_layer.append(('input_blocks.0.1', 'transformer_in'))
|
||||||
|
|
||||||
|
# Reference for the default settings
|
||||||
|
|
||||||
|
# "model_cfg": {
|
||||||
|
# "unet_in_dim": 4,
|
||||||
|
# "unet_dim": 320,
|
||||||
|
# "unet_y_dim": 768,
|
||||||
|
# "unet_context_dim": 1024,
|
||||||
|
# "unet_out_dim": 4,
|
||||||
|
# "unet_dim_mult": [1, 2, 4, 4],
|
||||||
|
# "unet_num_heads": 8,
|
||||||
|
# "unet_head_dim": 64,
|
||||||
|
# "unet_res_blocks": 2,
|
||||||
|
# "unet_attn_scales": [1, 0.5, 0.25],
|
||||||
|
# "unet_dropout": 0.1,
|
||||||
|
# "temporal_attention": "True",
|
||||||
|
# "num_timesteps": 1000,
|
||||||
|
# "mean_type": "eps",
|
||||||
|
# "var_type": "fixed_small",
|
||||||
|
# "loss_type": "mse"
|
||||||
|
# }
|
||||||
|
|
||||||
|
# hardcoded number of downblocks and resnets/attentions...
|
||||||
|
# would need smarter logic for other networks.
|
||||||
|
for i in range(4):
|
||||||
|
# loop over downblocks/upblocks
|
||||||
|
|
||||||
|
for j in range(2):
|
||||||
|
# loop over resnets/attentions for downblocks
|
||||||
|
|
||||||
|
# Spacial SD stuff
|
||||||
|
hf_down_res_prefix = f"down_blocks.{i}.resnets.{j}."
|
||||||
|
sd_down_res_prefix = f"input_blocks.{3*i + j + 1}.0."
|
||||||
|
unet_conversion_map_layer.append((sd_down_res_prefix, hf_down_res_prefix))
|
||||||
|
|
||||||
|
if i < 3:
|
||||||
|
# no attention layers in down_blocks.3
|
||||||
|
hf_down_atn_prefix = f"down_blocks.{i}.attentions.{j}."
|
||||||
|
sd_down_atn_prefix = f"input_blocks.{3*i + j + 1}.1."
|
||||||
|
unet_conversion_map_layer.append((sd_down_atn_prefix, hf_down_atn_prefix))
|
||||||
|
|
||||||
|
# Temporal MS stuff
|
||||||
|
hf_down_res_prefix = f"down_blocks.{i}.temp_convs.{j}."
|
||||||
|
sd_down_res_prefix = f"input_blocks.{3*i + j + 1}.0.temopral_conv."
|
||||||
|
unet_conversion_map_layer.append((sd_down_res_prefix, hf_down_res_prefix))
|
||||||
|
|
||||||
|
if i < 3:
|
||||||
|
# no attention layers in down_blocks.3
|
||||||
|
hf_down_atn_prefix = f"down_blocks.{i}.temp_attentions.{j}."
|
||||||
|
sd_down_atn_prefix = f"input_blocks.{3*i + j + 1}.2."
|
||||||
|
unet_conversion_map_layer.append((sd_down_atn_prefix, hf_down_atn_prefix))
|
||||||
|
|
||||||
|
for j in range(3):
|
||||||
|
# loop over resnets/attentions for upblocks
|
||||||
|
|
||||||
|
# Spacial SD stuff
|
||||||
|
hf_up_res_prefix = f"up_blocks.{i}.resnets.{j}."
|
||||||
|
sd_up_res_prefix = f"output_blocks.{3*i + j}.0."
|
||||||
|
unet_conversion_map_layer.append((sd_up_res_prefix, hf_up_res_prefix))
|
||||||
|
|
||||||
|
if i > 0:
|
||||||
|
# no attention layers in up_blocks.0
|
||||||
|
hf_up_atn_prefix = f"up_blocks.{i}.attentions.{j}."
|
||||||
|
sd_up_atn_prefix = f"output_blocks.{3*i + j}.1."
|
||||||
|
unet_conversion_map_layer.append((sd_up_atn_prefix, hf_up_atn_prefix))
|
||||||
|
|
||||||
|
# loop over resnets/attentions for upblocks
|
||||||
|
hf_up_res_prefix = f"up_blocks.{i}.temp_convs.{j}."
|
||||||
|
sd_up_res_prefix = f"output_blocks.{3*i + j}.0.temopral_conv."
|
||||||
|
unet_conversion_map_layer.append((sd_up_res_prefix, hf_up_res_prefix))
|
||||||
|
|
||||||
|
if i > 0:
|
||||||
|
# no attention layers in up_blocks.0
|
||||||
|
hf_up_atn_prefix = f"up_blocks.{i}.temp_attentions.{j}."
|
||||||
|
sd_up_atn_prefix = f"output_blocks.{3*i + j}.2."
|
||||||
|
unet_conversion_map_layer.append((sd_up_atn_prefix, hf_up_atn_prefix))
|
||||||
|
|
||||||
|
# Up/Downsamplers are 2D, so don't need to touch them
|
||||||
|
if i < 3:
|
||||||
|
# no downsample in down_blocks.3
|
||||||
|
hf_downsample_prefix = f"down_blocks.{i}.downsamplers.0.conv."
|
||||||
|
sd_downsample_prefix = f"input_blocks.{3*(i+1)}.op."
|
||||||
|
unet_conversion_map_layer.append((sd_downsample_prefix, hf_downsample_prefix))
|
||||||
|
|
||||||
|
# no upsample in up_blocks.3
|
||||||
|
hf_upsample_prefix = f"up_blocks.{i}.upsamplers.0."
|
||||||
|
sd_upsample_prefix = f"output_blocks.{3*i + 2}.{1 if i == 0 else 3}."
|
||||||
|
unet_conversion_map_layer.append((sd_upsample_prefix, hf_upsample_prefix))
|
||||||
|
|
||||||
|
|
||||||
|
# Handle the middle block
|
||||||
|
|
||||||
|
# Spacial
|
||||||
|
hf_mid_atn_prefix = "mid_block.attentions.0."
|
||||||
|
sd_mid_atn_prefix = "middle_block.1."
|
||||||
|
unet_conversion_map_layer.append((sd_mid_atn_prefix, hf_mid_atn_prefix))
|
||||||
|
|
||||||
|
for j in range(2):
|
||||||
|
hf_mid_res_prefix = f"mid_block.resnets.{j}."
|
||||||
|
sd_mid_res_prefix = f"middle_block.{3*j}."
|
||||||
|
unet_conversion_map_layer.append((sd_mid_res_prefix, hf_mid_res_prefix))
|
||||||
|
|
||||||
|
# Temporal
|
||||||
|
hf_mid_atn_prefix = "mid_block.temp_attentions.0."
|
||||||
|
sd_mid_atn_prefix = "middle_block.2."
|
||||||
|
unet_conversion_map_layer.append((sd_mid_atn_prefix, hf_mid_atn_prefix))
|
||||||
|
|
||||||
|
for j in range(2):
|
||||||
|
hf_mid_res_prefix = f"mid_block.temp_convs.{j}."
|
||||||
|
sd_mid_res_prefix = f"middle_block.{3*j}.temopral_conv."
|
||||||
|
unet_conversion_map_layer.append((sd_mid_res_prefix, hf_mid_res_prefix))
|
||||||
|
|
||||||
|
# The pipeline
|
||||||
|
def convert_unet_state_dict(unet_state_dict, strict_mapping=False):
|
||||||
|
print ('Converting the UNET')
|
||||||
|
# buyer beware: this is a *brittle* function,
|
||||||
|
# and correct output requires that all of these pieces interact in
|
||||||
|
# the exact order in which I have arranged them.
|
||||||
|
mapping = {k: k for k in unet_state_dict.keys()}
|
||||||
|
|
||||||
|
for sd_name, hf_name in unet_conversion_map:
|
||||||
|
if strict_mapping:
|
||||||
|
if hf_name in mapping:
|
||||||
|
mapping[hf_name] = sd_name
|
||||||
|
else:
|
||||||
|
mapping[hf_name] = sd_name
|
||||||
|
for k, v in mapping.items():
|
||||||
|
if "resnets" in k:
|
||||||
|
for sd_part, hf_part in unet_conversion_map_resnet:
|
||||||
|
v = v.replace(hf_part, sd_part)
|
||||||
|
mapping[k] = v
|
||||||
|
|
||||||
|
for k, v in mapping.items():
|
||||||
|
for sd_part, hf_part in unet_conversion_map_layer:
|
||||||
|
v = v.replace(hf_part, sd_part)
|
||||||
|
mapping[k] = v
|
||||||
|
|
||||||
|
|
||||||
|
# there must be a pattern, but I don't want to bother atm
|
||||||
|
do_not_unsqueeze = [f'output_blocks.{i}.1.proj_out.weight' for i in range(3, 12)] + [f'output_blocks.{i}.1.proj_in.weight' for i in range(3, 12)] + ['middle_block.1.proj_in.weight', 'middle_block.1.proj_out.weight'] + [f'input_blocks.{i}.1.proj_out.weight' for i in [1, 2, 4, 5, 7, 8]] + [f'input_blocks.{i}.1.proj_in.weight' for i in [1, 2, 4, 5, 7, 8]]
|
||||||
|
print (do_not_unsqueeze)
|
||||||
|
|
||||||
|
new_state_dict = {v: (unet_state_dict[k].unsqueeze(-1) if ('proj_' in k and ('bias' not in k) and (k not in do_not_unsqueeze)) else unet_state_dict[k]) for k, v in mapping.items()}
|
||||||
|
|
||||||
|
for k, v in new_state_dict.items():
|
||||||
|
has_k = False
|
||||||
|
for n in do_not_unsqueeze:
|
||||||
|
if k == n:
|
||||||
|
has_k = True
|
||||||
|
|
||||||
|
if has_k:
|
||||||
|
v = v.squeeze(-1)
|
||||||
|
new_state_dict[k] = v
|
||||||
|
|
||||||
|
return new_state_dict
|
||||||
|
|
||||||
|
# TODO: VAE conversion. We doesn't train it in the most cases, but may be handy for the future --kabachuha
|
||||||
|
# ================#
|
||||||
|
# VAE Conversion #
|
||||||
|
# ================#
|
||||||
|
|
||||||
|
vae_conversion_map = [
|
||||||
|
# (stable-diffusion, HF Diffusers)
|
||||||
|
("nin_shortcut", "conv_shortcut"),
|
||||||
|
("norm_out", "conv_norm_out"),
|
||||||
|
("mid.attn_1.", "mid_block.attentions.0."),
|
||||||
|
]
|
||||||
|
|
||||||
|
for i in range(4):
|
||||||
|
# down_blocks have two resnets
|
||||||
|
for j in range(2):
|
||||||
|
hf_down_prefix = f"encoder.down_blocks.{i}.resnets.{j}."
|
||||||
|
sd_down_prefix = f"encoder.down.{i}.block.{j}."
|
||||||
|
vae_conversion_map.append((sd_down_prefix, hf_down_prefix))
|
||||||
|
|
||||||
|
if i < 3:
|
||||||
|
hf_downsample_prefix = f"down_blocks.{i}.downsamplers.0."
|
||||||
|
sd_downsample_prefix = f"down.{i}.downsample."
|
||||||
|
vae_conversion_map.append((sd_downsample_prefix, hf_downsample_prefix))
|
||||||
|
|
||||||
|
hf_upsample_prefix = f"up_blocks.{i}.upsamplers.0."
|
||||||
|
sd_upsample_prefix = f"up.{3-i}.upsample."
|
||||||
|
vae_conversion_map.append((sd_upsample_prefix, hf_upsample_prefix))
|
||||||
|
|
||||||
|
# up_blocks have three resnets
|
||||||
|
# also, up blocks in hf are numbered in reverse from sd
|
||||||
|
for j in range(5):
|
||||||
|
hf_up_prefix = f"decoder.up_blocks.{i}.resnets.{j}."
|
||||||
|
sd_up_prefix = f"decoder.up.{3-i}.block.{j}."
|
||||||
|
vae_conversion_map.append((sd_up_prefix, hf_up_prefix))
|
||||||
|
|
||||||
|
# this part accounts for mid blocks in both the encoder and the decoder
|
||||||
|
for i in range(2):
|
||||||
|
hf_mid_res_prefix = f"mid_block.resnets.{i}."
|
||||||
|
sd_mid_res_prefix = f"mid.block_{i+1}."
|
||||||
|
vae_conversion_map.append((sd_mid_res_prefix, hf_mid_res_prefix))
|
||||||
|
|
||||||
|
|
||||||
|
vae_conversion_map_attn = [
|
||||||
|
# (stable-diffusion, HF Diffusers)
|
||||||
|
("norm.", "group_norm."),
|
||||||
|
("q.", "query."),
|
||||||
|
("k.", "key."),
|
||||||
|
("v.", "value."),
|
||||||
|
("proj_out.", "proj_attn."),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def reshape_weight_for_sd(w):
|
||||||
|
# convert HF linear weights to SD conv2d weights
|
||||||
|
return w.reshape(*w.shape, 1, 1)
|
||||||
|
|
||||||
|
|
||||||
|
def convert_vae_state_dict(vae_state_dict):
|
||||||
|
mapping = {k: k for k in vae_state_dict.keys()}
|
||||||
|
for k, v in mapping.items():
|
||||||
|
for sd_part, hf_part in vae_conversion_map:
|
||||||
|
v = v.replace(hf_part, sd_part)
|
||||||
|
mapping[k] = v
|
||||||
|
for k, v in mapping.items():
|
||||||
|
if "attentions" in k:
|
||||||
|
for sd_part, hf_part in vae_conversion_map_attn:
|
||||||
|
v = v.replace(hf_part, sd_part)
|
||||||
|
mapping[k] = v
|
||||||
|
new_state_dict = {v: vae_state_dict[k] for k, v in mapping.items()}
|
||||||
|
weights_to_convert = ["q", "k", "v", "proj_out"]
|
||||||
|
for k, v in new_state_dict.items():
|
||||||
|
for weight_name in weights_to_convert:
|
||||||
|
if f"mid.attn_1.{weight_name}.weight" in k:
|
||||||
|
print(f"Reshaping {k} for SD format")
|
||||||
|
new_state_dict[k] = reshape_weight_for_sd(v)
|
||||||
|
return new_state_dict
|
||||||
|
# =========================#
|
||||||
|
# Text Encoder Conversion #
|
||||||
|
# =========================#
|
||||||
|
|
||||||
|
# IT IS THE SAME CLIP ENCODER, SO JUST COPYPASTING IT --kabachuha
|
||||||
|
|
||||||
|
# =========================#
|
||||||
|
# Text Encoder Conversion #
|
||||||
|
# =========================#
|
||||||
|
|
||||||
|
|
||||||
|
textenc_conversion_lst = [
|
||||||
|
# (stable-diffusion, HF Diffusers)
|
||||||
|
("resblocks.", "text_model.encoder.layers."),
|
||||||
|
("ln_1", "layer_norm1"),
|
||||||
|
("ln_2", "layer_norm2"),
|
||||||
|
(".c_fc.", ".fc1."),
|
||||||
|
(".c_proj.", ".fc2."),
|
||||||
|
(".attn", ".self_attn"),
|
||||||
|
("ln_final.", "transformer.text_model.final_layer_norm."),
|
||||||
|
("token_embedding.weight", "transformer.text_model.embeddings.token_embedding.weight"),
|
||||||
|
("positional_embedding", "transformer.text_model.embeddings.position_embedding.weight"),
|
||||||
|
]
|
||||||
|
protected = {re.escape(x[1]): x[0] for x in textenc_conversion_lst}
|
||||||
|
textenc_pattern = re.compile("|".join(protected.keys()))
|
||||||
|
|
||||||
|
# Ordering is from https://github.com/pytorch/pytorch/blob/master/test/cpp/api/modules.cpp
|
||||||
|
code2idx = {"q": 0, "k": 1, "v": 2}
|
||||||
|
|
||||||
|
|
||||||
|
def convert_text_enc_state_dict_v20(text_enc_dict):
|
||||||
|
#print ('Converting the text encoder')
|
||||||
|
new_state_dict = {}
|
||||||
|
capture_qkv_weight = {}
|
||||||
|
capture_qkv_bias = {}
|
||||||
|
for k, v in text_enc_dict.items():
|
||||||
|
if (
|
||||||
|
k.endswith(".self_attn.q_proj.weight")
|
||||||
|
or k.endswith(".self_attn.k_proj.weight")
|
||||||
|
or k.endswith(".self_attn.v_proj.weight")
|
||||||
|
):
|
||||||
|
k_pre = k[: -len(".q_proj.weight")]
|
||||||
|
k_code = k[-len("q_proj.weight")]
|
||||||
|
if k_pre not in capture_qkv_weight:
|
||||||
|
capture_qkv_weight[k_pre] = [None, None, None]
|
||||||
|
capture_qkv_weight[k_pre][code2idx[k_code]] = v
|
||||||
|
continue
|
||||||
|
|
||||||
|
if (
|
||||||
|
k.endswith(".self_attn.q_proj.bias")
|
||||||
|
or k.endswith(".self_attn.k_proj.bias")
|
||||||
|
or k.endswith(".self_attn.v_proj.bias")
|
||||||
|
):
|
||||||
|
k_pre = k[: -len(".q_proj.bias")]
|
||||||
|
k_code = k[-len("q_proj.bias")]
|
||||||
|
if k_pre not in capture_qkv_bias:
|
||||||
|
capture_qkv_bias[k_pre] = [None, None, None]
|
||||||
|
capture_qkv_bias[k_pre][code2idx[k_code]] = v
|
||||||
|
continue
|
||||||
|
|
||||||
|
relabelled_key = textenc_pattern.sub(lambda m: protected[re.escape(m.group(0))], k)
|
||||||
|
new_state_dict[relabelled_key] = v
|
||||||
|
|
||||||
|
for k_pre, tensors in capture_qkv_weight.items():
|
||||||
|
if None in tensors:
|
||||||
|
raise Exception("CORRUPTED MODEL: one of the q-k-v values for the text encoder was missing")
|
||||||
|
relabelled_key = textenc_pattern.sub(lambda m: protected[re.escape(m.group(0))], k_pre)
|
||||||
|
new_state_dict[relabelled_key + ".in_proj_weight"] = torch.cat(tensors)
|
||||||
|
|
||||||
|
for k_pre, tensors in capture_qkv_bias.items():
|
||||||
|
if None in tensors:
|
||||||
|
raise Exception("CORRUPTED MODEL: one of the q-k-v values for the text encoder was missing")
|
||||||
|
relabelled_key = textenc_pattern.sub(lambda m: protected[re.escape(m.group(0))], k_pre)
|
||||||
|
new_state_dict[relabelled_key + ".in_proj_bias"] = torch.cat(tensors)
|
||||||
|
|
||||||
|
return new_state_dict
|
||||||
|
|
||||||
|
|
||||||
|
def convert_text_enc_state_dict(text_enc_dict):
|
||||||
|
return text_enc_dict
|
||||||
|
|
||||||
|
textenc_conversion_lst = [
|
||||||
|
# (stable-diffusion, HF Diffusers)
|
||||||
|
("resblocks.", "text_model.encoder.layers."),
|
||||||
|
("ln_1", "layer_norm1"),
|
||||||
|
("ln_2", "layer_norm2"),
|
||||||
|
(".c_fc.", ".fc1."),
|
||||||
|
(".c_proj.", ".fc2."),
|
||||||
|
(".attn", ".self_attn"),
|
||||||
|
("ln_final.", "transformer.text_model.final_layer_norm."),
|
||||||
|
("token_embedding.weight", "transformer.text_model.embeddings.token_embedding.weight"),
|
||||||
|
("positional_embedding", "transformer.text_model.embeddings.position_embedding.weight"),
|
||||||
|
]
|
||||||
|
protected = {re.escape(x[1]): x[0] for x in textenc_conversion_lst}
|
||||||
|
textenc_pattern = re.compile("|".join(protected.keys()))
|
||||||
|
|
||||||
|
# Ordering is from https://github.com/pytorch/pytorch/blob/master/test/cpp/api/modules.cpp
|
||||||
|
code2idx = {"q": 0, "k": 1, "v": 2}
|
||||||
|
|
||||||
|
|
||||||
|
def convert_text_enc_state_dict_v20(text_enc_dict):
|
||||||
|
new_state_dict = {}
|
||||||
|
capture_qkv_weight = {}
|
||||||
|
capture_qkv_bias = {}
|
||||||
|
for k, v in text_enc_dict.items():
|
||||||
|
if (
|
||||||
|
k.endswith(".self_attn.q_proj.weight")
|
||||||
|
or k.endswith(".self_attn.k_proj.weight")
|
||||||
|
or k.endswith(".self_attn.v_proj.weight")
|
||||||
|
):
|
||||||
|
k_pre = k[: -len(".q_proj.weight")]
|
||||||
|
k_code = k[-len("q_proj.weight")]
|
||||||
|
if k_pre not in capture_qkv_weight:
|
||||||
|
capture_qkv_weight[k_pre] = [None, None, None]
|
||||||
|
capture_qkv_weight[k_pre][code2idx[k_code]] = v
|
||||||
|
continue
|
||||||
|
|
||||||
|
if (
|
||||||
|
k.endswith(".self_attn.q_proj.bias")
|
||||||
|
or k.endswith(".self_attn.k_proj.bias")
|
||||||
|
or k.endswith(".self_attn.v_proj.bias")
|
||||||
|
):
|
||||||
|
k_pre = k[: -len(".q_proj.bias")]
|
||||||
|
k_code = k[-len("q_proj.bias")]
|
||||||
|
if k_pre not in capture_qkv_bias:
|
||||||
|
capture_qkv_bias[k_pre] = [None, None, None]
|
||||||
|
capture_qkv_bias[k_pre][code2idx[k_code]] = v
|
||||||
|
continue
|
||||||
|
|
||||||
|
relabelled_key = textenc_pattern.sub(lambda m: protected[re.escape(m.group(0))], k)
|
||||||
|
new_state_dict[relabelled_key] = v
|
||||||
|
|
||||||
|
for k_pre, tensors in capture_qkv_weight.items():
|
||||||
|
if None in tensors:
|
||||||
|
raise Exception("CORRUPTED MODEL: one of the q-k-v values for the text encoder was missing")
|
||||||
|
relabelled_key = textenc_pattern.sub(lambda m: protected[re.escape(m.group(0))], k_pre)
|
||||||
|
new_state_dict[relabelled_key + ".in_proj_weight"] = torch.cat(tensors)
|
||||||
|
|
||||||
|
for k_pre, tensors in capture_qkv_bias.items():
|
||||||
|
if None in tensors:
|
||||||
|
raise Exception("CORRUPTED MODEL: one of the q-k-v values for the text encoder was missing")
|
||||||
|
relabelled_key = textenc_pattern.sub(lambda m: protected[re.escape(m.group(0))], k_pre)
|
||||||
|
new_state_dict[relabelled_key + ".in_proj_bias"] = torch.cat(tensors)
|
||||||
|
|
||||||
|
return new_state_dict
|
||||||
|
|
||||||
|
|
||||||
|
def convert_text_enc_state_dict(text_enc_dict):
|
||||||
|
return text_enc_dict
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
|
||||||
|
parser.add_argument("--model_path", default=None, type=str, required=True, help="Path to the model to convert.")
|
||||||
|
parser.add_argument("--checkpoint_path", default=None, type=str, required=True, help="Path to the output model.")
|
||||||
|
parser.add_argument("--clip_checkpoint_path", default=None, type=str, help="Path to the output CLIP model.")
|
||||||
|
parser.add_argument("--half", action="store_true", help="Save weights in half precision.")
|
||||||
|
parser.add_argument(
|
||||||
|
"--use_safetensors", action="store_true", help="Save weights use safetensors, default is ckpt."
|
||||||
|
)
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
print ('Initializing the conversion map')
|
||||||
|
|
||||||
|
assert args.model_path is not None, "Must provide a model path!"
|
||||||
|
|
||||||
|
assert args.checkpoint_path is not None, "Must provide a checkpoint path!"
|
||||||
|
|
||||||
|
assert args.clip_checkpoint_path is not None, "Must provide a CLIP checkpoint path!"
|
||||||
|
|
||||||
|
# Path for safetensors
|
||||||
|
unet_path = osp.join(args.model_path, "unet", "diffusion_pytorch_model.safetensors")
|
||||||
|
vae_path = osp.join(args.model_path, "vae", "diffusion_pytorch_model.safetensors")
|
||||||
|
text_enc_path = osp.join(args.model_path, "text_encoder", "model.safetensors")
|
||||||
|
|
||||||
|
# Load models from safetensors if it exists, if it doesn't pytorch
|
||||||
|
if osp.exists(unet_path):
|
||||||
|
unet_state_dict = load_file(unet_path, device="cpu")
|
||||||
|
else:
|
||||||
|
unet_path = osp.join(args.model_path, "unet", "diffusion_pytorch_model.bin")
|
||||||
|
unet_state_dict = torch.load(unet_path, map_location="cpu")
|
||||||
|
|
||||||
|
if osp.exists(vae_path):
|
||||||
|
vae_state_dict = load_file(vae_path, device="cpu")
|
||||||
|
else:
|
||||||
|
vae_state_dict = None
|
||||||
|
|
||||||
|
if osp.exists(text_enc_path):
|
||||||
|
text_enc_dict = load_file(text_enc_path, device="cpu")
|
||||||
|
else:
|
||||||
|
text_enc_path = osp.join(args.model_path, "text_encoder", "pytorch_model.bin")
|
||||||
|
text_enc_dict = torch.load(text_enc_path, map_location="cpu")
|
||||||
|
|
||||||
|
# Convert the UNet model
|
||||||
|
unet_state_dict = convert_unet_state_dict(unet_state_dict)
|
||||||
|
|
||||||
|
# Convert the VAE model
|
||||||
|
vae_state_dict = convert_vae_state_dict(vae_state_dict)
|
||||||
|
vae_state_dict = {"first_stage_model." + k: v for k, v in vae_state_dict.items()}
|
||||||
|
|
||||||
|
# Easiest way to identify v2.0 model seems to be that the text encoder (OpenCLIP) is deeper
|
||||||
|
is_v20_model = "text_model.encoder.layers.22.layer_norm2.bias" in text_enc_dict
|
||||||
|
|
||||||
|
if is_v20_model:
|
||||||
|
|
||||||
|
# MODELSCOPE always uses the 2.X encoder, btw --kabachuha
|
||||||
|
|
||||||
|
# Need to add the tag 'transformer' in advance so we can knock it out from the final layer-norm
|
||||||
|
text_enc_dict = {"transformer." + k: v for k, v in text_enc_dict.items()}
|
||||||
|
text_enc_dict = convert_text_enc_state_dict_v20(text_enc_dict)
|
||||||
|
#text_enc_dict = {"cond_stage_model.model." + k: v for k, v in text_enc_dict.items()}
|
||||||
|
else:
|
||||||
|
text_enc_dict = convert_text_enc_state_dict(text_enc_dict)
|
||||||
|
|
||||||
|
# DON'T PUT TOGETHER FOR THE NEW CHECKPOINT AS MODELSCOPE USES THEM IN THE SPLITTED FORM --kabachuha
|
||||||
|
# Save CLIP and the Diffusion model to their own files
|
||||||
|
|
||||||
|
print ('Saving UNET')
|
||||||
|
state_dict = {**unet_state_dict}
|
||||||
|
|
||||||
|
if args.half:
|
||||||
|
state_dict = {k: v.half() for k, v in state_dict.items()}
|
||||||
|
if vae_state_dict is not None:
|
||||||
|
vae_state_dict = {k: v.half() for k, v in vae_state_dict.items()}
|
||||||
|
|
||||||
|
if args.use_safetensors:
|
||||||
|
save_file(state_dict, args.checkpoint_path)
|
||||||
|
|
||||||
|
if vae_state_dict is not None:
|
||||||
|
print("Saving VAE")
|
||||||
|
save_file(vae_state_dict, f"{args.checkpoint_path}.vae")
|
||||||
|
else:
|
||||||
|
if vae_state_dict is not None:
|
||||||
|
print("Saving VAE")
|
||||||
|
vae_state_dict = {'state_dict': vae_state_dict}
|
||||||
|
torch.save(vae_state_dict, f"{args.checkpoint_path}.vae")
|
||||||
|
|
||||||
|
torch.save(state_dict, args.checkpoint_path)
|
||||||
|
|
||||||
|
print('Operation successfull')
|
||||||
@@ -0,0 +1,959 @@
|
|||||||
|
# coding=utf-8
|
||||||
|
# Copyright 2023 The HuggingFace Inc. team.
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
""" Conversion script for the Stable Diffusion checkpoints."""
|
||||||
|
|
||||||
|
import re
|
||||||
|
from io import BytesIO
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
import requests
|
||||||
|
import torch
|
||||||
|
from transformers import (
|
||||||
|
AutoFeatureExtractor,
|
||||||
|
BertTokenizerFast,
|
||||||
|
CLIPImageProcessor,
|
||||||
|
CLIPTextModel,
|
||||||
|
CLIPTextModelWithProjection,
|
||||||
|
CLIPTokenizer,
|
||||||
|
CLIPVisionConfig,
|
||||||
|
CLIPVisionModelWithProjection,
|
||||||
|
)
|
||||||
|
|
||||||
|
from diffusers.models import (
|
||||||
|
AutoencoderKL,
|
||||||
|
PriorTransformer,
|
||||||
|
UNet2DConditionModel,
|
||||||
|
)
|
||||||
|
from diffusers.schedulers import (
|
||||||
|
DDIMScheduler,
|
||||||
|
DDPMScheduler,
|
||||||
|
DPMSolverMultistepScheduler,
|
||||||
|
EulerAncestralDiscreteScheduler,
|
||||||
|
EulerDiscreteScheduler,
|
||||||
|
HeunDiscreteScheduler,
|
||||||
|
LMSDiscreteScheduler,
|
||||||
|
PNDMScheduler,
|
||||||
|
UnCLIPScheduler,
|
||||||
|
)
|
||||||
|
from diffusers.utils.import_utils import BACKENDS_MAPPING
|
||||||
|
|
||||||
|
|
||||||
|
def shave_segments(path, n_shave_prefix_segments=1):
|
||||||
|
"""
|
||||||
|
Removes segments. Positive values shave the first segments, negative shave the last segments.
|
||||||
|
"""
|
||||||
|
if n_shave_prefix_segments >= 0:
|
||||||
|
return ".".join(path.split(".")[n_shave_prefix_segments:])
|
||||||
|
else:
|
||||||
|
return ".".join(path.split(".")[:n_shave_prefix_segments])
|
||||||
|
|
||||||
|
|
||||||
|
def renew_resnet_paths(old_list, n_shave_prefix_segments=0):
|
||||||
|
"""
|
||||||
|
Updates paths inside resnets to the new naming scheme (local renaming)
|
||||||
|
"""
|
||||||
|
mapping = []
|
||||||
|
for old_item in old_list:
|
||||||
|
new_item = old_item.replace("in_layers.0", "norm1")
|
||||||
|
new_item = new_item.replace("in_layers.2", "conv1")
|
||||||
|
|
||||||
|
new_item = new_item.replace("out_layers.0", "norm2")
|
||||||
|
new_item = new_item.replace("out_layers.3", "conv2")
|
||||||
|
|
||||||
|
new_item = new_item.replace("emb_layers.1", "time_emb_proj")
|
||||||
|
new_item = new_item.replace("skip_connection", "conv_shortcut")
|
||||||
|
|
||||||
|
new_item = shave_segments(new_item, n_shave_prefix_segments=n_shave_prefix_segments)
|
||||||
|
|
||||||
|
mapping.append({"old": old_item, "new": new_item})
|
||||||
|
|
||||||
|
return mapping
|
||||||
|
|
||||||
|
|
||||||
|
def renew_vae_resnet_paths(old_list, n_shave_prefix_segments=0):
|
||||||
|
"""
|
||||||
|
Updates paths inside resnets to the new naming scheme (local renaming)
|
||||||
|
"""
|
||||||
|
mapping = []
|
||||||
|
for old_item in old_list:
|
||||||
|
new_item = old_item
|
||||||
|
|
||||||
|
new_item = new_item.replace("nin_shortcut", "conv_shortcut")
|
||||||
|
new_item = shave_segments(new_item, n_shave_prefix_segments=n_shave_prefix_segments)
|
||||||
|
|
||||||
|
mapping.append({"old": old_item, "new": new_item})
|
||||||
|
|
||||||
|
return mapping
|
||||||
|
|
||||||
|
|
||||||
|
def renew_attention_paths(old_list, n_shave_prefix_segments=0):
|
||||||
|
"""
|
||||||
|
Updates paths inside attentions to the new naming scheme (local renaming)
|
||||||
|
"""
|
||||||
|
mapping = []
|
||||||
|
for old_item in old_list:
|
||||||
|
new_item = old_item
|
||||||
|
|
||||||
|
# new_item = new_item.replace('norm.weight', 'group_norm.weight')
|
||||||
|
# new_item = new_item.replace('norm.bias', 'group_norm.bias')
|
||||||
|
|
||||||
|
# new_item = new_item.replace('proj_out.weight', 'proj_attn.weight')
|
||||||
|
# new_item = new_item.replace('proj_out.bias', 'proj_attn.bias')
|
||||||
|
|
||||||
|
# new_item = shave_segments(new_item, n_shave_prefix_segments=n_shave_prefix_segments)
|
||||||
|
|
||||||
|
mapping.append({"old": old_item, "new": new_item})
|
||||||
|
|
||||||
|
return mapping
|
||||||
|
|
||||||
|
|
||||||
|
def renew_vae_attention_paths(old_list, n_shave_prefix_segments=0):
|
||||||
|
"""
|
||||||
|
Updates paths inside attentions to the new naming scheme (local renaming)
|
||||||
|
"""
|
||||||
|
mapping = []
|
||||||
|
for old_item in old_list:
|
||||||
|
new_item = old_item
|
||||||
|
|
||||||
|
new_item = new_item.replace("norm.weight", "group_norm.weight")
|
||||||
|
new_item = new_item.replace("norm.bias", "group_norm.bias")
|
||||||
|
|
||||||
|
new_item = new_item.replace("q.weight", "query.weight")
|
||||||
|
new_item = new_item.replace("q.bias", "query.bias")
|
||||||
|
|
||||||
|
new_item = new_item.replace("k.weight", "key.weight")
|
||||||
|
new_item = new_item.replace("k.bias", "key.bias")
|
||||||
|
|
||||||
|
new_item = new_item.replace("v.weight", "value.weight")
|
||||||
|
new_item = new_item.replace("v.bias", "value.bias")
|
||||||
|
|
||||||
|
new_item = new_item.replace("proj_out.weight", "proj_attn.weight")
|
||||||
|
new_item = new_item.replace("proj_out.bias", "proj_attn.bias")
|
||||||
|
|
||||||
|
new_item = shave_segments(new_item, n_shave_prefix_segments=n_shave_prefix_segments)
|
||||||
|
|
||||||
|
mapping.append({"old": old_item, "new": new_item})
|
||||||
|
|
||||||
|
return mapping
|
||||||
|
|
||||||
|
|
||||||
|
def assign_to_checkpoint(
|
||||||
|
paths, checkpoint, old_checkpoint, attention_paths_to_split=None, additional_replacements=None, config=None
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
This does the final conversion step: take locally converted weights and apply a global renaming to them. It splits
|
||||||
|
attention layers, and takes into account additional replacements that may arise.
|
||||||
|
|
||||||
|
Assigns the weights to the new checkpoint.
|
||||||
|
"""
|
||||||
|
assert isinstance(paths, list), "Paths should be a list of dicts containing 'old' and 'new' keys."
|
||||||
|
|
||||||
|
# Splits the attention layers into three variables.
|
||||||
|
if attention_paths_to_split is not None:
|
||||||
|
for path, path_map in attention_paths_to_split.items():
|
||||||
|
old_tensor = old_checkpoint[path]
|
||||||
|
channels = old_tensor.shape[0] // 3
|
||||||
|
|
||||||
|
target_shape = (-1, channels) if len(old_tensor.shape) == 3 else (-1)
|
||||||
|
|
||||||
|
num_heads = old_tensor.shape[0] // config["num_head_channels"] // 3
|
||||||
|
|
||||||
|
old_tensor = old_tensor.reshape((num_heads, 3 * channels // num_heads) + old_tensor.shape[1:])
|
||||||
|
query, key, value = old_tensor.split(channels // num_heads, dim=1)
|
||||||
|
|
||||||
|
checkpoint[path_map["query"]] = query.reshape(target_shape)
|
||||||
|
checkpoint[path_map["key"]] = key.reshape(target_shape)
|
||||||
|
checkpoint[path_map["value"]] = value.reshape(target_shape)
|
||||||
|
|
||||||
|
for path in paths:
|
||||||
|
new_path = path["new"]
|
||||||
|
|
||||||
|
# These have already been assigned
|
||||||
|
if attention_paths_to_split is not None and new_path in attention_paths_to_split:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Global renaming happens here
|
||||||
|
new_path = new_path.replace("middle_block.0", "mid_block.resnets.0")
|
||||||
|
new_path = new_path.replace("middle_block.1", "mid_block.attentions.0")
|
||||||
|
new_path = new_path.replace("middle_block.2", "mid_block.resnets.1")
|
||||||
|
|
||||||
|
if additional_replacements is not None:
|
||||||
|
for replacement in additional_replacements:
|
||||||
|
new_path = new_path.replace(replacement["old"], replacement["new"])
|
||||||
|
|
||||||
|
# proj_attn.weight has to be converted from conv 1D to linear
|
||||||
|
if "proj_attn.weight" in new_path:
|
||||||
|
checkpoint[new_path] = old_checkpoint[path["old"]][:, :, 0]
|
||||||
|
else:
|
||||||
|
checkpoint[new_path] = old_checkpoint[path["old"]]
|
||||||
|
|
||||||
|
|
||||||
|
def conv_attn_to_linear(checkpoint):
|
||||||
|
keys = list(checkpoint.keys())
|
||||||
|
attn_keys = ["query.weight", "key.weight", "value.weight"]
|
||||||
|
for key in keys:
|
||||||
|
if ".".join(key.split(".")[-2:]) in attn_keys:
|
||||||
|
if checkpoint[key].ndim > 2:
|
||||||
|
checkpoint[key] = checkpoint[key][:, :, 0, 0]
|
||||||
|
elif "proj_attn.weight" in key:
|
||||||
|
if checkpoint[key].ndim > 2:
|
||||||
|
checkpoint[key] = checkpoint[key][:, :, 0]
|
||||||
|
|
||||||
|
|
||||||
|
def create_unet_diffusers_config(original_config, image_size: int, controlnet=False):
|
||||||
|
"""
|
||||||
|
Creates a config for the diffusers based on the config of the LDM model.
|
||||||
|
"""
|
||||||
|
if controlnet:
|
||||||
|
unet_params = original_config.model.params.control_stage_config.params
|
||||||
|
else:
|
||||||
|
unet_params = original_config.model.params.unet_config.params
|
||||||
|
|
||||||
|
vae_params = original_config.model.params.first_stage_config.params.ddconfig
|
||||||
|
|
||||||
|
block_out_channels = [unet_params.model_channels * mult for mult in unet_params.channel_mult]
|
||||||
|
|
||||||
|
down_block_types = []
|
||||||
|
resolution = 1
|
||||||
|
for i in range(len(block_out_channels)):
|
||||||
|
block_type = "CrossAttnDownBlock2D" if resolution in unet_params.attention_resolutions else "DownBlock2D"
|
||||||
|
down_block_types.append(block_type)
|
||||||
|
if i != len(block_out_channels) - 1:
|
||||||
|
resolution *= 2
|
||||||
|
|
||||||
|
up_block_types = []
|
||||||
|
for i in range(len(block_out_channels)):
|
||||||
|
block_type = "CrossAttnUpBlock2D" if resolution in unet_params.attention_resolutions else "UpBlock2D"
|
||||||
|
up_block_types.append(block_type)
|
||||||
|
resolution //= 2
|
||||||
|
|
||||||
|
vae_scale_factor = 2 ** (len(vae_params.ch_mult) - 1)
|
||||||
|
|
||||||
|
head_dim = unet_params.num_heads if "num_heads" in unet_params else None
|
||||||
|
use_linear_projection = (
|
||||||
|
unet_params.use_linear_in_transformer if "use_linear_in_transformer" in unet_params else False
|
||||||
|
)
|
||||||
|
if use_linear_projection:
|
||||||
|
# stable diffusion 2-base-512 and 2-768
|
||||||
|
if head_dim is None:
|
||||||
|
head_dim = [5, 10, 20, 20]
|
||||||
|
|
||||||
|
class_embed_type = None
|
||||||
|
projection_class_embeddings_input_dim = None
|
||||||
|
|
||||||
|
if "num_classes" in unet_params:
|
||||||
|
if unet_params.num_classes == "sequential":
|
||||||
|
class_embed_type = "projection"
|
||||||
|
assert "adm_in_channels" in unet_params
|
||||||
|
projection_class_embeddings_input_dim = unet_params.adm_in_channels
|
||||||
|
else:
|
||||||
|
raise NotImplementedError(f"Unknown conditional unet num_classes config: {unet_params.num_classes}")
|
||||||
|
|
||||||
|
config = {
|
||||||
|
"sample_size": image_size // vae_scale_factor,
|
||||||
|
"in_channels": unet_params.in_channels,
|
||||||
|
"down_block_types": tuple(down_block_types),
|
||||||
|
"block_out_channels": tuple(block_out_channels),
|
||||||
|
"layers_per_block": unet_params.num_res_blocks,
|
||||||
|
"cross_attention_dim": unet_params.context_dim,
|
||||||
|
"attention_head_dim": head_dim,
|
||||||
|
"use_linear_projection": use_linear_projection,
|
||||||
|
"class_embed_type": class_embed_type,
|
||||||
|
"projection_class_embeddings_input_dim": projection_class_embeddings_input_dim,
|
||||||
|
}
|
||||||
|
|
||||||
|
if not controlnet:
|
||||||
|
config["out_channels"] = unet_params.out_channels
|
||||||
|
config["up_block_types"] = tuple(up_block_types)
|
||||||
|
|
||||||
|
return config
|
||||||
|
|
||||||
|
|
||||||
|
def create_vae_diffusers_config(original_config, image_size: int):
|
||||||
|
"""
|
||||||
|
Creates a config for the diffusers based on the config of the LDM model.
|
||||||
|
"""
|
||||||
|
vae_params = original_config.model.params.first_stage_config.params.ddconfig
|
||||||
|
_ = original_config.model.params.first_stage_config.params.embed_dim
|
||||||
|
|
||||||
|
block_out_channels = [vae_params.ch * mult for mult in vae_params.ch_mult]
|
||||||
|
down_block_types = ["DownEncoderBlock2D"] * len(block_out_channels)
|
||||||
|
up_block_types = ["UpDecoderBlock2D"] * len(block_out_channels)
|
||||||
|
|
||||||
|
config = {
|
||||||
|
"sample_size": image_size,
|
||||||
|
"in_channels": vae_params.in_channels,
|
||||||
|
"out_channels": vae_params.out_ch,
|
||||||
|
"down_block_types": tuple(down_block_types),
|
||||||
|
"up_block_types": tuple(up_block_types),
|
||||||
|
"block_out_channels": tuple(block_out_channels),
|
||||||
|
"latent_channels": vae_params.z_channels,
|
||||||
|
"layers_per_block": vae_params.num_res_blocks,
|
||||||
|
}
|
||||||
|
return config
|
||||||
|
|
||||||
|
|
||||||
|
def create_diffusers_schedular(original_config):
|
||||||
|
schedular = DDIMScheduler(
|
||||||
|
num_train_timesteps=original_config.model.params.timesteps,
|
||||||
|
beta_start=original_config.model.params.linear_start,
|
||||||
|
beta_end=original_config.model.params.linear_end,
|
||||||
|
beta_schedule="scaled_linear",
|
||||||
|
)
|
||||||
|
return schedular
|
||||||
|
|
||||||
|
|
||||||
|
def create_ldm_bert_config(original_config):
|
||||||
|
bert_params = original_config.model.parms.cond_stage_config.params
|
||||||
|
config = LDMBertConfig(
|
||||||
|
d_model=bert_params.n_embed,
|
||||||
|
encoder_layers=bert_params.n_layer,
|
||||||
|
encoder_ffn_dim=bert_params.n_embed * 4,
|
||||||
|
)
|
||||||
|
return config
|
||||||
|
|
||||||
|
|
||||||
|
def convert_ldm_unet_checkpoint(checkpoint, config, path=None, extract_ema=False, controlnet=False):
|
||||||
|
"""
|
||||||
|
Takes a state dict and a config, and returns a converted checkpoint.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# extract state_dict for UNet
|
||||||
|
unet_state_dict = {}
|
||||||
|
keys = list(checkpoint.keys())
|
||||||
|
|
||||||
|
if controlnet:
|
||||||
|
unet_key = "control_model."
|
||||||
|
else:
|
||||||
|
unet_key = "model.diffusion_model."
|
||||||
|
|
||||||
|
# at least a 100 parameters have to start with `model_ema` in order for the checkpoint to be EMA
|
||||||
|
if sum(k.startswith("model_ema") for k in keys) > 100 and extract_ema:
|
||||||
|
print(f"Checkpoint {path} has both EMA and non-EMA weights.")
|
||||||
|
print(
|
||||||
|
"In this conversion only the EMA weights are extracted. If you want to instead extract the non-EMA"
|
||||||
|
" weights (useful to continue fine-tuning), please make sure to remove the `--extract_ema` flag."
|
||||||
|
)
|
||||||
|
for key in keys:
|
||||||
|
if key.startswith("model.diffusion_model"):
|
||||||
|
flat_ema_key = "model_ema." + "".join(key.split(".")[1:])
|
||||||
|
unet_state_dict[key.replace(unet_key, "")] = checkpoint.pop(flat_ema_key)
|
||||||
|
else:
|
||||||
|
if sum(k.startswith("model_ema") for k in keys) > 100:
|
||||||
|
print(
|
||||||
|
"In this conversion only the non-EMA weights are extracted. If you want to instead extract the EMA"
|
||||||
|
" weights (usually better for inference), please make sure to add the `--extract_ema` flag."
|
||||||
|
)
|
||||||
|
|
||||||
|
for key in keys:
|
||||||
|
if key.startswith(unet_key):
|
||||||
|
unet_state_dict[key.replace(unet_key, "")] = checkpoint.pop(key)
|
||||||
|
|
||||||
|
new_checkpoint = {}
|
||||||
|
|
||||||
|
new_checkpoint["time_embedding.linear_1.weight"] = unet_state_dict["time_embed.0.weight"]
|
||||||
|
new_checkpoint["time_embedding.linear_1.bias"] = unet_state_dict["time_embed.0.bias"]
|
||||||
|
new_checkpoint["time_embedding.linear_2.weight"] = unet_state_dict["time_embed.2.weight"]
|
||||||
|
new_checkpoint["time_embedding.linear_2.bias"] = unet_state_dict["time_embed.2.bias"]
|
||||||
|
|
||||||
|
if config["class_embed_type"] is None:
|
||||||
|
# No parameters to port
|
||||||
|
...
|
||||||
|
elif config["class_embed_type"] == "timestep" or config["class_embed_type"] == "projection":
|
||||||
|
new_checkpoint["class_embedding.linear_1.weight"] = unet_state_dict["label_emb.0.0.weight"]
|
||||||
|
new_checkpoint["class_embedding.linear_1.bias"] = unet_state_dict["label_emb.0.0.bias"]
|
||||||
|
new_checkpoint["class_embedding.linear_2.weight"] = unet_state_dict["label_emb.0.2.weight"]
|
||||||
|
new_checkpoint["class_embedding.linear_2.bias"] = unet_state_dict["label_emb.0.2.bias"]
|
||||||
|
else:
|
||||||
|
raise NotImplementedError(f"Not implemented `class_embed_type`: {config['class_embed_type']}")
|
||||||
|
|
||||||
|
new_checkpoint["conv_in.weight"] = unet_state_dict["input_blocks.0.0.weight"]
|
||||||
|
new_checkpoint["conv_in.bias"] = unet_state_dict["input_blocks.0.0.bias"]
|
||||||
|
|
||||||
|
if not controlnet:
|
||||||
|
new_checkpoint["conv_norm_out.weight"] = unet_state_dict["out.0.weight"]
|
||||||
|
new_checkpoint["conv_norm_out.bias"] = unet_state_dict["out.0.bias"]
|
||||||
|
new_checkpoint["conv_out.weight"] = unet_state_dict["out.2.weight"]
|
||||||
|
new_checkpoint["conv_out.bias"] = unet_state_dict["out.2.bias"]
|
||||||
|
|
||||||
|
# Retrieves the keys for the input blocks only
|
||||||
|
num_input_blocks = len({".".join(layer.split(".")[:2]) for layer in unet_state_dict if "input_blocks" in layer})
|
||||||
|
input_blocks = {
|
||||||
|
layer_id: [key for key in unet_state_dict if f"input_blocks.{layer_id}" in key]
|
||||||
|
for layer_id in range(num_input_blocks)
|
||||||
|
}
|
||||||
|
|
||||||
|
# Retrieves the keys for the middle blocks only
|
||||||
|
num_middle_blocks = len({".".join(layer.split(".")[:2]) for layer in unet_state_dict if "middle_block" in layer})
|
||||||
|
middle_blocks = {
|
||||||
|
layer_id: [key for key in unet_state_dict if f"middle_block.{layer_id}" in key]
|
||||||
|
for layer_id in range(num_middle_blocks)
|
||||||
|
}
|
||||||
|
|
||||||
|
# Retrieves the keys for the output blocks only
|
||||||
|
num_output_blocks = len({".".join(layer.split(".")[:2]) for layer in unet_state_dict if "output_blocks" in layer})
|
||||||
|
output_blocks = {
|
||||||
|
layer_id: [key for key in unet_state_dict if f"output_blocks.{layer_id}" in key]
|
||||||
|
for layer_id in range(num_output_blocks)
|
||||||
|
}
|
||||||
|
|
||||||
|
for i in range(1, num_input_blocks):
|
||||||
|
block_id = (i - 1) // (config["layers_per_block"] + 1)
|
||||||
|
layer_in_block_id = (i - 1) % (config["layers_per_block"] + 1)
|
||||||
|
|
||||||
|
resnets = [
|
||||||
|
key for key in input_blocks[i] if f"input_blocks.{i}.0" in key and f"input_blocks.{i}.0.op" not in key
|
||||||
|
]
|
||||||
|
attentions = [key for key in input_blocks[i] if f"input_blocks.{i}.1" in key]
|
||||||
|
|
||||||
|
if f"input_blocks.{i}.0.op.weight" in unet_state_dict:
|
||||||
|
new_checkpoint[f"down_blocks.{block_id}.downsamplers.0.conv.weight"] = unet_state_dict.pop(
|
||||||
|
f"input_blocks.{i}.0.op.weight"
|
||||||
|
)
|
||||||
|
new_checkpoint[f"down_blocks.{block_id}.downsamplers.0.conv.bias"] = unet_state_dict.pop(
|
||||||
|
f"input_blocks.{i}.0.op.bias"
|
||||||
|
)
|
||||||
|
|
||||||
|
paths = renew_resnet_paths(resnets)
|
||||||
|
meta_path = {"old": f"input_blocks.{i}.0", "new": f"down_blocks.{block_id}.resnets.{layer_in_block_id}"}
|
||||||
|
assign_to_checkpoint(
|
||||||
|
paths, new_checkpoint, unet_state_dict, additional_replacements=[meta_path], config=config
|
||||||
|
)
|
||||||
|
|
||||||
|
if len(attentions):
|
||||||
|
paths = renew_attention_paths(attentions)
|
||||||
|
meta_path = {"old": f"input_blocks.{i}.1", "new": f"down_blocks.{block_id}.attentions.{layer_in_block_id}"}
|
||||||
|
assign_to_checkpoint(
|
||||||
|
paths, new_checkpoint, unet_state_dict, additional_replacements=[meta_path], config=config
|
||||||
|
)
|
||||||
|
|
||||||
|
resnet_0 = middle_blocks[0]
|
||||||
|
attentions = middle_blocks[1]
|
||||||
|
resnet_1 = middle_blocks[2]
|
||||||
|
|
||||||
|
resnet_0_paths = renew_resnet_paths(resnet_0)
|
||||||
|
assign_to_checkpoint(resnet_0_paths, new_checkpoint, unet_state_dict, config=config)
|
||||||
|
|
||||||
|
resnet_1_paths = renew_resnet_paths(resnet_1)
|
||||||
|
assign_to_checkpoint(resnet_1_paths, new_checkpoint, unet_state_dict, config=config)
|
||||||
|
|
||||||
|
attentions_paths = renew_attention_paths(attentions)
|
||||||
|
meta_path = {"old": "middle_block.1", "new": "mid_block.attentions.0"}
|
||||||
|
assign_to_checkpoint(
|
||||||
|
attentions_paths, new_checkpoint, unet_state_dict, additional_replacements=[meta_path], config=config
|
||||||
|
)
|
||||||
|
|
||||||
|
for i in range(num_output_blocks):
|
||||||
|
block_id = i // (config["layers_per_block"] + 1)
|
||||||
|
layer_in_block_id = i % (config["layers_per_block"] + 1)
|
||||||
|
output_block_layers = [shave_segments(name, 2) for name in output_blocks[i]]
|
||||||
|
output_block_list = {}
|
||||||
|
|
||||||
|
for layer in output_block_layers:
|
||||||
|
layer_id, layer_name = layer.split(".")[0], shave_segments(layer, 1)
|
||||||
|
if layer_id in output_block_list:
|
||||||
|
output_block_list[layer_id].append(layer_name)
|
||||||
|
else:
|
||||||
|
output_block_list[layer_id] = [layer_name]
|
||||||
|
|
||||||
|
if len(output_block_list) > 1:
|
||||||
|
resnets = [key for key in output_blocks[i] if f"output_blocks.{i}.0" in key]
|
||||||
|
attentions = [key for key in output_blocks[i] if f"output_blocks.{i}.1" in key]
|
||||||
|
|
||||||
|
resnet_0_paths = renew_resnet_paths(resnets)
|
||||||
|
paths = renew_resnet_paths(resnets)
|
||||||
|
|
||||||
|
meta_path = {"old": f"output_blocks.{i}.0", "new": f"up_blocks.{block_id}.resnets.{layer_in_block_id}"}
|
||||||
|
assign_to_checkpoint(
|
||||||
|
paths, new_checkpoint, unet_state_dict, additional_replacements=[meta_path], config=config
|
||||||
|
)
|
||||||
|
|
||||||
|
output_block_list = {k: sorted(v) for k, v in output_block_list.items()}
|
||||||
|
if ["conv.bias", "conv.weight"] in output_block_list.values():
|
||||||
|
index = list(output_block_list.values()).index(["conv.bias", "conv.weight"])
|
||||||
|
new_checkpoint[f"up_blocks.{block_id}.upsamplers.0.conv.weight"] = unet_state_dict[
|
||||||
|
f"output_blocks.{i}.{index}.conv.weight"
|
||||||
|
]
|
||||||
|
new_checkpoint[f"up_blocks.{block_id}.upsamplers.0.conv.bias"] = unet_state_dict[
|
||||||
|
f"output_blocks.{i}.{index}.conv.bias"
|
||||||
|
]
|
||||||
|
|
||||||
|
# Clear attentions as they have been attributed above.
|
||||||
|
if len(attentions) == 2:
|
||||||
|
attentions = []
|
||||||
|
|
||||||
|
if len(attentions):
|
||||||
|
paths = renew_attention_paths(attentions)
|
||||||
|
meta_path = {
|
||||||
|
"old": f"output_blocks.{i}.1",
|
||||||
|
"new": f"up_blocks.{block_id}.attentions.{layer_in_block_id}",
|
||||||
|
}
|
||||||
|
assign_to_checkpoint(
|
||||||
|
paths, new_checkpoint, unet_state_dict, additional_replacements=[meta_path], config=config
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
resnet_0_paths = renew_resnet_paths(output_block_layers, n_shave_prefix_segments=1)
|
||||||
|
for path in resnet_0_paths:
|
||||||
|
old_path = ".".join(["output_blocks", str(i), path["old"]])
|
||||||
|
new_path = ".".join(["up_blocks", str(block_id), "resnets", str(layer_in_block_id), path["new"]])
|
||||||
|
|
||||||
|
new_checkpoint[new_path] = unet_state_dict[old_path]
|
||||||
|
|
||||||
|
if controlnet:
|
||||||
|
# conditioning embedding
|
||||||
|
|
||||||
|
orig_index = 0
|
||||||
|
|
||||||
|
new_checkpoint["controlnet_cond_embedding.conv_in.weight"] = unet_state_dict.pop(
|
||||||
|
f"input_hint_block.{orig_index}.weight"
|
||||||
|
)
|
||||||
|
new_checkpoint["controlnet_cond_embedding.conv_in.bias"] = unet_state_dict.pop(
|
||||||
|
f"input_hint_block.{orig_index}.bias"
|
||||||
|
)
|
||||||
|
|
||||||
|
orig_index += 2
|
||||||
|
|
||||||
|
diffusers_index = 0
|
||||||
|
|
||||||
|
while diffusers_index < 6:
|
||||||
|
new_checkpoint[f"controlnet_cond_embedding.blocks.{diffusers_index}.weight"] = unet_state_dict.pop(
|
||||||
|
f"input_hint_block.{orig_index}.weight"
|
||||||
|
)
|
||||||
|
new_checkpoint[f"controlnet_cond_embedding.blocks.{diffusers_index}.bias"] = unet_state_dict.pop(
|
||||||
|
f"input_hint_block.{orig_index}.bias"
|
||||||
|
)
|
||||||
|
diffusers_index += 1
|
||||||
|
orig_index += 2
|
||||||
|
|
||||||
|
new_checkpoint["controlnet_cond_embedding.conv_out.weight"] = unet_state_dict.pop(
|
||||||
|
f"input_hint_block.{orig_index}.weight"
|
||||||
|
)
|
||||||
|
new_checkpoint["controlnet_cond_embedding.conv_out.bias"] = unet_state_dict.pop(
|
||||||
|
f"input_hint_block.{orig_index}.bias"
|
||||||
|
)
|
||||||
|
|
||||||
|
# down blocks
|
||||||
|
for i in range(num_input_blocks):
|
||||||
|
new_checkpoint[f"controlnet_down_blocks.{i}.weight"] = unet_state_dict.pop(f"zero_convs.{i}.0.weight")
|
||||||
|
new_checkpoint[f"controlnet_down_blocks.{i}.bias"] = unet_state_dict.pop(f"zero_convs.{i}.0.bias")
|
||||||
|
|
||||||
|
# mid block
|
||||||
|
new_checkpoint["controlnet_mid_block.weight"] = unet_state_dict.pop("middle_block_out.0.weight")
|
||||||
|
new_checkpoint["controlnet_mid_block.bias"] = unet_state_dict.pop("middle_block_out.0.bias")
|
||||||
|
|
||||||
|
return new_checkpoint
|
||||||
|
|
||||||
|
|
||||||
|
def convert_ldm_vae_checkpoint(checkpoint, config):
|
||||||
|
# extract state dict for VAE
|
||||||
|
vae_state_dict = {}
|
||||||
|
vae_key = "first_stage_model."
|
||||||
|
keys = list(checkpoint.keys())
|
||||||
|
for key in keys:
|
||||||
|
if key.startswith(vae_key):
|
||||||
|
vae_state_dict[key.replace(vae_key, "")] = checkpoint.get(key)
|
||||||
|
|
||||||
|
new_checkpoint = {}
|
||||||
|
|
||||||
|
new_checkpoint["encoder.conv_in.weight"] = vae_state_dict["encoder.conv_in.weight"]
|
||||||
|
new_checkpoint["encoder.conv_in.bias"] = vae_state_dict["encoder.conv_in.bias"]
|
||||||
|
new_checkpoint["encoder.conv_out.weight"] = vae_state_dict["encoder.conv_out.weight"]
|
||||||
|
new_checkpoint["encoder.conv_out.bias"] = vae_state_dict["encoder.conv_out.bias"]
|
||||||
|
new_checkpoint["encoder.conv_norm_out.weight"] = vae_state_dict["encoder.norm_out.weight"]
|
||||||
|
new_checkpoint["encoder.conv_norm_out.bias"] = vae_state_dict["encoder.norm_out.bias"]
|
||||||
|
|
||||||
|
new_checkpoint["decoder.conv_in.weight"] = vae_state_dict["decoder.conv_in.weight"]
|
||||||
|
new_checkpoint["decoder.conv_in.bias"] = vae_state_dict["decoder.conv_in.bias"]
|
||||||
|
new_checkpoint["decoder.conv_out.weight"] = vae_state_dict["decoder.conv_out.weight"]
|
||||||
|
new_checkpoint["decoder.conv_out.bias"] = vae_state_dict["decoder.conv_out.bias"]
|
||||||
|
new_checkpoint["decoder.conv_norm_out.weight"] = vae_state_dict["decoder.norm_out.weight"]
|
||||||
|
new_checkpoint["decoder.conv_norm_out.bias"] = vae_state_dict["decoder.norm_out.bias"]
|
||||||
|
|
||||||
|
new_checkpoint["quant_conv.weight"] = vae_state_dict["quant_conv.weight"]
|
||||||
|
new_checkpoint["quant_conv.bias"] = vae_state_dict["quant_conv.bias"]
|
||||||
|
new_checkpoint["post_quant_conv.weight"] = vae_state_dict["post_quant_conv.weight"]
|
||||||
|
new_checkpoint["post_quant_conv.bias"] = vae_state_dict["post_quant_conv.bias"]
|
||||||
|
|
||||||
|
# Retrieves the keys for the encoder down blocks only
|
||||||
|
num_down_blocks = len({".".join(layer.split(".")[:3]) for layer in vae_state_dict if "encoder.down" in layer})
|
||||||
|
down_blocks = {
|
||||||
|
layer_id: [key for key in vae_state_dict if f"down.{layer_id}" in key] for layer_id in range(num_down_blocks)
|
||||||
|
}
|
||||||
|
|
||||||
|
# Retrieves the keys for the decoder up blocks only
|
||||||
|
num_up_blocks = len({".".join(layer.split(".")[:3]) for layer in vae_state_dict if "decoder.up" in layer})
|
||||||
|
up_blocks = {
|
||||||
|
layer_id: [key for key in vae_state_dict if f"up.{layer_id}" in key] for layer_id in range(num_up_blocks)
|
||||||
|
}
|
||||||
|
|
||||||
|
for i in range(num_down_blocks):
|
||||||
|
resnets = [key for key in down_blocks[i] if f"down.{i}" in key and f"down.{i}.downsample" not in key]
|
||||||
|
|
||||||
|
if f"encoder.down.{i}.downsample.conv.weight" in vae_state_dict:
|
||||||
|
new_checkpoint[f"encoder.down_blocks.{i}.downsamplers.0.conv.weight"] = vae_state_dict.pop(
|
||||||
|
f"encoder.down.{i}.downsample.conv.weight"
|
||||||
|
)
|
||||||
|
new_checkpoint[f"encoder.down_blocks.{i}.downsamplers.0.conv.bias"] = vae_state_dict.pop(
|
||||||
|
f"encoder.down.{i}.downsample.conv.bias"
|
||||||
|
)
|
||||||
|
|
||||||
|
paths = renew_vae_resnet_paths(resnets)
|
||||||
|
meta_path = {"old": f"down.{i}.block", "new": f"down_blocks.{i}.resnets"}
|
||||||
|
assign_to_checkpoint(paths, new_checkpoint, vae_state_dict, additional_replacements=[meta_path], config=config)
|
||||||
|
|
||||||
|
mid_resnets = [key for key in vae_state_dict if "encoder.mid.block" in key]
|
||||||
|
num_mid_res_blocks = 2
|
||||||
|
for i in range(1, num_mid_res_blocks + 1):
|
||||||
|
resnets = [key for key in mid_resnets if f"encoder.mid.block_{i}" in key]
|
||||||
|
|
||||||
|
paths = renew_vae_resnet_paths(resnets)
|
||||||
|
meta_path = {"old": f"mid.block_{i}", "new": f"mid_block.resnets.{i - 1}"}
|
||||||
|
assign_to_checkpoint(paths, new_checkpoint, vae_state_dict, additional_replacements=[meta_path], config=config)
|
||||||
|
|
||||||
|
mid_attentions = [key for key in vae_state_dict if "encoder.mid.attn" in key]
|
||||||
|
paths = renew_vae_attention_paths(mid_attentions)
|
||||||
|
meta_path = {"old": "mid.attn_1", "new": "mid_block.attentions.0"}
|
||||||
|
assign_to_checkpoint(paths, new_checkpoint, vae_state_dict, additional_replacements=[meta_path], config=config)
|
||||||
|
conv_attn_to_linear(new_checkpoint)
|
||||||
|
|
||||||
|
for i in range(num_up_blocks):
|
||||||
|
block_id = num_up_blocks - 1 - i
|
||||||
|
resnets = [
|
||||||
|
key for key in up_blocks[block_id] if f"up.{block_id}" in key and f"up.{block_id}.upsample" not in key
|
||||||
|
]
|
||||||
|
|
||||||
|
if f"decoder.up.{block_id}.upsample.conv.weight" in vae_state_dict:
|
||||||
|
new_checkpoint[f"decoder.up_blocks.{i}.upsamplers.0.conv.weight"] = vae_state_dict[
|
||||||
|
f"decoder.up.{block_id}.upsample.conv.weight"
|
||||||
|
]
|
||||||
|
new_checkpoint[f"decoder.up_blocks.{i}.upsamplers.0.conv.bias"] = vae_state_dict[
|
||||||
|
f"decoder.up.{block_id}.upsample.conv.bias"
|
||||||
|
]
|
||||||
|
|
||||||
|
paths = renew_vae_resnet_paths(resnets)
|
||||||
|
meta_path = {"old": f"up.{block_id}.block", "new": f"up_blocks.{i}.resnets"}
|
||||||
|
assign_to_checkpoint(paths, new_checkpoint, vae_state_dict, additional_replacements=[meta_path], config=config)
|
||||||
|
|
||||||
|
mid_resnets = [key for key in vae_state_dict if "decoder.mid.block" in key]
|
||||||
|
num_mid_res_blocks = 2
|
||||||
|
for i in range(1, num_mid_res_blocks + 1):
|
||||||
|
resnets = [key for key in mid_resnets if f"decoder.mid.block_{i}" in key]
|
||||||
|
|
||||||
|
paths = renew_vae_resnet_paths(resnets)
|
||||||
|
meta_path = {"old": f"mid.block_{i}", "new": f"mid_block.resnets.{i - 1}"}
|
||||||
|
assign_to_checkpoint(paths, new_checkpoint, vae_state_dict, additional_replacements=[meta_path], config=config)
|
||||||
|
|
||||||
|
mid_attentions = [key for key in vae_state_dict if "decoder.mid.attn" in key]
|
||||||
|
paths = renew_vae_attention_paths(mid_attentions)
|
||||||
|
meta_path = {"old": "mid.attn_1", "new": "mid_block.attentions.0"}
|
||||||
|
assign_to_checkpoint(paths, new_checkpoint, vae_state_dict, additional_replacements=[meta_path], config=config)
|
||||||
|
conv_attn_to_linear(new_checkpoint)
|
||||||
|
return new_checkpoint
|
||||||
|
|
||||||
|
|
||||||
|
def convert_ldm_bert_checkpoint(checkpoint, config):
|
||||||
|
def _copy_attn_layer(hf_attn_layer, pt_attn_layer):
|
||||||
|
hf_attn_layer.q_proj.weight.data = pt_attn_layer.to_q.weight
|
||||||
|
hf_attn_layer.k_proj.weight.data = pt_attn_layer.to_k.weight
|
||||||
|
hf_attn_layer.v_proj.weight.data = pt_attn_layer.to_v.weight
|
||||||
|
|
||||||
|
hf_attn_layer.out_proj.weight = pt_attn_layer.to_out.weight
|
||||||
|
hf_attn_layer.out_proj.bias = pt_attn_layer.to_out.bias
|
||||||
|
|
||||||
|
def _copy_linear(hf_linear, pt_linear):
|
||||||
|
hf_linear.weight = pt_linear.weight
|
||||||
|
hf_linear.bias = pt_linear.bias
|
||||||
|
|
||||||
|
def _copy_layer(hf_layer, pt_layer):
|
||||||
|
# copy layer norms
|
||||||
|
_copy_linear(hf_layer.self_attn_layer_norm, pt_layer[0][0])
|
||||||
|
_copy_linear(hf_layer.final_layer_norm, pt_layer[1][0])
|
||||||
|
|
||||||
|
# copy attn
|
||||||
|
_copy_attn_layer(hf_layer.self_attn, pt_layer[0][1])
|
||||||
|
|
||||||
|
# copy MLP
|
||||||
|
pt_mlp = pt_layer[1][1]
|
||||||
|
_copy_linear(hf_layer.fc1, pt_mlp.net[0][0])
|
||||||
|
_copy_linear(hf_layer.fc2, pt_mlp.net[2])
|
||||||
|
|
||||||
|
def _copy_layers(hf_layers, pt_layers):
|
||||||
|
for i, hf_layer in enumerate(hf_layers):
|
||||||
|
if i != 0:
|
||||||
|
i += i
|
||||||
|
pt_layer = pt_layers[i : i + 2]
|
||||||
|
_copy_layer(hf_layer, pt_layer)
|
||||||
|
|
||||||
|
hf_model = LDMBertModel(config).eval()
|
||||||
|
|
||||||
|
# copy embeds
|
||||||
|
hf_model.model.embed_tokens.weight = checkpoint.transformer.token_emb.weight
|
||||||
|
hf_model.model.embed_positions.weight.data = checkpoint.transformer.pos_emb.emb.weight
|
||||||
|
|
||||||
|
# copy layer norm
|
||||||
|
_copy_linear(hf_model.model.layer_norm, checkpoint.transformer.norm)
|
||||||
|
|
||||||
|
# copy hidden layers
|
||||||
|
_copy_layers(hf_model.model.layers, checkpoint.transformer.attn_layers.layers)
|
||||||
|
|
||||||
|
_copy_linear(hf_model.to_logits, checkpoint.transformer.to_logits)
|
||||||
|
|
||||||
|
return hf_model
|
||||||
|
|
||||||
|
|
||||||
|
def convert_ldm_clip_checkpoint(checkpoint):
|
||||||
|
text_model = CLIPTextModel.from_pretrained("openai/clip-vit-large-patch14")
|
||||||
|
keys = list(checkpoint.keys())
|
||||||
|
|
||||||
|
text_model_dict = {}
|
||||||
|
|
||||||
|
for key in keys:
|
||||||
|
if key.startswith("cond_stage_model.transformer"):
|
||||||
|
text_model_dict[key[len("cond_stage_model.transformer.") :]] = checkpoint[key]
|
||||||
|
|
||||||
|
text_model.load_state_dict(text_model_dict)
|
||||||
|
|
||||||
|
return text_model
|
||||||
|
|
||||||
|
|
||||||
|
textenc_conversion_lst = [
|
||||||
|
("cond_stage_model.model.positional_embedding", "text_model.embeddings.position_embedding.weight"),
|
||||||
|
("cond_stage_model.model.token_embedding.weight", "text_model.embeddings.token_embedding.weight"),
|
||||||
|
("cond_stage_model.model.ln_final.weight", "text_model.final_layer_norm.weight"),
|
||||||
|
("cond_stage_model.model.ln_final.bias", "text_model.final_layer_norm.bias"),
|
||||||
|
]
|
||||||
|
textenc_conversion_map = {x[0]: x[1] for x in textenc_conversion_lst}
|
||||||
|
|
||||||
|
textenc_transformer_conversion_lst = [
|
||||||
|
# (stable-diffusion, HF Diffusers)
|
||||||
|
("resblocks.", "text_model.encoder.layers."),
|
||||||
|
("ln_1", "layer_norm1"),
|
||||||
|
("ln_2", "layer_norm2"),
|
||||||
|
(".c_fc.", ".fc1."),
|
||||||
|
(".c_proj.", ".fc2."),
|
||||||
|
(".attn", ".self_attn"),
|
||||||
|
("ln_final.", "transformer.text_model.final_layer_norm."),
|
||||||
|
("token_embedding.weight", "transformer.text_model.embeddings.token_embedding.weight"),
|
||||||
|
("positional_embedding", "transformer.text_model.embeddings.position_embedding.weight"),
|
||||||
|
]
|
||||||
|
protected = {re.escape(x[0]): x[1] for x in textenc_transformer_conversion_lst}
|
||||||
|
textenc_pattern = re.compile("|".join(protected.keys()))
|
||||||
|
|
||||||
|
|
||||||
|
def convert_paint_by_example_checkpoint(checkpoint):
|
||||||
|
config = CLIPVisionConfig.from_pretrained("openai/clip-vit-large-patch14")
|
||||||
|
model = PaintByExampleImageEncoder(config)
|
||||||
|
|
||||||
|
keys = list(checkpoint.keys())
|
||||||
|
|
||||||
|
text_model_dict = {}
|
||||||
|
|
||||||
|
for key in keys:
|
||||||
|
if key.startswith("cond_stage_model.transformer"):
|
||||||
|
text_model_dict[key[len("cond_stage_model.transformer.") :]] = checkpoint[key]
|
||||||
|
|
||||||
|
# load clip vision
|
||||||
|
model.model.load_state_dict(text_model_dict)
|
||||||
|
|
||||||
|
# load mapper
|
||||||
|
keys_mapper = {
|
||||||
|
k[len("cond_stage_model.mapper.res") :]: v
|
||||||
|
for k, v in checkpoint.items()
|
||||||
|
if k.startswith("cond_stage_model.mapper")
|
||||||
|
}
|
||||||
|
|
||||||
|
MAPPING = {
|
||||||
|
"attn.c_qkv": ["attn1.to_q", "attn1.to_k", "attn1.to_v"],
|
||||||
|
"attn.c_proj": ["attn1.to_out.0"],
|
||||||
|
"ln_1": ["norm1"],
|
||||||
|
"ln_2": ["norm3"],
|
||||||
|
"mlp.c_fc": ["ff.net.0.proj"],
|
||||||
|
"mlp.c_proj": ["ff.net.2"],
|
||||||
|
}
|
||||||
|
|
||||||
|
mapped_weights = {}
|
||||||
|
for key, value in keys_mapper.items():
|
||||||
|
prefix = key[: len("blocks.i")]
|
||||||
|
suffix = key.split(prefix)[-1].split(".")[-1]
|
||||||
|
name = key.split(prefix)[-1].split(suffix)[0][1:-1]
|
||||||
|
mapped_names = MAPPING[name]
|
||||||
|
|
||||||
|
num_splits = len(mapped_names)
|
||||||
|
for i, mapped_name in enumerate(mapped_names):
|
||||||
|
new_name = ".".join([prefix, mapped_name, suffix])
|
||||||
|
shape = value.shape[0] // num_splits
|
||||||
|
mapped_weights[new_name] = value[i * shape : (i + 1) * shape]
|
||||||
|
|
||||||
|
model.mapper.load_state_dict(mapped_weights)
|
||||||
|
|
||||||
|
# load final layer norm
|
||||||
|
model.final_layer_norm.load_state_dict(
|
||||||
|
{
|
||||||
|
"bias": checkpoint["cond_stage_model.final_ln.bias"],
|
||||||
|
"weight": checkpoint["cond_stage_model.final_ln.weight"],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
# load final proj
|
||||||
|
model.proj_out.load_state_dict(
|
||||||
|
{
|
||||||
|
"bias": checkpoint["proj_out.bias"],
|
||||||
|
"weight": checkpoint["proj_out.weight"],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
# load uncond vector
|
||||||
|
model.uncond_vector.data = torch.nn.Parameter(checkpoint["learnable_vector"])
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
def convert_open_clip_checkpoint(checkpoint):
|
||||||
|
text_model = CLIPTextModel.from_pretrained("stabilityai/stable-diffusion-2", subfolder="text_encoder")
|
||||||
|
|
||||||
|
keys = list(checkpoint.keys())
|
||||||
|
|
||||||
|
text_model_dict = {}
|
||||||
|
|
||||||
|
if "cond_stage_model.model.text_projection" in checkpoint:
|
||||||
|
d_model = int(checkpoint["cond_stage_model.model.text_projection"].shape[0])
|
||||||
|
else:
|
||||||
|
d_model = 1024
|
||||||
|
|
||||||
|
text_model_dict["text_model.embeddings.position_ids"] = text_model.text_model.embeddings.get_buffer("position_ids")
|
||||||
|
|
||||||
|
for key in keys:
|
||||||
|
if "resblocks.23" in key: # Diffusers drops the final layer and only uses the penultimate layer
|
||||||
|
continue
|
||||||
|
if key in textenc_conversion_map:
|
||||||
|
text_model_dict[textenc_conversion_map[key]] = checkpoint[key]
|
||||||
|
if key.startswith("cond_stage_model.model.transformer."):
|
||||||
|
new_key = key[len("cond_stage_model.model.transformer.") :]
|
||||||
|
if new_key.endswith(".in_proj_weight"):
|
||||||
|
new_key = new_key[: -len(".in_proj_weight")]
|
||||||
|
new_key = textenc_pattern.sub(lambda m: protected[re.escape(m.group(0))], new_key)
|
||||||
|
text_model_dict[new_key + ".q_proj.weight"] = checkpoint[key][:d_model, :]
|
||||||
|
text_model_dict[new_key + ".k_proj.weight"] = checkpoint[key][d_model : d_model * 2, :]
|
||||||
|
text_model_dict[new_key + ".v_proj.weight"] = checkpoint[key][d_model * 2 :, :]
|
||||||
|
elif new_key.endswith(".in_proj_bias"):
|
||||||
|
new_key = new_key[: -len(".in_proj_bias")]
|
||||||
|
new_key = textenc_pattern.sub(lambda m: protected[re.escape(m.group(0))], new_key)
|
||||||
|
text_model_dict[new_key + ".q_proj.bias"] = checkpoint[key][:d_model]
|
||||||
|
text_model_dict[new_key + ".k_proj.bias"] = checkpoint[key][d_model : d_model * 2]
|
||||||
|
text_model_dict[new_key + ".v_proj.bias"] = checkpoint[key][d_model * 2 :]
|
||||||
|
else:
|
||||||
|
new_key = textenc_pattern.sub(lambda m: protected[re.escape(m.group(0))], new_key)
|
||||||
|
|
||||||
|
text_model_dict[new_key] = checkpoint[key]
|
||||||
|
|
||||||
|
text_model.load_state_dict(text_model_dict)
|
||||||
|
|
||||||
|
return text_model
|
||||||
|
|
||||||
|
|
||||||
|
def stable_unclip_image_encoder(original_config):
|
||||||
|
"""
|
||||||
|
Returns the image processor and clip image encoder for the img2img unclip pipeline.
|
||||||
|
|
||||||
|
We currently know of two types of stable unclip models which separately use the clip and the openclip image
|
||||||
|
encoders.
|
||||||
|
"""
|
||||||
|
|
||||||
|
image_embedder_config = original_config.model.params.embedder_config
|
||||||
|
|
||||||
|
sd_clip_image_embedder_class = image_embedder_config.target
|
||||||
|
sd_clip_image_embedder_class = sd_clip_image_embedder_class.split(".")[-1]
|
||||||
|
|
||||||
|
if sd_clip_image_embedder_class == "ClipImageEmbedder":
|
||||||
|
clip_model_name = image_embedder_config.params.model
|
||||||
|
|
||||||
|
if clip_model_name == "ViT-L/14":
|
||||||
|
feature_extractor = CLIPImageProcessor()
|
||||||
|
image_encoder = CLIPVisionModelWithProjection.from_pretrained("openai/clip-vit-large-patch14")
|
||||||
|
else:
|
||||||
|
raise NotImplementedError(f"Unknown CLIP checkpoint name in stable diffusion checkpoint {clip_model_name}")
|
||||||
|
|
||||||
|
elif sd_clip_image_embedder_class == "FrozenOpenCLIPImageEmbedder":
|
||||||
|
feature_extractor = CLIPImageProcessor()
|
||||||
|
image_encoder = CLIPVisionModelWithProjection.from_pretrained("laion/CLIP-ViT-H-14-laion2B-s32B-b79K")
|
||||||
|
else:
|
||||||
|
raise NotImplementedError(
|
||||||
|
f"Unknown CLIP image embedder class in stable diffusion checkpoint {sd_clip_image_embedder_class}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return feature_extractor, image_encoder
|
||||||
|
|
||||||
|
|
||||||
|
def stable_unclip_image_noising_components(
|
||||||
|
original_config, clip_stats_path: Optional[str] = None, device: Optional[str] = None
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Returns the noising components for the img2img and txt2img unclip pipelines.
|
||||||
|
|
||||||
|
Converts the stability noise augmentor into
|
||||||
|
1. a `StableUnCLIPImageNormalizer` for holding the CLIP stats
|
||||||
|
2. a `DDPMScheduler` for holding the noise schedule
|
||||||
|
|
||||||
|
If the noise augmentor config specifies a clip stats path, the `clip_stats_path` must be provided.
|
||||||
|
"""
|
||||||
|
noise_aug_config = original_config.model.params.noise_aug_config
|
||||||
|
noise_aug_class = noise_aug_config.target
|
||||||
|
noise_aug_class = noise_aug_class.split(".")[-1]
|
||||||
|
|
||||||
|
if noise_aug_class == "CLIPEmbeddingNoiseAugmentation":
|
||||||
|
noise_aug_config = noise_aug_config.params
|
||||||
|
embedding_dim = noise_aug_config.timestep_dim
|
||||||
|
max_noise_level = noise_aug_config.noise_schedule_config.timesteps
|
||||||
|
beta_schedule = noise_aug_config.noise_schedule_config.beta_schedule
|
||||||
|
|
||||||
|
image_normalizer = StableUnCLIPImageNormalizer(embedding_dim=embedding_dim)
|
||||||
|
image_noising_scheduler = DDPMScheduler(num_train_timesteps=max_noise_level, beta_schedule=beta_schedule)
|
||||||
|
|
||||||
|
if "clip_stats_path" in noise_aug_config:
|
||||||
|
if clip_stats_path is None:
|
||||||
|
raise ValueError("This stable unclip config requires a `clip_stats_path`")
|
||||||
|
|
||||||
|
clip_mean, clip_std = torch.load(clip_stats_path, map_location=device)
|
||||||
|
clip_mean = clip_mean[None, :]
|
||||||
|
clip_std = clip_std[None, :]
|
||||||
|
|
||||||
|
clip_stats_state_dict = {
|
||||||
|
"mean": clip_mean,
|
||||||
|
"std": clip_std,
|
||||||
|
}
|
||||||
|
|
||||||
|
image_normalizer.load_state_dict(clip_stats_state_dict)
|
||||||
|
else:
|
||||||
|
raise NotImplementedError(f"Unknown noise augmentor class: {noise_aug_class}")
|
||||||
|
|
||||||
|
return image_normalizer, image_noising_scheduler
|
||||||
|
|
||||||
|
|
||||||
|
def convert_controlnet_checkpoint(
|
||||||
|
checkpoint, original_config, checkpoint_path, image_size, upcast_attention, extract_ema
|
||||||
|
):
|
||||||
|
ctrlnet_config = create_unet_diffusers_config(original_config, image_size=image_size, controlnet=True)
|
||||||
|
ctrlnet_config["upcast_attention"] = upcast_attention
|
||||||
|
|
||||||
|
ctrlnet_config.pop("sample_size")
|
||||||
|
|
||||||
|
controlnet_model = ControlNetModel(**ctrlnet_config)
|
||||||
|
|
||||||
|
converted_ctrl_checkpoint = convert_ldm_unet_checkpoint(
|
||||||
|
checkpoint, ctrlnet_config, path=checkpoint_path, extract_ema=extract_ema, controlnet=True
|
||||||
|
)
|
||||||
|
|
||||||
|
controlnet_model.load_state_dict(converted_ctrl_checkpoint)
|
||||||
|
|
||||||
|
return controlnet_model
|
||||||
@@ -0,0 +1,152 @@
|
|||||||
|
# coding=utf-8
|
||||||
|
# Copyright 2023, Haofan Wang, Qixun Wang, All rights reserved.
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
#
|
||||||
|
# Changes were made to this source code by Yuwei Guo.
|
||||||
|
""" Conversion script for the LoRA's safetensors checkpoints. """
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from safetensors.torch import load_file
|
||||||
|
|
||||||
|
from diffusers import StableDiffusionPipeline
|
||||||
|
|
||||||
|
|
||||||
|
def load_diffusers_lora(pipeline, state_dict, alpha=1.0):
|
||||||
|
# directly update weight in diffusers model
|
||||||
|
for key in state_dict:
|
||||||
|
# only process lora down key
|
||||||
|
if "up." in key: continue
|
||||||
|
|
||||||
|
up_key = key.replace(".down.", ".up.")
|
||||||
|
model_key = key.replace("processor.", "").replace("_lora", "").replace("down.", "").replace("up.", "")
|
||||||
|
model_key = model_key.replace("to_out.", "to_out.0.")
|
||||||
|
layer_infos = model_key.split(".")[:-1]
|
||||||
|
|
||||||
|
curr_layer = pipeline.unet
|
||||||
|
while len(layer_infos) > 0:
|
||||||
|
temp_name = layer_infos.pop(0)
|
||||||
|
curr_layer = curr_layer.__getattr__(temp_name)
|
||||||
|
|
||||||
|
weight_down = state_dict[key]
|
||||||
|
weight_up = state_dict[up_key]
|
||||||
|
curr_layer.weight.data += alpha * torch.mm(weight_up, weight_down).to(curr_layer.weight.data.device)
|
||||||
|
|
||||||
|
return pipeline
|
||||||
|
|
||||||
|
|
||||||
|
def convert_lora(pipeline, state_dict, LORA_PREFIX_UNET="lora_unet", LORA_PREFIX_TEXT_ENCODER="lora_te", alpha=0.6):
|
||||||
|
# load base model
|
||||||
|
# pipeline = StableDiffusionPipeline.from_pretrained(base_model_path, torch_dtype=torch.float32)
|
||||||
|
|
||||||
|
# load LoRA weight from .safetensors
|
||||||
|
# state_dict = load_file(checkpoint_path)
|
||||||
|
|
||||||
|
visited = []
|
||||||
|
|
||||||
|
# directly update weight in diffusers model
|
||||||
|
for key in state_dict:
|
||||||
|
# it is suggested to print out the key, it usually will be something like below
|
||||||
|
# "lora_te_text_model_encoder_layers_0_self_attn_k_proj.lora_down.weight"
|
||||||
|
|
||||||
|
# as we have set the alpha beforehand, so just skip
|
||||||
|
if ".alpha" in key or key in visited:
|
||||||
|
continue
|
||||||
|
|
||||||
|
if "text" in key:
|
||||||
|
layer_infos = key.split(".")[0].split(LORA_PREFIX_TEXT_ENCODER + "_")[-1].split("_")
|
||||||
|
curr_layer = pipeline.text_encoder
|
||||||
|
else:
|
||||||
|
layer_infos = key.split(".")[0].split(LORA_PREFIX_UNET + "_")[-1].split("_")
|
||||||
|
curr_layer = pipeline.unet
|
||||||
|
|
||||||
|
# find the target layer
|
||||||
|
temp_name = layer_infos.pop(0)
|
||||||
|
while len(layer_infos) > -1:
|
||||||
|
try:
|
||||||
|
curr_layer = curr_layer.__getattr__(temp_name)
|
||||||
|
if len(layer_infos) > 0:
|
||||||
|
temp_name = layer_infos.pop(0)
|
||||||
|
elif len(layer_infos) == 0:
|
||||||
|
break
|
||||||
|
except Exception:
|
||||||
|
if len(temp_name) > 0:
|
||||||
|
temp_name += "_" + layer_infos.pop(0)
|
||||||
|
else:
|
||||||
|
temp_name = layer_infos.pop(0)
|
||||||
|
|
||||||
|
pair_keys = []
|
||||||
|
if "lora_down" in key:
|
||||||
|
pair_keys.append(key.replace("lora_down", "lora_up"))
|
||||||
|
pair_keys.append(key)
|
||||||
|
else:
|
||||||
|
pair_keys.append(key)
|
||||||
|
pair_keys.append(key.replace("lora_up", "lora_down"))
|
||||||
|
|
||||||
|
# update weight
|
||||||
|
if len(state_dict[pair_keys[0]].shape) == 4:
|
||||||
|
weight_up = state_dict[pair_keys[0]].squeeze(3).squeeze(2).to(torch.float32)
|
||||||
|
weight_down = state_dict[pair_keys[1]].squeeze(3).squeeze(2).to(torch.float32)
|
||||||
|
curr_layer.weight.data += alpha * torch.mm(weight_up, weight_down).unsqueeze(2).unsqueeze(3).to(curr_layer.weight.data.device)
|
||||||
|
else:
|
||||||
|
weight_up = state_dict[pair_keys[0]].to(torch.float32)
|
||||||
|
weight_down = state_dict[pair_keys[1]].to(torch.float32)
|
||||||
|
curr_layer.weight.data += alpha * torch.mm(weight_up, weight_down).to(curr_layer.weight.data.device)
|
||||||
|
|
||||||
|
# update visited list
|
||||||
|
for item in pair_keys:
|
||||||
|
visited.append(item)
|
||||||
|
|
||||||
|
return pipeline
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
|
||||||
|
parser.add_argument(
|
||||||
|
"--base_model_path", default=None, type=str, required=True, help="Path to the base model in diffusers format."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--checkpoint_path", default=None, type=str, required=True, help="Path to the checkpoint to convert."
|
||||||
|
)
|
||||||
|
parser.add_argument("--dump_path", default=None, type=str, required=True, help="Path to the output model.")
|
||||||
|
parser.add_argument(
|
||||||
|
"--lora_prefix_unet", default="lora_unet", type=str, help="The prefix of UNet weight in safetensors"
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--lora_prefix_text_encoder",
|
||||||
|
default="lora_te",
|
||||||
|
type=str,
|
||||||
|
help="The prefix of text encoder weight in safetensors",
|
||||||
|
)
|
||||||
|
parser.add_argument("--alpha", default=0.75, type=float, help="The merging ratio in W = W0 + alpha * deltaW")
|
||||||
|
parser.add_argument(
|
||||||
|
"--to_safetensors", action="store_true", help="Whether to store pipeline in safetensors format or not."
|
||||||
|
)
|
||||||
|
parser.add_argument("--device", type=str, help="Device to use (e.g. cpu, cuda:0, cuda:1, etc.)")
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
base_model_path = args.base_model_path
|
||||||
|
checkpoint_path = args.checkpoint_path
|
||||||
|
dump_path = args.dump_path
|
||||||
|
lora_prefix_unet = args.lora_prefix_unet
|
||||||
|
lora_prefix_text_encoder = args.lora_prefix_text_encoder
|
||||||
|
alpha = args.alpha
|
||||||
|
|
||||||
|
pipe = convert(base_model_path, checkpoint_path, lora_prefix_unet, lora_prefix_text_encoder, alpha)
|
||||||
|
|
||||||
|
pipe = pipe.to(args.device)
|
||||||
|
pipe.save_pretrained(args.dump_path, safe_serialization=args.to_safetensors)
|
||||||
@@ -0,0 +1,751 @@
|
|||||||
|
import os
|
||||||
|
import decord
|
||||||
|
import numpy as np
|
||||||
|
import random
|
||||||
|
import json
|
||||||
|
import torchvision
|
||||||
|
import torchvision.transforms as T
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from glob import glob
|
||||||
|
from PIL import Image
|
||||||
|
from itertools import islice
|
||||||
|
from pathlib import Path
|
||||||
|
from .bucketing import sensible_buckets
|
||||||
|
|
||||||
|
decord.bridge.set_bridge('torch')
|
||||||
|
|
||||||
|
from torch.utils.data import Dataset
|
||||||
|
from einops import rearrange, repeat
|
||||||
|
|
||||||
|
TRAIN_DATA_VARS = ['train_data', 'frames', 'image_dir', 'video_files']
|
||||||
|
VID_TYPES = (".mp4", ".avi", ".mov", ".webm", ".flv", ".mjpeg")
|
||||||
|
|
||||||
|
def get_prompt_ids(prompt, tokenizer):
|
||||||
|
prompt_ids = tokenizer(
|
||||||
|
prompt,
|
||||||
|
truncation=True,
|
||||||
|
padding='max_length',
|
||||||
|
max_length=tokenizer.model_max_length,
|
||||||
|
return_tensors="pt",
|
||||||
|
).input_ids
|
||||||
|
return prompt_ids
|
||||||
|
|
||||||
|
def read_caption_file(caption_file):
|
||||||
|
with open(caption_file, 'r', encoding="utf8") as t:
|
||||||
|
return t.read()
|
||||||
|
|
||||||
|
def get_text_prompt(
|
||||||
|
text_prompt: str = '',
|
||||||
|
fallback_prompt: str= '',
|
||||||
|
file_path:str = '',
|
||||||
|
ext_types=['.mp4'],
|
||||||
|
use_caption=False
|
||||||
|
):
|
||||||
|
try:
|
||||||
|
if use_caption:
|
||||||
|
if len(text_prompt) > 1: return text_prompt
|
||||||
|
caption_file = ''
|
||||||
|
# Use caption on per-video basis (One caption PER video)
|
||||||
|
for ext in ext_types:
|
||||||
|
maybe_file = file_path.replace(ext, '.txt')
|
||||||
|
if maybe_file.endswith(ext_types): continue
|
||||||
|
if os.path.exists(maybe_file):
|
||||||
|
caption_file = maybe_file
|
||||||
|
break
|
||||||
|
|
||||||
|
if os.path.exists(caption_file):
|
||||||
|
return read_caption_file(caption_file)
|
||||||
|
|
||||||
|
# Return fallback prompt if no conditions are met.
|
||||||
|
return fallback_prompt
|
||||||
|
|
||||||
|
return text_prompt
|
||||||
|
except:
|
||||||
|
print(f"Couldn't read prompt caption for {file_path}. Using fallback.")
|
||||||
|
return fallback_prompt
|
||||||
|
|
||||||
|
|
||||||
|
def get_video_frames(vr, start_idx, sample_rate=1, max_frames=24):
|
||||||
|
max_range = len(vr)
|
||||||
|
frame_number = sorted((start_idx, max_range))[1]
|
||||||
|
|
||||||
|
frame_range = range(frame_number, max_range, sample_rate)
|
||||||
|
frame_range_indices = list(frame_range)[:max_frames]
|
||||||
|
|
||||||
|
return frame_range_indices
|
||||||
|
|
||||||
|
def process_video(
|
||||||
|
vid_path,
|
||||||
|
use_bucketing,
|
||||||
|
w,
|
||||||
|
h,
|
||||||
|
get_frame_buckets,
|
||||||
|
get_frame_batch,
|
||||||
|
callback=None
|
||||||
|
):
|
||||||
|
resized_h = None
|
||||||
|
resized_w = None
|
||||||
|
|
||||||
|
if use_bucketing:
|
||||||
|
vr = decord.VideoReader(vid_path)
|
||||||
|
resize, height, width = get_frame_buckets(vr)
|
||||||
|
video = get_frame_batch(vr, resize=resize)
|
||||||
|
resized_h, resized_w = height, width
|
||||||
|
else:
|
||||||
|
vr = decord.VideoReader(vid_path, width=w, height=h)
|
||||||
|
video = get_frame_batch(vr)
|
||||||
|
resized_h, resized_w = w, h
|
||||||
|
|
||||||
|
if callback is not None:
|
||||||
|
callback(resized_h, resized_w)
|
||||||
|
|
||||||
|
return video, vr
|
||||||
|
|
||||||
|
|
||||||
|
class DatasetProcessor(object):
|
||||||
|
def __init__(self, cond_processor=None, cond_processor_kwargs={}):
|
||||||
|
self.condition_processor_model_loaded = False
|
||||||
|
self.condition_processor_kwargs = cond_processor_kwargs
|
||||||
|
self.condition_processor_name = ""
|
||||||
|
self.condition_enabled = False
|
||||||
|
self.resized_w = 0
|
||||||
|
self.resized_h = 0
|
||||||
|
|
||||||
|
def get_frame_range(self, vr):
|
||||||
|
return get_video_frames(
|
||||||
|
vr,
|
||||||
|
self.sample_start_idx,
|
||||||
|
self.frame_step,
|
||||||
|
self.n_sample_frames
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_frame_buckets(self, vr):
|
||||||
|
h, w, c = vr[0].shape
|
||||||
|
width, height = sensible_buckets(
|
||||||
|
self.width,
|
||||||
|
self.height,
|
||||||
|
w,
|
||||||
|
h,
|
||||||
|
extra_simple=False,
|
||||||
|
min_size=256
|
||||||
|
)
|
||||||
|
resize = T.transforms.Resize(
|
||||||
|
(height, width),
|
||||||
|
interpolation=torchvision.transforms.InterpolationMode.BILINEAR
|
||||||
|
)
|
||||||
|
return resize, height, width
|
||||||
|
|
||||||
|
def set_resize_props(self, resized_h, resized_w, *args, **kwargs):
|
||||||
|
self.resized_h = resized_h
|
||||||
|
self.resized_w = resized_w
|
||||||
|
|
||||||
|
def process_video_wrapper(self, vid_path):
|
||||||
|
video, vr = process_video(
|
||||||
|
vid_path,
|
||||||
|
self.use_bucketing,
|
||||||
|
self.width,
|
||||||
|
self.height,
|
||||||
|
self.get_frame_buckets,
|
||||||
|
self.get_frame_batch,
|
||||||
|
callback=self.set_resize_props
|
||||||
|
)
|
||||||
|
|
||||||
|
return video, vr
|
||||||
|
|
||||||
|
def chunk(self, it, size):
|
||||||
|
it = iter(it)
|
||||||
|
return iter(lambda: tuple(islice(it, size)), ())
|
||||||
|
|
||||||
|
def create_video_chunks(
|
||||||
|
self,
|
||||||
|
video_path: str,
|
||||||
|
fps: int,
|
||||||
|
frame_step: int,
|
||||||
|
n_sample_frames: int,
|
||||||
|
max_chunks: int,
|
||||||
|
start_idx: int
|
||||||
|
):
|
||||||
|
# Create a list of frames separated by sample frames
|
||||||
|
# [(1,2,3), (4,5,6), ...]
|
||||||
|
vr = decord.VideoReader(video_path)
|
||||||
|
|
||||||
|
frame_step = min(self.get_avg_fps(vr, fps), 3) if fps > 0 else frame_step
|
||||||
|
vr_range = range(start_idx, len(vr), frame_step)
|
||||||
|
|
||||||
|
frames = list(self.chunk(vr_range, n_sample_frames))
|
||||||
|
|
||||||
|
# Delete any list that contains an out of range index.
|
||||||
|
frames = list(
|
||||||
|
filter(lambda x: len(x) == n_sample_frames, frames)
|
||||||
|
)
|
||||||
|
|
||||||
|
return frames[:self.max_video_clips(frames, max_chunks)]
|
||||||
|
|
||||||
|
def get_avg_fps(self, vr: decord.VideoReader, fps: int = 0):
|
||||||
|
native_fps = vr.get_avg_fps()
|
||||||
|
|
||||||
|
every_nth_frame = max(1, round(native_fps / fps))
|
||||||
|
every_nth_frame = min(len(vr), every_nth_frame)
|
||||||
|
|
||||||
|
return every_nth_frame
|
||||||
|
|
||||||
|
def max_video_clips(self, frames: int, max_chunks: int):
|
||||||
|
return len(frames) if max_chunks == 0 else max_chunks
|
||||||
|
|
||||||
|
# Inspired by the VideoMAE repository.
|
||||||
|
def normalize_input(
|
||||||
|
self,
|
||||||
|
item,
|
||||||
|
mean=[0.485, 0.456, 0.406],
|
||||||
|
std=[0.229, 0.224, 0.225],
|
||||||
|
use_simple_norm=False
|
||||||
|
):
|
||||||
|
if item.dtype == torch.uint8 and not use_simple_norm:
|
||||||
|
import warnings
|
||||||
|
warnings.warn("Using norm based off of ImageNet.")
|
||||||
|
|
||||||
|
item = rearrange(item, 'f c h w -> f h w c')
|
||||||
|
|
||||||
|
item = item.float() / 255.0
|
||||||
|
mean = torch.tensor(mean)
|
||||||
|
std = torch.tensor(std)
|
||||||
|
|
||||||
|
out = rearrange((item - mean) / std, 'f h w c -> f c h w')
|
||||||
|
|
||||||
|
return out
|
||||||
|
else:
|
||||||
|
item = item.float() / 255.
|
||||||
|
item = torchvision.transforms.Normalize([0.5] * 3, [0.5] * 3)(item)
|
||||||
|
|
||||||
|
return item
|
||||||
|
|
||||||
|
def _example(self, item, prompt_ids, prompt):
|
||||||
|
example = {
|
||||||
|
"pixel_values": self.normalize_input(item, use_simple_norm=True),
|
||||||
|
"resized_h": self.resized_h,
|
||||||
|
"resized_w": self.resized_w,
|
||||||
|
"prompt_ids": prompt_ids,
|
||||||
|
"text_prompt": prompt,
|
||||||
|
'dataset': self.__getname__(),
|
||||||
|
}
|
||||||
|
|
||||||
|
self.resized_h, self.resized_w = 0, 0
|
||||||
|
|
||||||
|
return example
|
||||||
|
|
||||||
|
# https://github.com/ExponentialML/Video-BLIP2-Preprocessor
|
||||||
|
class VideoJsonDataset(DatasetProcessor, Dataset):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
tokenizer = None,
|
||||||
|
width: int = 256,
|
||||||
|
height: int = 256,
|
||||||
|
n_sample_frames: int = 4,
|
||||||
|
sample_start_idx: int = 1,
|
||||||
|
frame_step: int = 1,
|
||||||
|
json_path: str ="",
|
||||||
|
json_data = None,
|
||||||
|
vid_data_key: str = "video_path",
|
||||||
|
preprocessed: bool = False,
|
||||||
|
use_bucketing: bool = False,
|
||||||
|
condition_processor = None,
|
||||||
|
**kwargs
|
||||||
|
):
|
||||||
|
DatasetProcessor.__init__(self, condition_processor, kwargs.get('cond_processor_kwargs', {}))
|
||||||
|
self.vid_types = VID_TYPES
|
||||||
|
self.use_bucketing = use_bucketing
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
self.preprocessed = preprocessed
|
||||||
|
|
||||||
|
self.vid_data_key = vid_data_key
|
||||||
|
self.train_data = self.load_from_json(json_path, json_data)
|
||||||
|
|
||||||
|
self.width = width
|
||||||
|
self.height = height
|
||||||
|
|
||||||
|
self.n_sample_frames = n_sample_frames
|
||||||
|
self.sample_start_idx = sample_start_idx
|
||||||
|
self.frame_step = frame_step
|
||||||
|
|
||||||
|
def build_json(self, json_data):
|
||||||
|
extended_data = []
|
||||||
|
for data in json_data['data']:
|
||||||
|
for nested_data in data['data']:
|
||||||
|
self.build_json_dict(
|
||||||
|
data,
|
||||||
|
nested_data,
|
||||||
|
extended_data
|
||||||
|
)
|
||||||
|
json_data = extended_data
|
||||||
|
return json_data
|
||||||
|
|
||||||
|
def build_json_dict(self, data, nested_data, extended_data):
|
||||||
|
clip_path = nested_data['clip_path'] if 'clip_path' in nested_data else None
|
||||||
|
|
||||||
|
extended_data.append({
|
||||||
|
self.vid_data_key: data[self.vid_data_key],
|
||||||
|
'frame_index': nested_data['frame_index'],
|
||||||
|
'prompt': nested_data['prompt'],
|
||||||
|
'clip_path': clip_path
|
||||||
|
})
|
||||||
|
|
||||||
|
def load_from_json(self, path, json_data):
|
||||||
|
try:
|
||||||
|
with open(path) as jpath:
|
||||||
|
print(f"Loading JSON from {path}")
|
||||||
|
json_data = json.load(jpath)
|
||||||
|
|
||||||
|
return self.build_json(json_data)
|
||||||
|
|
||||||
|
except:
|
||||||
|
self.train_data = []
|
||||||
|
print("Non-existant JSON path. Skipping.")
|
||||||
|
|
||||||
|
def validate_json(self, base_path, path):
|
||||||
|
return os.path.exists(f"{base_path}/{path}")
|
||||||
|
|
||||||
|
def train_data_batch(self, index):
|
||||||
|
|
||||||
|
# If we are training on individual clips.
|
||||||
|
if 'clip_path' in self.train_data[index] and \
|
||||||
|
self.train_data[index]['clip_path'] is not None:
|
||||||
|
|
||||||
|
vid_data = self.train_data[index]
|
||||||
|
|
||||||
|
clip_path = vid_data['clip_path']
|
||||||
|
|
||||||
|
# Get video prompt
|
||||||
|
prompt = vid_data['prompt']
|
||||||
|
|
||||||
|
video, _ = self.process_video_wrapper(clip_path)
|
||||||
|
|
||||||
|
prompt_ids = get_prompt_ids(prompt, self.tokenizer)
|
||||||
|
|
||||||
|
return video, prompt, prompt_ids
|
||||||
|
|
||||||
|
# Assign train data
|
||||||
|
train_data = self.train_data[index]
|
||||||
|
|
||||||
|
# Get the frame of the current index.
|
||||||
|
self.sample_start_idx = train_data['frame_index']
|
||||||
|
|
||||||
|
# Initialize resize
|
||||||
|
resize = None
|
||||||
|
|
||||||
|
video, vr = self.process_video_wrapper(train_data[self.vid_data_key])
|
||||||
|
|
||||||
|
# Get video prompt
|
||||||
|
prompt = train_data['prompt']
|
||||||
|
vr.seek(0)
|
||||||
|
|
||||||
|
prompt_ids = get_prompt_ids(prompt, self.tokenizer)
|
||||||
|
|
||||||
|
return video, prompt, prompt_ids
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def __getname__(): return 'json'
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
if self.train_data is not None:
|
||||||
|
return len(self.train_data)
|
||||||
|
else:
|
||||||
|
return 0
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
|
||||||
|
# Initialize variables
|
||||||
|
video = None
|
||||||
|
prompt = None
|
||||||
|
prompt_ids = None
|
||||||
|
|
||||||
|
# Use default JSON training
|
||||||
|
if self.train_data is not None:
|
||||||
|
video, prompt, prompt_ids = self.train_data_batch(index)
|
||||||
|
|
||||||
|
return self._example(video, prompt_ids, prompt)
|
||||||
|
|
||||||
|
|
||||||
|
class SingleVideoDataset(DatasetProcessor, Dataset):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
tokenizer = None,
|
||||||
|
width: int = 256,
|
||||||
|
height: int = 256,
|
||||||
|
n_sample_frames: int = 4,
|
||||||
|
fps: int = 0,
|
||||||
|
frame_step: int = 1,
|
||||||
|
single_video_path: str = "",
|
||||||
|
single_video_prompt: str = "",
|
||||||
|
use_caption: bool = False,
|
||||||
|
use_bucketing: bool = False,
|
||||||
|
condition_processor = None,
|
||||||
|
max_chunks: int = 0,
|
||||||
|
sample_start_idx: int = 0,
|
||||||
|
**kwargs
|
||||||
|
):
|
||||||
|
DatasetProcessor.__init__(self, condition_processor, kwargs.get('cond_processor_kwargs', {}))
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
self.use_bucketing = use_bucketing
|
||||||
|
self.frames = []
|
||||||
|
self.index = 1
|
||||||
|
self.vid_types = (".mp4", ".avi", ".mov", ".webm", ".flv", ".mjpeg")
|
||||||
|
self.n_sample_frames = n_sample_frames
|
||||||
|
self.fps = fps
|
||||||
|
self.frame_step = frame_step
|
||||||
|
self.max_chunks = max_chunks
|
||||||
|
self.sample_start_idx = sample_start_idx
|
||||||
|
|
||||||
|
self.single_video_path = single_video_path
|
||||||
|
self.single_video_prompt = single_video_prompt
|
||||||
|
self.frames = self.create_video_chunks(
|
||||||
|
single_video_path,
|
||||||
|
fps,
|
||||||
|
frame_step,
|
||||||
|
n_sample_frames,
|
||||||
|
max_chunks,
|
||||||
|
sample_start_idx
|
||||||
|
)
|
||||||
|
|
||||||
|
self.width = width
|
||||||
|
self.height = height
|
||||||
|
|
||||||
|
def get_frame_batch(self, vr, resize=None):
|
||||||
|
index = self.index
|
||||||
|
|
||||||
|
frames = vr.get_batch(self.frames[self.index])
|
||||||
|
video = rearrange(frames, "f h w c -> f c h w")
|
||||||
|
|
||||||
|
if resize is not None:
|
||||||
|
video = resize(video)
|
||||||
|
return video
|
||||||
|
|
||||||
|
def get_prompt(self):
|
||||||
|
vid_ext = self.single_video_path.split(".")[-1]
|
||||||
|
video_path = self.single_video_path
|
||||||
|
|
||||||
|
maybe_text_file = video_path.replace(f".{vid_ext}", ".txt")
|
||||||
|
|
||||||
|
if os.path.exists(maybe_text_file):
|
||||||
|
with open(maybe_text_file, "r") as f:
|
||||||
|
prompt = f.read()
|
||||||
|
else:
|
||||||
|
prompt = self.single_video_prompt
|
||||||
|
|
||||||
|
return prompt
|
||||||
|
|
||||||
|
def single_video_batch(self, index):
|
||||||
|
train_data = self.single_video_path
|
||||||
|
self.index = index
|
||||||
|
|
||||||
|
if train_data.endswith(self.vid_types):
|
||||||
|
video, _ = self.process_video_wrapper(train_data)
|
||||||
|
|
||||||
|
prompt = self.get_prompt()
|
||||||
|
prompt_ids = get_prompt_ids(prompt, self.tokenizer)
|
||||||
|
|
||||||
|
return video, prompt, prompt_ids
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Single video is not a video type. Types: {self.vid_types}")
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def __getname__():
|
||||||
|
return 'single_video'
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return len(self.frames)
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
video, prompt, prompt_ids = self.single_video_batch(index)
|
||||||
|
return self._example(video, prompt_ids, prompt)
|
||||||
|
|
||||||
|
class ImageDataset(DatasetProcessor, Dataset):
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
tokenizer = None,
|
||||||
|
width: int = 256,
|
||||||
|
height: int = 256,
|
||||||
|
base_width: int = 256,
|
||||||
|
base_height: int = 256,
|
||||||
|
use_caption: bool = False,
|
||||||
|
image_dir: str = '',
|
||||||
|
single_img_prompt: str = '',
|
||||||
|
use_bucketing: bool = False,
|
||||||
|
fallback_prompt: str = '',
|
||||||
|
condition_processor = None,
|
||||||
|
**kwargs
|
||||||
|
):
|
||||||
|
DatasetProcessor.__init__(self, condition_processor, kwargs.get('cond_processor_kwargs', {}))
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
self.img_types = (".png", ".jpg", ".jpeg", '.bmp')
|
||||||
|
self.use_bucketing = use_bucketing
|
||||||
|
|
||||||
|
self.image_dir = self.get_images_list(image_dir)
|
||||||
|
self.fallback_prompt = fallback_prompt
|
||||||
|
|
||||||
|
self.use_caption = use_caption
|
||||||
|
self.single_img_prompt = single_img_prompt
|
||||||
|
|
||||||
|
self.width = width
|
||||||
|
self.height = height
|
||||||
|
|
||||||
|
def get_images_list(self, image_dir):
|
||||||
|
if os.path.exists(image_dir):
|
||||||
|
imgs = [x for x in os.listdir(image_dir) if x.endswith(self.img_types)]
|
||||||
|
full_img_dir = []
|
||||||
|
|
||||||
|
for img in imgs:
|
||||||
|
full_img_dir.append(f"{image_dir}/{img}")
|
||||||
|
|
||||||
|
return sorted(full_img_dir)
|
||||||
|
|
||||||
|
return ['']
|
||||||
|
|
||||||
|
def image_batch(self, index):
|
||||||
|
train_data = self.image_dir[index]
|
||||||
|
img = train_data
|
||||||
|
|
||||||
|
try:
|
||||||
|
img = torchvision.io.read_image(img, mode=torchvision.io.ImageReadMode.RGB)
|
||||||
|
except:
|
||||||
|
img = T.transforms.PILToTensor()(Image.open(img).convert("RGB"))
|
||||||
|
|
||||||
|
width = self.width
|
||||||
|
height = self.height
|
||||||
|
|
||||||
|
if self.use_bucketing:
|
||||||
|
_, h, w = img.shape
|
||||||
|
width, height = sensible_buckets(width, height, w, h, extra_simple=False)
|
||||||
|
|
||||||
|
resize = T.transforms.Resize((height, width), antialias=True)
|
||||||
|
|
||||||
|
img = resize(img)
|
||||||
|
img = repeat(img, 'c h w -> f c h w', f=1)
|
||||||
|
|
||||||
|
prompt = get_text_prompt(
|
||||||
|
file_path=train_data,
|
||||||
|
text_prompt=self.single_img_prompt,
|
||||||
|
fallback_prompt=self.fallback_prompt,
|
||||||
|
ext_types=self.img_types,
|
||||||
|
use_caption=True
|
||||||
|
)
|
||||||
|
prompt_ids = get_prompt_ids(prompt, self.tokenizer)
|
||||||
|
|
||||||
|
return img, prompt, prompt_ids
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def __getname__(): return 'image'
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
# Image directory
|
||||||
|
if os.path.exists(self.image_dir[0]):
|
||||||
|
return len(self.image_dir)
|
||||||
|
else:
|
||||||
|
return 0
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
img, prompt, prompt_ids = self.image_batch(index)
|
||||||
|
|
||||||
|
return self._example(img, prompt_ids, prompt)
|
||||||
|
|
||||||
|
# NOTE: This is currently unused in this repository. All videos are processed with SingleVideoDataset.
|
||||||
|
# If you are doing folder based training, all single videos are concatenated into a single dataset using ConcatDataset.
|
||||||
|
# The VideoFolderDataset class is still usable, but must be manually set and modified in your training script.
|
||||||
|
class VideoFolderDataset(DatasetProcessor, Dataset):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
tokenizer=None,
|
||||||
|
width: int = 256,
|
||||||
|
height: int = 256,
|
||||||
|
n_sample_frames: int = 16,
|
||||||
|
fps: int = 8,
|
||||||
|
path: str = "./data",
|
||||||
|
fallback_prompt: str = "",
|
||||||
|
use_bucketing: bool = False,
|
||||||
|
condition_processor = None,
|
||||||
|
sample_start_idx: int = 0,
|
||||||
|
**kwargs
|
||||||
|
):
|
||||||
|
DatasetProcessor.__init__(self, condition_processor, kwargs.get('cond_processor_kwargs', {}))
|
||||||
|
self.tokenizer = tokenizer
|
||||||
|
self.use_bucketing = use_bucketing
|
||||||
|
|
||||||
|
self.fallback_prompt = fallback_prompt
|
||||||
|
|
||||||
|
self.video_files = glob(f"{path}/*.mp4")
|
||||||
|
|
||||||
|
self.width = width
|
||||||
|
self.height = height
|
||||||
|
|
||||||
|
self.sample_start_idx = sample_start_idx
|
||||||
|
self.n_sample_frames = n_sample_frames
|
||||||
|
self.fps = fps
|
||||||
|
|
||||||
|
def get_frame_batch(self, vr, resize=None):
|
||||||
|
n_sample_frames = self.n_sample_frames
|
||||||
|
native_fps = vr.get_avg_fps()
|
||||||
|
|
||||||
|
every_nth_frame = max(1, round(native_fps / self.fps))
|
||||||
|
every_nth_frame = min(len(vr), every_nth_frame)
|
||||||
|
|
||||||
|
effective_length = len(vr) // every_nth_frame
|
||||||
|
if effective_length < n_sample_frames:
|
||||||
|
n_sample_frames = effective_length
|
||||||
|
|
||||||
|
effective_idx = random.randint(0, (effective_length - n_sample_frames))
|
||||||
|
idxs = every_nth_frame * np.arange(effective_idx, effective_idx + n_sample_frames)
|
||||||
|
|
||||||
|
video = vr.get_batch(idxs)
|
||||||
|
video = rearrange(video, "f h w c -> f c h w")
|
||||||
|
|
||||||
|
if resize is not None: video = resize(video)
|
||||||
|
return video, vr
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def __getname__(): return 'folder'
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return len(self.video_files)
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
|
||||||
|
video, _ = self.process_video_wrapper(self.video_files[index])
|
||||||
|
|
||||||
|
if os.path.exists(self.video_files[index].replace(".mp4", ".txt")):
|
||||||
|
with open(self.video_files[index].replace(".mp4", ".txt"), "r") as f:
|
||||||
|
prompt = f.read()
|
||||||
|
else:
|
||||||
|
prompt = self.fallback_prompt
|
||||||
|
|
||||||
|
prompt_ids = get_prompt_ids(prompt, self.tokenizer)
|
||||||
|
|
||||||
|
return self._example(video[0], prompt_ids, prompt)
|
||||||
|
|
||||||
|
class CachedDataset(DatasetProcessor, Dataset):
|
||||||
|
def __init__(self, cache_dir: str = ''):
|
||||||
|
DatasetProcessor.__init__(self)
|
||||||
|
self.cache_dir = cache_dir
|
||||||
|
self.cached_data_list = self.get_files_list()
|
||||||
|
|
||||||
|
def get_files_list(self):
|
||||||
|
tensors_list = [f"{self.cache_dir}/{x}" for x in os.listdir(self.cache_dir) if x.endswith('.pt')]
|
||||||
|
return sorted(tensors_list)
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return len(self.cached_data_list)
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
cached_latent = torch.load(self.cached_data_list[index], map_location='cpu')
|
||||||
|
|
||||||
|
return cached_latent
|
||||||
|
|
||||||
|
class ConcatInterleavedDataset(Dataset):
|
||||||
|
def __init__(self, datasets):
|
||||||
|
self.datasets = datasets
|
||||||
|
self.train_data_vars = TRAIN_DATA_VARS
|
||||||
|
|
||||||
|
self.interleave_datasets()
|
||||||
|
|
||||||
|
def get_parent_dataset(self):
|
||||||
|
|
||||||
|
# There's a chance that the subset images may be bigger than the video if doing text training.
|
||||||
|
# If it has the attribute "is_subset", we can simply ignore it to ensure it isn't the biggest
|
||||||
|
# length.
|
||||||
|
dataset_lengths = [d.__len__() if not hasattr(d, 'is_subset') else 0 for d in self.datasets]
|
||||||
|
max_dataset_index = dataset_lengths.index(max(dataset_lengths))
|
||||||
|
|
||||||
|
parent_dataset = self.datasets[max_dataset_index]
|
||||||
|
|
||||||
|
return parent_dataset, max_dataset_index
|
||||||
|
|
||||||
|
def process_dataset(self, dataset):
|
||||||
|
processed_dataset = []
|
||||||
|
train_data_var_name = self.get_dataset_data_var_name(dataset)[0]
|
||||||
|
train_data_var = getattr(dataset, train_data_var_name)
|
||||||
|
|
||||||
|
for idx, item in enumerate(train_data_var):
|
||||||
|
if isinstance(item, dict) and 'idx_modulo' in item:
|
||||||
|
ref_idx = item['idx_modulo']
|
||||||
|
already_processed_item = processed_dataset[ref_idx]
|
||||||
|
|
||||||
|
# Dataset items are assumed to be of type Dict
|
||||||
|
already_processed_item['reference_idx'] = ref_idx
|
||||||
|
processed_dataset.append(already_processed_item)
|
||||||
|
else:
|
||||||
|
processed_dataset.append(dataset[idx])
|
||||||
|
|
||||||
|
return processed_dataset
|
||||||
|
|
||||||
|
def get_dataset_data_var_name(self, dataset):
|
||||||
|
return [v for v in self.train_data_vars if v in dataset.__dict__.keys()]
|
||||||
|
|
||||||
|
def create_data_val_dict(self, val, idx, length, idx_modulo):
|
||||||
|
return dict(
|
||||||
|
value=val,
|
||||||
|
idx=idx,
|
||||||
|
length=length,
|
||||||
|
idx_modulo=idx_modulo
|
||||||
|
)
|
||||||
|
|
||||||
|
def interleave_datasets(self):
|
||||||
|
parent_dataset, parent_dataset_index = self.get_parent_dataset()
|
||||||
|
child_datasets = self.datasets.copy()
|
||||||
|
child_datasets.pop(parent_dataset_index)
|
||||||
|
|
||||||
|
parent_dataset_length = parent_dataset.__len__()
|
||||||
|
|
||||||
|
for dataset in child_datasets:
|
||||||
|
if dataset.__len__() <= 0:
|
||||||
|
del dataset
|
||||||
|
continue
|
||||||
|
|
||||||
|
var_name = self.get_dataset_data_var_name(dataset)
|
||||||
|
var_name = var_name[0] if len(var_name) == 1 else None
|
||||||
|
|
||||||
|
if var_name is None:
|
||||||
|
continue
|
||||||
|
|
||||||
|
original_dataset_length = dataset.__len__()
|
||||||
|
|
||||||
|
train_data_var = getattr(dataset, var_name)
|
||||||
|
train_data_var *= parent_dataset_length
|
||||||
|
new_train_data_val = train_data_var[:parent_dataset_length]
|
||||||
|
|
||||||
|
# Do this to reference items that were already accessed.
|
||||||
|
# Since some __getitem__ functions are heavy (numpy computations, video reads, etc.),
|
||||||
|
# we want to avoid performing the same expensive function multiple times.
|
||||||
|
# We simply point to the corresponding index so that when we interleave, we can just copy the __getitem__ result.
|
||||||
|
for i, val in enumerate(new_train_data_val):
|
||||||
|
if i >= original_dataset_length:
|
||||||
|
clamped_idx = i % original_dataset_length
|
||||||
|
new_train_data_val[i] = self.create_data_val_dict(
|
||||||
|
val,
|
||||||
|
i,
|
||||||
|
original_dataset_length,
|
||||||
|
clamped_idx
|
||||||
|
)
|
||||||
|
|
||||||
|
setattr(dataset, var_name, new_train_data_val)
|
||||||
|
|
||||||
|
from itertools import chain
|
||||||
|
|
||||||
|
print("Interleaving Datasets. Please wait...")
|
||||||
|
train_datasets = [parent_dataset] + child_datasets
|
||||||
|
|
||||||
|
# Zip all of the items in the datasets. We do this to __get_item__ all of our data.
|
||||||
|
# Example (d == Dataset): [(d1_item1, d2_item1, d3_item1), (d1_item2, d2_item2, d3_item2), (...)]
|
||||||
|
interleave_datasets = zip(*[self.process_dataset(d) for d in train_datasets])
|
||||||
|
|
||||||
|
# Now we flatten it as a new Dataset iterable Dataset to be concatenated.
|
||||||
|
# Example: [d1_item1, d2_item1, d3_item1, d2_item1, d2_item2, d2_item3, ...]
|
||||||
|
InterLeavedDataset = list(chain(*interleave_datasets))
|
||||||
|
self.datasets = InterLeavedDataset
|
||||||
|
|
||||||
|
print("Finished interleaving datasets.")
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return len(self.datasets)
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
return self.datasets[index]
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,435 @@
|
|||||||
|
import os
|
||||||
|
from logging import warnings
|
||||||
|
import torch
|
||||||
|
from typing import Union
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from ...animatediff.models.unet import UNet3DConditionModel
|
||||||
|
from transformers import CLIPTextModel
|
||||||
|
from ...animatediff.utils.convert_diffusers_to_original_ms_text_to_video import convert_unet_state_dict, convert_text_enc_state_dict_v20
|
||||||
|
|
||||||
|
from .lora import (
|
||||||
|
extract_lora_ups_down,
|
||||||
|
inject_trainable_lora_extended,
|
||||||
|
save_lora_weight,
|
||||||
|
save_lora_safetensors,
|
||||||
|
train_patch_pipe,
|
||||||
|
monkeypatch_or_replace_lora,
|
||||||
|
monkeypatch_or_replace_lora_extended
|
||||||
|
)
|
||||||
|
|
||||||
|
from ...animatediff.stable_lora.lora import (
|
||||||
|
activate_lora_train,
|
||||||
|
add_lora_to,
|
||||||
|
save_lora,
|
||||||
|
load_lora,
|
||||||
|
set_mode_group
|
||||||
|
)
|
||||||
|
|
||||||
|
FILE_BASENAMES = ['unet', 'text_encoder']
|
||||||
|
LORA_FILE_TYPES = ['.pt', '.safetensors']
|
||||||
|
CLONE_OF_SIMO_KEYS = ['model', 'loras', 'target_replace_module', 'r']
|
||||||
|
STABLE_LORA_KEYS = [
|
||||||
|
'model',
|
||||||
|
'target_module',
|
||||||
|
'search_class',
|
||||||
|
'r',
|
||||||
|
'dropout',
|
||||||
|
'lora_bias',
|
||||||
|
'scale'
|
||||||
|
]
|
||||||
|
|
||||||
|
lora_versions = dict(
|
||||||
|
stable_lora = "stable_lora",
|
||||||
|
cloneofsimo = "cloneofsimo"
|
||||||
|
)
|
||||||
|
|
||||||
|
lora_func_types = dict(
|
||||||
|
loader = "loader",
|
||||||
|
injector = "injector"
|
||||||
|
)
|
||||||
|
|
||||||
|
lora_args = dict(
|
||||||
|
model = None,
|
||||||
|
loras = None,
|
||||||
|
target_replace_module = [],
|
||||||
|
target_module = [],
|
||||||
|
r = 4,
|
||||||
|
search_class = [torch.nn.Linear],
|
||||||
|
dropout = 0,
|
||||||
|
lora_bias = 'none',
|
||||||
|
scale = 0
|
||||||
|
)
|
||||||
|
|
||||||
|
LoraVersions = SimpleNamespace(**lora_versions)
|
||||||
|
LoraFuncTypes = SimpleNamespace(**lora_func_types)
|
||||||
|
|
||||||
|
LORA_VERSIONS = [LoraVersions.stable_lora, LoraVersions.cloneofsimo]
|
||||||
|
LORA_FUNC_TYPES = [LoraFuncTypes.loader, LoraFuncTypes.injector]
|
||||||
|
|
||||||
|
def filter_dict(_dict, keys=[]):
|
||||||
|
if len(keys) == 0:
|
||||||
|
assert "Keys cannot empty for filtering return dict."
|
||||||
|
|
||||||
|
for k in keys:
|
||||||
|
if k not in lora_args.keys():
|
||||||
|
assert f"{k} does not exist in available LoRA arguments"
|
||||||
|
|
||||||
|
return {k: v for k, v in _dict.items() if k in keys}
|
||||||
|
|
||||||
|
class LoraHandler(object):
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
version: LORA_VERSIONS = LoraVersions.cloneofsimo,
|
||||||
|
use_unet_lora: bool = False,
|
||||||
|
use_text_lora: bool = False,
|
||||||
|
save_for_webui: bool = False,
|
||||||
|
only_for_webui: bool = False,
|
||||||
|
lora_bias: str = 'none',
|
||||||
|
unet_replace_modules: list = ['UNet3DConditionModel'],
|
||||||
|
text_encoder_replace_modules: list = ['CLIPEncoderLayer']
|
||||||
|
):
|
||||||
|
self.version = version
|
||||||
|
self.lora_loader = self.get_lora_func(func_type=LoraFuncTypes.loader)
|
||||||
|
self.lora_injector = self.get_lora_func(func_type=LoraFuncTypes.injector)
|
||||||
|
self.lora_bias = lora_bias
|
||||||
|
self.use_unet_lora = use_unet_lora
|
||||||
|
self.use_text_lora = use_text_lora
|
||||||
|
self.save_for_webui = save_for_webui
|
||||||
|
self.only_for_webui = only_for_webui
|
||||||
|
self.unet_replace_modules = unet_replace_modules
|
||||||
|
self.text_encoder_replace_modules = text_encoder_replace_modules
|
||||||
|
self.use_lora = any([use_text_lora, use_unet_lora])
|
||||||
|
|
||||||
|
if self.use_lora:
|
||||||
|
print(f"Using LoRA Version: {self.version}")
|
||||||
|
|
||||||
|
def is_cloneofsimo_lora(self):
|
||||||
|
return self.version == LoraVersions.cloneofsimo
|
||||||
|
|
||||||
|
def is_stable_lora(self):
|
||||||
|
return self.version == LoraVersions.stable_lora
|
||||||
|
|
||||||
|
def get_lora_func(self, func_type: LORA_FUNC_TYPES = LoraFuncTypes.loader):
|
||||||
|
|
||||||
|
if self.is_cloneofsimo_lora():
|
||||||
|
|
||||||
|
if func_type == LoraFuncTypes.loader:
|
||||||
|
return monkeypatch_or_replace_lora_extended
|
||||||
|
|
||||||
|
if func_type == LoraFuncTypes.injector:
|
||||||
|
return inject_trainable_lora_extended
|
||||||
|
|
||||||
|
if self.is_stable_lora():
|
||||||
|
|
||||||
|
if func_type == LoraFuncTypes.loader:
|
||||||
|
return load_lora
|
||||||
|
|
||||||
|
if func_type == LoraFuncTypes.injector:
|
||||||
|
return add_lora_to
|
||||||
|
|
||||||
|
assert "LoRA Version does not exist."
|
||||||
|
|
||||||
|
def check_lora_ext(self, lora_file: str):
|
||||||
|
return lora_file.endswith(tuple(LORA_FILE_TYPES))
|
||||||
|
|
||||||
|
def get_lora_file_path(
|
||||||
|
self,
|
||||||
|
lora_path: str,
|
||||||
|
model: Union[UNet3DConditionModel, CLIPTextModel]
|
||||||
|
):
|
||||||
|
if os.path.exists(lora_path):
|
||||||
|
lora_filenames = [fns for fns in os.listdir(lora_path)]
|
||||||
|
is_lora = self.check_lora_ext(lora_path)
|
||||||
|
|
||||||
|
is_unet = isinstance(model, UNet3DConditionModel)
|
||||||
|
is_text = isinstance(model, CLIPTextModel)
|
||||||
|
idx = 0 if is_unet else 1
|
||||||
|
|
||||||
|
base_name = FILE_BASENAMES[idx]
|
||||||
|
|
||||||
|
for lora_filename in lora_filenames:
|
||||||
|
is_lora = self.check_lora_ext(lora_filename)
|
||||||
|
if not is_lora:
|
||||||
|
continue
|
||||||
|
|
||||||
|
if base_name in lora_filename:
|
||||||
|
return os.path.join(lora_path, lora_filename)
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
def handle_lora_load(self, file_name:str, lora_loader_args: dict = None):
|
||||||
|
self.lora_loader(**lora_loader_args)
|
||||||
|
print(f"Successfully loaded LoRA from: {file_name}")
|
||||||
|
|
||||||
|
def load_lora(self, model, lora_path: str = '', lora_loader_args: dict = None, *args, **kwargs):
|
||||||
|
try:
|
||||||
|
lora_file = self.get_lora_file_path(lora_path, model)
|
||||||
|
|
||||||
|
if lora_file is not None:
|
||||||
|
lora_loader_args.update({"lora_path": lora_file})
|
||||||
|
self.handle_lora_load(lora_file, lora_loader_args)
|
||||||
|
|
||||||
|
else:
|
||||||
|
print(f"Could not load LoRAs for {model.__class__.__name__}. Injecting new ones instead...")
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
print(f"An error occured while loading a LoRA file: {e}")
|
||||||
|
|
||||||
|
def get_lora_func_args(
|
||||||
|
self,
|
||||||
|
lora_path,
|
||||||
|
use_lora,
|
||||||
|
model,
|
||||||
|
replace_modules,
|
||||||
|
r,
|
||||||
|
dropout,
|
||||||
|
lora_bias,
|
||||||
|
scale
|
||||||
|
):
|
||||||
|
return_dict = lora_args.copy()
|
||||||
|
|
||||||
|
if self.is_cloneofsimo_lora():
|
||||||
|
return_dict = filter_dict(return_dict, keys=CLONE_OF_SIMO_KEYS)
|
||||||
|
return_dict.update({
|
||||||
|
"model": model,
|
||||||
|
"loras": self.get_lora_file_path(lora_path, model),
|
||||||
|
"target_replace_module": replace_modules,
|
||||||
|
"r": r
|
||||||
|
})
|
||||||
|
|
||||||
|
if self.is_stable_lora():
|
||||||
|
KEYS = ['model', 'lora_path', 'scale']
|
||||||
|
return_dict = filter_dict(return_dict, KEYS)
|
||||||
|
|
||||||
|
return_dict.update({'model': model, 'lora_path': lora_path, 'scale': scale})
|
||||||
|
|
||||||
|
return return_dict
|
||||||
|
|
||||||
|
def do_lora_injection(
|
||||||
|
self,
|
||||||
|
model,
|
||||||
|
replace_modules,
|
||||||
|
bias='none',
|
||||||
|
dropout=0,
|
||||||
|
r=4,
|
||||||
|
scale=0,
|
||||||
|
lora_loader_args=None,
|
||||||
|
):
|
||||||
|
REPLACE_MODULES = replace_modules
|
||||||
|
|
||||||
|
params = None
|
||||||
|
negation = None
|
||||||
|
is_injection_hybrid = False
|
||||||
|
|
||||||
|
if self.is_cloneofsimo_lora():
|
||||||
|
is_injection_hybrid = True
|
||||||
|
injector_args = lora_loader_args
|
||||||
|
|
||||||
|
params, negation = self.lora_injector(**injector_args)
|
||||||
|
for _up, _down in extract_lora_ups_down(
|
||||||
|
model,
|
||||||
|
target_replace_module=REPLACE_MODULES):
|
||||||
|
|
||||||
|
if all(x is not None for x in [_up, _down]):
|
||||||
|
print(f"Lora successfully injected into {model.__class__.__name__}.")
|
||||||
|
|
||||||
|
break
|
||||||
|
|
||||||
|
return params, negation, is_injection_hybrid
|
||||||
|
|
||||||
|
if self.is_stable_lora():
|
||||||
|
injector_args = lora_args.copy()
|
||||||
|
injector_args = filter_dict(injector_args, keys=STABLE_LORA_KEYS)
|
||||||
|
|
||||||
|
SEARCH_CLASS = [torch.nn.Linear, torch.nn.Conv2d, torch.nn.Conv3d, torch.nn.Embedding]
|
||||||
|
|
||||||
|
injector_args.update({
|
||||||
|
"model": model,
|
||||||
|
"target_module": REPLACE_MODULES,
|
||||||
|
"search_class": SEARCH_CLASS,
|
||||||
|
"r": r,
|
||||||
|
"dropout": dropout,
|
||||||
|
"lora_bias": self.lora_bias,
|
||||||
|
"scale": scale
|
||||||
|
})
|
||||||
|
activator = self.lora_injector(**injector_args)
|
||||||
|
activator()
|
||||||
|
|
||||||
|
return params, negation, is_injection_hybrid
|
||||||
|
|
||||||
|
def add_lora_to_model(self, use_lora, model, replace_modules, dropout=0.0, lora_path='', r=16, scale=0):
|
||||||
|
|
||||||
|
params = None
|
||||||
|
negation = None
|
||||||
|
|
||||||
|
lora_loader_args = self.get_lora_func_args(
|
||||||
|
lora_path,
|
||||||
|
use_lora,
|
||||||
|
model,
|
||||||
|
replace_modules,
|
||||||
|
r,
|
||||||
|
dropout,
|
||||||
|
self.lora_bias,
|
||||||
|
scale
|
||||||
|
)
|
||||||
|
|
||||||
|
if use_lora:
|
||||||
|
params, negation, is_injection_hybrid = self.do_lora_injection(
|
||||||
|
model,
|
||||||
|
replace_modules,
|
||||||
|
bias=self.lora_bias,
|
||||||
|
lora_loader_args=lora_loader_args,
|
||||||
|
dropout=dropout,
|
||||||
|
r=r,
|
||||||
|
scale=scale
|
||||||
|
)
|
||||||
|
|
||||||
|
if not is_injection_hybrid:
|
||||||
|
self.load_lora(model, lora_path=lora_path, lora_loader_args=lora_loader_args)
|
||||||
|
|
||||||
|
params = model if params is None else params
|
||||||
|
return params, negation
|
||||||
|
|
||||||
|
|
||||||
|
def deactivate_lora_train(self, models, deactivate=True):
|
||||||
|
"""
|
||||||
|
Usage: Use before and after sampling previews.
|
||||||
|
Currently only available for Stable LoRA.
|
||||||
|
"""
|
||||||
|
if self.is_stable_lora():
|
||||||
|
set_mode_group(models, not deactivate)
|
||||||
|
|
||||||
|
def save_cloneofsimo_lora(
|
||||||
|
self,
|
||||||
|
model,
|
||||||
|
save_path,
|
||||||
|
step,
|
||||||
|
use_safetensors=True,
|
||||||
|
lora_rank="",
|
||||||
|
lora_name="",
|
||||||
|
use_motion_lora_format=False
|
||||||
|
):
|
||||||
|
|
||||||
|
# Same arguments as top level method
|
||||||
|
def save_lora(
|
||||||
|
model,
|
||||||
|
name,
|
||||||
|
condition,
|
||||||
|
replace_modules,
|
||||||
|
step,
|
||||||
|
save_path,
|
||||||
|
use_safetensors=True,
|
||||||
|
lora_rank="",
|
||||||
|
use_motion_lora_format=False
|
||||||
|
):
|
||||||
|
if condition and replace_modules is not None:
|
||||||
|
|
||||||
|
save_path = f"{save_path}/{step}_{name}"
|
||||||
|
|
||||||
|
if not use_safetensors:
|
||||||
|
save_lora_weight(
|
||||||
|
model,
|
||||||
|
save_path + ".pt",
|
||||||
|
replace_modules,
|
||||||
|
self.lora_r,
|
||||||
|
use_motion_lora_format=use_motion_lora_format
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
save_lora_safetensors(
|
||||||
|
model,
|
||||||
|
save_path + ".safetensors",
|
||||||
|
target_replace_module=replace_modules,
|
||||||
|
lora_rank=lora_rank,
|
||||||
|
use_motion_lora_format=use_motion_lora_format
|
||||||
|
)
|
||||||
|
|
||||||
|
save_lora(
|
||||||
|
model.unet,
|
||||||
|
f"{lora_name}_{FILE_BASENAMES[0]}",
|
||||||
|
self.use_unet_lora,
|
||||||
|
self.unet_replace_modules,
|
||||||
|
step,
|
||||||
|
save_path,
|
||||||
|
use_safetensors,
|
||||||
|
lora_rank,
|
||||||
|
use_motion_lora_format
|
||||||
|
)
|
||||||
|
save_lora(
|
||||||
|
model.text_encoder,
|
||||||
|
f"{lora_name}_{FILE_BASENAMES[1]}",
|
||||||
|
self.use_text_lora,
|
||||||
|
self.text_encoder_replace_modules,
|
||||||
|
step,
|
||||||
|
save_path,
|
||||||
|
use_safetensors,
|
||||||
|
lora_rank,
|
||||||
|
use_motion_lora_format
|
||||||
|
)
|
||||||
|
|
||||||
|
train_patch_pipe(model, self.use_unet_lora, self.use_text_lora)
|
||||||
|
|
||||||
|
def save_stable_lora(
|
||||||
|
self,
|
||||||
|
model,
|
||||||
|
step,
|
||||||
|
name,
|
||||||
|
save_path = '',
|
||||||
|
save_for_webui=False,
|
||||||
|
only_for_webui=False
|
||||||
|
):
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
save_filename = f"{step}_{name}"
|
||||||
|
lora_metadata = metadata = {
|
||||||
|
"stable_lora_text_to_video": "v1",
|
||||||
|
"lora_name": name + "_" + uuid.uuid4().hex.lower()[:5]
|
||||||
|
}
|
||||||
|
save_lora(
|
||||||
|
unet=model.unet,
|
||||||
|
text_encoder=model.text_encoder,
|
||||||
|
save_text_weights=self.use_text_lora,
|
||||||
|
output_dir=save_path,
|
||||||
|
lora_filename=save_filename,
|
||||||
|
lora_bias=self.lora_bias,
|
||||||
|
save_for_webui=self.save_for_webui,
|
||||||
|
only_webui=self.only_for_webui,
|
||||||
|
metadata=lora_metadata,
|
||||||
|
unet_dict_converter=convert_unet_state_dict,
|
||||||
|
text_dict_converter=convert_text_enc_state_dict_v20
|
||||||
|
)
|
||||||
|
|
||||||
|
def save_lora_weights(
|
||||||
|
self,
|
||||||
|
model: None,
|
||||||
|
save_path: str ='',
|
||||||
|
step: str = '',
|
||||||
|
use_safetensors: bool = True,
|
||||||
|
lora_rank="Not Logged",
|
||||||
|
lora_name="",
|
||||||
|
use_motion_lora_format=False
|
||||||
|
):
|
||||||
|
save_path = f"{save_path}"
|
||||||
|
os.makedirs(save_path, exist_ok=True)
|
||||||
|
|
||||||
|
if self.is_cloneofsimo_lora():
|
||||||
|
if any([self.save_for_webui, self.only_for_webui]):
|
||||||
|
warnings.warn(
|
||||||
|
"""
|
||||||
|
You have 'save_for_webui' enabled, but are using cloneofsimo's LoRA implemention.
|
||||||
|
Only 'stable_lora' is supported for saving to a compatible webui file.
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
self.save_cloneofsimo_lora(
|
||||||
|
model,
|
||||||
|
save_path,
|
||||||
|
step,
|
||||||
|
use_safetensors=use_safetensors,
|
||||||
|
lora_rank=lora_rank,
|
||||||
|
lora_name=lora_name,
|
||||||
|
use_motion_lora_format=use_motion_lora_format
|
||||||
|
)
|
||||||
|
|
||||||
|
if self.is_stable_lora():
|
||||||
|
name = 'lora_text_to_video'
|
||||||
|
self.save_stable_lora(model, step, name, save_path)
|
||||||
|
|
||||||
@@ -0,0 +1,172 @@
|
|||||||
|
import os
|
||||||
|
import imageio
|
||||||
|
import numpy as np
|
||||||
|
from typing import Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torchvision
|
||||||
|
import torch.distributed as dist
|
||||||
|
|
||||||
|
from safetensors import safe_open
|
||||||
|
from tqdm import tqdm
|
||||||
|
from einops import rearrange
|
||||||
|
from ...animatediff.utils.convert_from_ckpt import convert_ldm_unet_checkpoint, convert_ldm_clip_checkpoint, convert_ldm_vae_checkpoint
|
||||||
|
from ...animatediff.utils.convert_lora_safetensor_to_diffusers import convert_lora, load_diffusers_lora
|
||||||
|
|
||||||
|
|
||||||
|
def zero_rank_print(s):
|
||||||
|
if (not dist.is_initialized()) and (dist.is_initialized() and dist.get_rank() == 0): print("### " + s)
|
||||||
|
|
||||||
|
|
||||||
|
def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=6, fps=8):
|
||||||
|
videos = rearrange(videos, "b c t h w -> t b c h w")
|
||||||
|
outputs = []
|
||||||
|
for x in videos:
|
||||||
|
x = torchvision.utils.make_grid(x, nrow=n_rows)
|
||||||
|
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||||
|
if rescale:
|
||||||
|
x = (x + 1.0) / 2.0 # -1,1 -> 0,1
|
||||||
|
x = (x * 255).numpy().astype(np.uint8)
|
||||||
|
outputs.append(x)
|
||||||
|
|
||||||
|
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||||
|
imageio.mimsave(path, outputs, fps=fps)
|
||||||
|
|
||||||
|
|
||||||
|
# DDIM Inversion
|
||||||
|
@torch.no_grad()
|
||||||
|
def init_prompt(prompt, pipeline):
|
||||||
|
uncond_input = pipeline.tokenizer(
|
||||||
|
[""], padding="max_length", max_length=pipeline.tokenizer.model_max_length,
|
||||||
|
return_tensors="pt"
|
||||||
|
)
|
||||||
|
uncond_embeddings = pipeline.text_encoder(uncond_input.input_ids.to(pipeline.device))[0]
|
||||||
|
text_input = pipeline.tokenizer(
|
||||||
|
[prompt],
|
||||||
|
padding="max_length",
|
||||||
|
max_length=pipeline.tokenizer.model_max_length,
|
||||||
|
truncation=True,
|
||||||
|
return_tensors="pt",
|
||||||
|
)
|
||||||
|
text_embeddings = pipeline.text_encoder(text_input.input_ids.to(pipeline.device))[0]
|
||||||
|
context = torch.cat([uncond_embeddings, text_embeddings])
|
||||||
|
|
||||||
|
return context
|
||||||
|
|
||||||
|
|
||||||
|
def next_step(model_output: Union[torch.FloatTensor, np.ndarray], timestep: int,
|
||||||
|
sample: Union[torch.FloatTensor, np.ndarray], ddim_scheduler):
|
||||||
|
timestep, next_timestep = min(
|
||||||
|
timestep - ddim_scheduler.config.num_train_timesteps // ddim_scheduler.num_inference_steps, 999), timestep
|
||||||
|
alpha_prod_t = ddim_scheduler.alphas_cumprod[timestep] if timestep >= 0 else ddim_scheduler.final_alpha_cumprod
|
||||||
|
alpha_prod_t_next = ddim_scheduler.alphas_cumprod[next_timestep]
|
||||||
|
beta_prod_t = 1 - alpha_prod_t
|
||||||
|
next_original_sample = (sample - beta_prod_t ** 0.5 * model_output) / alpha_prod_t ** 0.5
|
||||||
|
next_sample_direction = (1 - alpha_prod_t_next) ** 0.5 * model_output
|
||||||
|
next_sample = alpha_prod_t_next ** 0.5 * next_original_sample + next_sample_direction
|
||||||
|
return next_sample
|
||||||
|
|
||||||
|
|
||||||
|
def get_noise_pred_single(latents, t, context, unet):
|
||||||
|
noise_pred = unet(latents, t, encoder_hidden_states=context)["sample"]
|
||||||
|
return noise_pred
|
||||||
|
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def ddim_loop(pipeline, ddim_scheduler, latent, num_inv_steps, prompt):
|
||||||
|
context = init_prompt(prompt, pipeline)
|
||||||
|
uncond_embeddings, cond_embeddings = context.chunk(2)
|
||||||
|
all_latent = [latent]
|
||||||
|
latent = latent.clone().detach()
|
||||||
|
for i in tqdm(range(num_inv_steps)):
|
||||||
|
t = ddim_scheduler.timesteps[len(ddim_scheduler.timesteps) - i - 1]
|
||||||
|
noise_pred = get_noise_pred_single(latent, t, cond_embeddings, pipeline.unet)
|
||||||
|
latent = next_step(noise_pred, t, latent, ddim_scheduler)
|
||||||
|
all_latent.append(latent)
|
||||||
|
return all_latent
|
||||||
|
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def ddim_inversion(pipeline, ddim_scheduler, video_latent, num_inv_steps, prompt=""):
|
||||||
|
ddim_latents = ddim_loop(pipeline, ddim_scheduler, video_latent, num_inv_steps, prompt)
|
||||||
|
return ddim_latents
|
||||||
|
|
||||||
|
def load_weights(
|
||||||
|
animation_pipeline,
|
||||||
|
# motion module
|
||||||
|
motion_module_path = "",
|
||||||
|
motion_module_lora_configs = [],
|
||||||
|
# domain adapter
|
||||||
|
adapter_lora_path = "",
|
||||||
|
adapter_lora_scale = 1.0,
|
||||||
|
# image layers
|
||||||
|
dreambooth_model_path = "",
|
||||||
|
lora_model_path = "",
|
||||||
|
lora_alpha = 0.8,
|
||||||
|
):
|
||||||
|
# motion module
|
||||||
|
unet_state_dict = {}
|
||||||
|
if motion_module_path != "":
|
||||||
|
print(f"load motion module from {motion_module_path}")
|
||||||
|
motion_module_state_dict = torch.load(motion_module_path, map_location="cpu")
|
||||||
|
motion_module_state_dict = motion_module_state_dict["state_dict"] if "state_dict" in motion_module_state_dict else motion_module_state_dict
|
||||||
|
unet_state_dict.update({name: param for name, param in motion_module_state_dict.items() if "motion_modules." in name})
|
||||||
|
unet_state_dict.pop("animatediff_config", "")
|
||||||
|
|
||||||
|
missing, unexpected = animation_pipeline.unet.load_state_dict(unet_state_dict, strict=False)
|
||||||
|
assert len(unexpected) == 0
|
||||||
|
del unet_state_dict
|
||||||
|
|
||||||
|
# base model
|
||||||
|
if dreambooth_model_path != "":
|
||||||
|
print(f"load dreambooth model from {dreambooth_model_path}")
|
||||||
|
if dreambooth_model_path.endswith(".safetensors"):
|
||||||
|
dreambooth_state_dict = {}
|
||||||
|
with safe_open(dreambooth_model_path, framework="pt", device="cpu") as f:
|
||||||
|
for key in f.keys():
|
||||||
|
dreambooth_state_dict[key] = f.get_tensor(key)
|
||||||
|
elif dreambooth_model_path.endswith(".ckpt"):
|
||||||
|
dreambooth_state_dict = torch.load(dreambooth_model_path, map_location="cpu")
|
||||||
|
|
||||||
|
# 1. vae
|
||||||
|
converted_vae_checkpoint = convert_ldm_vae_checkpoint(dreambooth_state_dict, animation_pipeline.vae.config)
|
||||||
|
animation_pipeline.vae.load_state_dict(converted_vae_checkpoint)
|
||||||
|
# 2. unet
|
||||||
|
converted_unet_checkpoint = convert_ldm_unet_checkpoint(dreambooth_state_dict, animation_pipeline.unet.config)
|
||||||
|
animation_pipeline.unet.load_state_dict(converted_unet_checkpoint, strict=False)
|
||||||
|
# 3. text_model
|
||||||
|
animation_pipeline.text_encoder = convert_ldm_clip_checkpoint(dreambooth_state_dict)
|
||||||
|
del dreambooth_state_dict
|
||||||
|
|
||||||
|
# lora layers
|
||||||
|
if lora_model_path != "":
|
||||||
|
print(f"load lora model from {lora_model_path}")
|
||||||
|
assert lora_model_path.endswith(".safetensors")
|
||||||
|
lora_state_dict = {}
|
||||||
|
with safe_open(lora_model_path, framework="pt", device="cpu") as f:
|
||||||
|
for key in f.keys():
|
||||||
|
lora_state_dict[key] = f.get_tensor(key)
|
||||||
|
|
||||||
|
animation_pipeline = convert_lora(animation_pipeline, lora_state_dict, alpha=lora_alpha)
|
||||||
|
del lora_state_dict
|
||||||
|
|
||||||
|
# domain adapter lora
|
||||||
|
if adapter_lora_path != "":
|
||||||
|
print(f"load domain lora from {adapter_lora_path}")
|
||||||
|
domain_lora_state_dict = torch.load(adapter_lora_path, map_location="cpu")
|
||||||
|
domain_lora_state_dict = domain_lora_state_dict["state_dict"] if "state_dict" in domain_lora_state_dict else domain_lora_state_dict
|
||||||
|
domain_lora_state_dict.pop("animatediff_config", "")
|
||||||
|
|
||||||
|
animation_pipeline = load_diffusers_lora(animation_pipeline, domain_lora_state_dict, alpha=adapter_lora_scale)
|
||||||
|
|
||||||
|
# motion module lora
|
||||||
|
for motion_module_lora_config in motion_module_lora_configs:
|
||||||
|
path, alpha = motion_module_lora_config["path"], motion_module_lora_config["alpha"]
|
||||||
|
print(f"load motion LoRA from {path}")
|
||||||
|
motion_lora_state_dict = torch.load(path, map_location="cpu")
|
||||||
|
motion_lora_state_dict = motion_lora_state_dict["state_dict"] if "state_dict" in motion_lora_state_dict else motion_lora_state_dict
|
||||||
|
motion_lora_state_dict.pop("animatediff_config", "")
|
||||||
|
|
||||||
|
animation_pipeline = load_diffusers_lora(animation_pipeline, motion_lora_state_dict, alpha)
|
||||||
|
|
||||||
|
return animation_pipeline
|
||||||
@@ -0,0 +1,55 @@
|
|||||||
|
# Model From Huggingface Diffusers
|
||||||
|
pretrained_model_path: "diffusers/stable-diffusion-v1-5"
|
||||||
|
|
||||||
|
# Model in CKPT format. This is the CKPT that you download from CivitAI or use in A111 Comfy, etc.
|
||||||
|
# In most cases, leave this blank.
|
||||||
|
unet_checkpoint_path: ""
|
||||||
|
|
||||||
|
# Must be CKPT from https://huggingface.co/guoyww/animatediff/tree/main
|
||||||
|
motion_module_path: "v3_sd15_mm.ckpt"
|
||||||
|
|
||||||
|
# Must be CKPT from https://huggingface.co/guoyww/animatediff/tree/main
|
||||||
|
# Optional for training, but highly recommended as a starting point.
|
||||||
|
domain_adapter_path: "" #"v3_sd15_adapter.ckpt"
|
||||||
|
|
||||||
|
# ["single_video", "folder"]
|
||||||
|
# single_video = path/my_video.mp4
|
||||||
|
# folder = path/my_videos
|
||||||
|
|
||||||
|
# You can have .txt file with prompt in the same folder.
|
||||||
|
# Eg. path/my_videos/1.mp4 path/my_videos/1.txt
|
||||||
|
# Otherwise, every video will have the same training prompt.
|
||||||
|
mode_type: "single_video"
|
||||||
|
|
||||||
|
video:
|
||||||
|
# Your local video path
|
||||||
|
path: "examples/pexels-cottonbro-5319934 (2160p).mp4" # Or just path/to/folder_of_videos_with_.txt_files/ (set mode_type to "folder")
|
||||||
|
|
||||||
|
# Optional custom start frame (idx). Leave this at 0 if your video is already trimmed.
|
||||||
|
# This is only recommended for single videos.
|
||||||
|
start_time: 0
|
||||||
|
|
||||||
|
# If your video is longer than 16 frames, it will be chunked into multiple parts.
|
||||||
|
# This is advanced usage, so leave this at 1.
|
||||||
|
max_chunks: 1
|
||||||
|
|
||||||
|
# Describe your action with a simple prompt.
|
||||||
|
training_prompt: "a man is running on a bridge"
|
||||||
|
|
||||||
|
# A custom prompt that will generate during to see your training progress.
|
||||||
|
validation_prompt: "a highly realistic video of batman running in a mystic forest, depth of field, epic lights, high quality, trending on artstation"
|
||||||
|
|
||||||
|
# The name of the LoRA file (will save to ./results/{save_name}...)
|
||||||
|
save_name: "man_running"
|
||||||
|
|
||||||
|
# Quality
|
||||||
|
# Quality of training: ["low", "preferred", "best"]
|
||||||
|
# low = Save the most memory, preferred = Most optimal, best = Memory intensive
|
||||||
|
quality: "preferred"
|
||||||
|
|
||||||
|
# Advanced users only. Only modify this if you're familiar with traning models.
|
||||||
|
# You can call / modify this config directly if you know what you're doing.
|
||||||
|
training_config: "configs/training/motion_director/training.yaml"
|
||||||
|
|
||||||
|
# Do not change this. Refer to the above for advanced training.
|
||||||
|
simple_mode: True
|
||||||
@@ -0,0 +1,109 @@
|
|||||||
|
image_finetune: false
|
||||||
|
output_dir: "outputs"
|
||||||
|
|
||||||
|
pretrained_model_path: ""
|
||||||
|
motion_module_path: ""
|
||||||
|
domain_adapter_path: ""
|
||||||
|
|
||||||
|
unet_additional_kwargs:
|
||||||
|
use_inflated_groupnorm: true
|
||||||
|
use_motion_module: true
|
||||||
|
motion_module_resolutions: [1,2,4,8]
|
||||||
|
motion_module_mid_block: false
|
||||||
|
motion_module_type: Vanilla
|
||||||
|
|
||||||
|
motion_module_kwargs:
|
||||||
|
num_attention_heads: 8
|
||||||
|
num_transformer_block: 1
|
||||||
|
attention_block_types: [ "Temporal_Self", "Temporal_Self" ]
|
||||||
|
temporal_position_encoding: true
|
||||||
|
temporal_position_encoding_max_len: 32
|
||||||
|
temporal_attention_dim_div: 1
|
||||||
|
zero_initialize: true
|
||||||
|
|
||||||
|
noise_scheduler_kwargs:
|
||||||
|
num_train_timesteps: 1000
|
||||||
|
beta_start: 0.00085
|
||||||
|
beta_end: 0.012
|
||||||
|
beta_schedule: "linear"
|
||||||
|
clip_sample: False
|
||||||
|
|
||||||
|
use_text_augmenter: False
|
||||||
|
dataset_types: ["single_video"]
|
||||||
|
|
||||||
|
cfg_random_null_text_ratio: 0
|
||||||
|
|
||||||
|
train_data:
|
||||||
|
manual_sample_size: False
|
||||||
|
sample_size: [320, 512]
|
||||||
|
# The width and height in which you want your training data to be resized to.
|
||||||
|
width: 256
|
||||||
|
height: 256
|
||||||
|
|
||||||
|
# This will find the closest aspect ratio to your input width and height.
|
||||||
|
# For example, 512x512 width and height with a video of resolution 1280x720 will be resized to 512x256
|
||||||
|
use_bucketing: False
|
||||||
|
|
||||||
|
# The start frame index where your videos should start (Leave this at one for json and folder based training).
|
||||||
|
sample_start_idx: 0
|
||||||
|
|
||||||
|
# Used for 'folder'. The rate at which your frames are sampled.
|
||||||
|
fps: 0
|
||||||
|
|
||||||
|
# For 'single_video' and 'json'. The number of frames to "step" (1,2,3,4) (frame_step=2) -> (1,3,5,7, ...).
|
||||||
|
frame_step: 3
|
||||||
|
|
||||||
|
# The number of frames to sample. The higher this number, the higher the VRAM (acts similar to batch size).
|
||||||
|
n_sample_frames: 16
|
||||||
|
|
||||||
|
# For validation
|
||||||
|
sample_n_frames: 16
|
||||||
|
|
||||||
|
# 'single_video'
|
||||||
|
single_video_path: ""
|
||||||
|
|
||||||
|
# The prompt when using a single video file
|
||||||
|
single_video_prompt: ""
|
||||||
|
|
||||||
|
fallback_prompt: ""
|
||||||
|
|
||||||
|
max_chunks: 1
|
||||||
|
|
||||||
|
validation_data:
|
||||||
|
prompts:
|
||||||
|
- ""
|
||||||
|
num_inference_steps: 25
|
||||||
|
guidance_scale: 9
|
||||||
|
spatial_scale: 0.5
|
||||||
|
validation_seed: 44
|
||||||
|
|
||||||
|
lora_name: ""
|
||||||
|
use_motion_lora_format: True
|
||||||
|
lora_rank: 32
|
||||||
|
lora_unet_dropout: 0.1
|
||||||
|
single_spatial_lora: True
|
||||||
|
train_sample_validation: False
|
||||||
|
unet_checkpoint_path: ""
|
||||||
|
|
||||||
|
learning_rate: 5e-4
|
||||||
|
learning_rate_spatial: 1e-4
|
||||||
|
adam_weight_decay: 1e-2
|
||||||
|
cache_latents: true
|
||||||
|
train_batch_size: 1
|
||||||
|
use_lion_optim: True
|
||||||
|
use_offset_noise: False
|
||||||
|
|
||||||
|
max_train_epoch: 503
|
||||||
|
max_train_steps: -1
|
||||||
|
checkpointing_epochs: -1
|
||||||
|
checkpointing_steps: 100
|
||||||
|
|
||||||
|
validation_steps: 50
|
||||||
|
validation_steps_tuple: [2, 50]
|
||||||
|
|
||||||
|
global_seed: 33
|
||||||
|
mixed_precision_training: true
|
||||||
|
enable_xformers_memory_efficient_attention: True
|
||||||
|
gradient_checkpointing: True
|
||||||
|
|
||||||
|
is_debug: False
|
||||||
Binary file not shown.
@@ -0,0 +1,900 @@
|
|||||||
|
import os
|
||||||
|
import math
|
||||||
|
import random
|
||||||
|
import logging
|
||||||
|
import inspect
|
||||||
|
import datetime
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
from tqdm.auto import tqdm
|
||||||
|
from einops import rearrange
|
||||||
|
from omegaconf import OmegaConf
|
||||||
|
#from typing import Dict, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torchvision
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
#import diffusers
|
||||||
|
from diffusers import AutoencoderKL, DDIMScheduler, DDPMScheduler
|
||||||
|
#from diffusers.models import UNet2DConditionModel
|
||||||
|
#from diffusers.pipelines import StableDiffusionPipeline
|
||||||
|
from diffusers.optimization import get_scheduler
|
||||||
|
#from diffusers.utils import check_min_version
|
||||||
|
|
||||||
|
#import transformers
|
||||||
|
from transformers import CLIPTextModel, CLIPTokenizer
|
||||||
|
|
||||||
|
from .animatediff.models.unet import UNet3DConditionModel
|
||||||
|
from .animatediff.pipelines.pipeline_animation import AnimationPipeline
|
||||||
|
from .animatediff.utils.util import save_videos_grid, load_diffusers_lora, load_weights
|
||||||
|
from .animatediff.utils.lora_handler import LoraHandler
|
||||||
|
from .animatediff.utils.lora import extract_lora_child_module
|
||||||
|
|
||||||
|
#from .animatediff.utils.configs import get_simple_config
|
||||||
|
from lion_pytorch import Lion
|
||||||
|
import comfy.model_management
|
||||||
|
import comfy.utils
|
||||||
|
|
||||||
|
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
|
||||||
|
augment_text_list = [
|
||||||
|
"a video of",
|
||||||
|
"a high quality video of",
|
||||||
|
"a good video of",
|
||||||
|
"a nice video of",
|
||||||
|
"a great video of",
|
||||||
|
"a video showing",
|
||||||
|
"video of",
|
||||||
|
"video clip of",
|
||||||
|
"great video of",
|
||||||
|
"cool video of",
|
||||||
|
"best video of",
|
||||||
|
"streamed video of",
|
||||||
|
"excellent video of",
|
||||||
|
"new video of",
|
||||||
|
"new video clip of",
|
||||||
|
"high quality video of",
|
||||||
|
"a video showing of",
|
||||||
|
"a clear video showing",
|
||||||
|
"video clip showing",
|
||||||
|
"a clear video showing",
|
||||||
|
"a nice video showing",
|
||||||
|
"a good video showing",
|
||||||
|
"video, high quality,"
|
||||||
|
"high quality, video, video clip,",
|
||||||
|
"nice video, clear quality,",
|
||||||
|
"clear quality video of"
|
||||||
|
]
|
||||||
|
|
||||||
|
def create_save_paths(output_dir: str):
|
||||||
|
lora_path = f"{output_dir}/lora"
|
||||||
|
|
||||||
|
directories = [
|
||||||
|
output_dir,
|
||||||
|
f"{output_dir}/samples",
|
||||||
|
f"{output_dir}/sanity_check",
|
||||||
|
lora_path
|
||||||
|
]
|
||||||
|
|
||||||
|
for directory in directories:
|
||||||
|
os.makedirs(directory, exist_ok=True)
|
||||||
|
|
||||||
|
return lora_path
|
||||||
|
|
||||||
|
def do_sanity_check(
|
||||||
|
pixel_values: torch.Tensor,
|
||||||
|
cache_latents: bool,
|
||||||
|
validation_pipeline: AnimationPipeline,
|
||||||
|
device: str,
|
||||||
|
image_finetune: bool=False,
|
||||||
|
output_dir: str = "",
|
||||||
|
text_prompt: str = ""
|
||||||
|
):
|
||||||
|
pixel_values, texts = pixel_values.cpu(), text_prompt
|
||||||
|
|
||||||
|
if cache_latents:
|
||||||
|
pixel_values = validation_pipeline.decode_latents(pixel_values.to(device))
|
||||||
|
to_torch = torch.from_numpy(pixel_values)
|
||||||
|
pixel_values = rearrange(to_torch, 'b c f h w -> b f c h w')
|
||||||
|
|
||||||
|
if not image_finetune:
|
||||||
|
pixel_values = rearrange(pixel_values, "b f c h w -> b c f h w")
|
||||||
|
for idx, (pixel_value, text) in enumerate(zip(pixel_values, texts)):
|
||||||
|
pixel_value = pixel_value[None, ...]
|
||||||
|
text = text
|
||||||
|
save_name = f"{'-'.join(text.replace('/', '').split()[:10]) if not text == '' else f'-{idx}'}.mp4"
|
||||||
|
save_videos_grid(pixel_value, f"{output_dir}/sanity_check/{save_name}", rescale=not cache_latents)
|
||||||
|
else:
|
||||||
|
for idx, (pixel_value, text) in enumerate(zip(pixel_values, texts)):
|
||||||
|
pixel_value = pixel_value / 2. + 0.5
|
||||||
|
text = text
|
||||||
|
save_name = f"{'-'.join(text.replace('/', '').split()[:10]) if not text == '' else f'-{idx}'}.png"
|
||||||
|
torchvision.utils.save_image(pixel_value, f"{output_dir}/sanity_check/{save_name}")
|
||||||
|
|
||||||
|
def sample_noise(latents, noise_strength, use_offset_noise=False):
|
||||||
|
b, c, f, *_ = latents.shape
|
||||||
|
noise_latents = torch.randn_like(latents, device=latents.device)
|
||||||
|
|
||||||
|
if use_offset_noise:
|
||||||
|
offset_noise = torch.randn(b, c, f, 1, 1, device=latents.device)
|
||||||
|
noise_latents = noise_latents + noise_strength * offset_noise
|
||||||
|
|
||||||
|
return noise_latents
|
||||||
|
|
||||||
|
def param_optim(model, condition, extra_params=None, is_lora=False, negation=None):
|
||||||
|
extra_params = extra_params if len(extra_params.keys()) > 0 else None
|
||||||
|
return {
|
||||||
|
"model": model,
|
||||||
|
"condition": condition,
|
||||||
|
'extra_params': extra_params,
|
||||||
|
'is_lora': is_lora,
|
||||||
|
"negation": negation
|
||||||
|
}
|
||||||
|
|
||||||
|
def create_optim_params(name='param', params=None, lr=5e-6, extra_params=None):
|
||||||
|
params = {
|
||||||
|
"name": name,
|
||||||
|
"params": params,
|
||||||
|
"lr": lr
|
||||||
|
}
|
||||||
|
if extra_params is not None:
|
||||||
|
for k, v in extra_params.items():
|
||||||
|
params[k] = v
|
||||||
|
|
||||||
|
return params
|
||||||
|
|
||||||
|
def create_optimizer_params(model_list, lr):
|
||||||
|
import itertools
|
||||||
|
optimizer_params = []
|
||||||
|
|
||||||
|
for optim in model_list:
|
||||||
|
model, condition, extra_params, is_lora, negation = optim.values()
|
||||||
|
# Check if we are doing LoRA training.
|
||||||
|
if is_lora and condition and isinstance(model, list):
|
||||||
|
params = create_optim_params(
|
||||||
|
params=itertools.chain(*model),
|
||||||
|
extra_params=extra_params
|
||||||
|
)
|
||||||
|
optimizer_params.append(params)
|
||||||
|
continue
|
||||||
|
|
||||||
|
if is_lora and condition and not isinstance(model, list):
|
||||||
|
for n, p in model.named_parameters():
|
||||||
|
if 'lora' in n:
|
||||||
|
params = create_optim_params(n, p, lr, extra_params)
|
||||||
|
optimizer_params.append(params)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# If this is true, we can train it.
|
||||||
|
if condition:
|
||||||
|
for n, p in model.named_parameters():
|
||||||
|
should_negate = 'lora' in n and not is_lora
|
||||||
|
if should_negate: continue
|
||||||
|
|
||||||
|
params = create_optim_params(n, p, lr, extra_params)
|
||||||
|
optimizer_params.append(params)
|
||||||
|
|
||||||
|
return optimizer_params
|
||||||
|
|
||||||
|
def scale_loras(lora_list: list, scale: float, step=None, spatial_lora_num=None):
|
||||||
|
|
||||||
|
# Assumed enumerator
|
||||||
|
if step is not None and spatial_lora_num is not None:
|
||||||
|
process_list = range(0, len(lora_list), spatial_lora_num)
|
||||||
|
else:
|
||||||
|
process_list = lora_list
|
||||||
|
|
||||||
|
for lora_i in process_list:
|
||||||
|
if step is not None:
|
||||||
|
lora_list[lora_i].scale = scale
|
||||||
|
else:
|
||||||
|
lora_i.scale = scale
|
||||||
|
|
||||||
|
def tensor_to_vae_latent(t, vae):
|
||||||
|
video_length = t.shape[1]
|
||||||
|
|
||||||
|
t = rearrange(t, "b f c h w -> (b f) c h w")
|
||||||
|
latents = vae.encode(t).latent_dist.sample()
|
||||||
|
latents = rearrange(latents, "(b f) c h w -> b c f h w", f=video_length)
|
||||||
|
latents = latents * 0.18215
|
||||||
|
|
||||||
|
return latents
|
||||||
|
|
||||||
|
def get_spatial_latents(
|
||||||
|
pixel_values: torch.Tensor,
|
||||||
|
random_hflip_img: int,
|
||||||
|
cache_latents: bool,
|
||||||
|
noisy_latents:torch.Tensor,
|
||||||
|
target: torch.Tensor,
|
||||||
|
timesteps: torch.Tensor,
|
||||||
|
noise_scheduler: DDPMScheduler
|
||||||
|
):
|
||||||
|
ran_idx = torch.randint(0, pixel_values.shape[2], (1,)).item()
|
||||||
|
use_hflip = random.uniform(0, 1) < random_hflip_img
|
||||||
|
|
||||||
|
noisy_latents_input = None
|
||||||
|
target_spatial = None
|
||||||
|
|
||||||
|
if use_hflip:
|
||||||
|
pixel_values_spatial = torchvision.transforms.functional.hflip(
|
||||||
|
pixel_values[:, ran_idx, :, :, :] if not cache_latents else\
|
||||||
|
pixel_values[:, :, ran_idx, :, :]
|
||||||
|
).unsqueeze(1)
|
||||||
|
|
||||||
|
latents_spatial = (
|
||||||
|
tensor_to_vae_latent(pixel_values_spatial, vae) if not cache_latents
|
||||||
|
else
|
||||||
|
pixel_values_spatial
|
||||||
|
)
|
||||||
|
|
||||||
|
noise_spatial = sample_noise(latents_spatial, 0, use_offset_noise=False)
|
||||||
|
noisy_latents_input = noise_scheduler.add_noise(latents_spatial, noise_spatial, timesteps)
|
||||||
|
|
||||||
|
target_spatial = noise_spatial
|
||||||
|
else:
|
||||||
|
noisy_latents_input = noisy_latents[:, :, ran_idx, :, :]
|
||||||
|
target_spatial = target[:, :, ran_idx, :, :]
|
||||||
|
|
||||||
|
return noisy_latents_input, target_spatial, use_hflip
|
||||||
|
|
||||||
|
def create_ad_temporal_loss(
|
||||||
|
model_pred: torch.Tensor,
|
||||||
|
loss_temporal: torch.Tensor,
|
||||||
|
target: torch.Tensor
|
||||||
|
):
|
||||||
|
|
||||||
|
beta = 1
|
||||||
|
alpha = (beta ** 2 + 1) ** 0.5
|
||||||
|
|
||||||
|
ran_idx = torch.randint(0, model_pred.shape[2], (1,)).item()
|
||||||
|
|
||||||
|
model_pred_decent = alpha * model_pred - beta * model_pred[:, :, ran_idx, :, :].unsqueeze(2)
|
||||||
|
target_decent = alpha * target - beta * target[:, :, ran_idx, :, :].unsqueeze(2)
|
||||||
|
|
||||||
|
loss_ad_temporal = F.mse_loss(model_pred_decent.float(), target_decent.float(), reduction="mean")
|
||||||
|
loss_temporal = loss_temporal + loss_ad_temporal
|
||||||
|
|
||||||
|
return loss_temporal
|
||||||
|
|
||||||
|
class AD_MotionDirector_train:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"validation_models": ("VALIDATION_MODELS", ),
|
||||||
|
"unet": ("MODEL", ),
|
||||||
|
"clip": ("CLIP", ),
|
||||||
|
"tokenizer": ("TOKENIZER", ),
|
||||||
|
"vae": ("VAE", ),
|
||||||
|
"lora_name": ("STRING", {"multiline": False, "default": "motiondirectorlora",}),
|
||||||
|
"images": ("IMAGE", ),
|
||||||
|
"prompt": ("STRING", {"multiline": True, "default": "",}),
|
||||||
|
"validation_prompt": ("STRING", {"multiline": True, "default": "",}),
|
||||||
|
"max_train_epoch": ("INT", {"default": 300, "min": -1, "max": 10000, "step": 1}),
|
||||||
|
"max_train_steps": ("INT", {"default": -1, "min": -1, "max": 10000, "step": 1}),
|
||||||
|
"learning_rate": ("FLOAT", {"default": 5e-4, "min": 0, "max": 10000, "step": 0.00001}),
|
||||||
|
"learning_rate_spatial": ("FLOAT", {"default": 1e-4, "min": 0, "max": 10000, "step": 0.00001}),
|
||||||
|
"checkpointing_steps": ("INT", {"default": 100, "min": -1, "max": 10000, "step": 1}),
|
||||||
|
"checkpointing_epochs": ("INT", {"default": -1, "min": -1, "max": 10000, "step": 1}),
|
||||||
|
"lora_rank": ("INT", {"default": 32, "min": 8, "max": 4096, "step": 8}),
|
||||||
|
"use_xformers": ("BOOLEAN", {"default": False}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("IMAGE",)
|
||||||
|
RETURN_NAMES =("image",)
|
||||||
|
FUNCTION = "process"
|
||||||
|
|
||||||
|
CATEGORY = "AD_MotionDirector"
|
||||||
|
|
||||||
|
def process(self, validation_models, unet, clip, tokenizer, vae, images, prompt, validation_prompt,
|
||||||
|
lora_name, max_train_epoch, max_train_steps, learning_rate, learning_rate_spatial, checkpointing_steps, checkpointing_epochs, lora_rank, use_xformers):
|
||||||
|
with torch.inference_mode(False):
|
||||||
|
|
||||||
|
motion_module_path, domain_adapter_path, unet_checkpoint_path = validation_models
|
||||||
|
|
||||||
|
input_height, input_width = images.shape[1], images.shape[2]
|
||||||
|
images = images * 2.0 - 1.0 #normalize to the expected range (-1, 1)
|
||||||
|
pixel_values = images.clone().requires_grad_(True)
|
||||||
|
pixel_values = pixel_values.permute(0, 3, 1, 2).unsqueeze(0)#B,H,W,C to B,F,C,H,W
|
||||||
|
|
||||||
|
text_encoder = clip
|
||||||
|
text_prompt = []
|
||||||
|
text_prompt.append(prompt)
|
||||||
|
|
||||||
|
device = comfy.model_management.get_torch_device()
|
||||||
|
config = OmegaConf.load(os.path.join(script_directory, f"configs/training/motion_director/training.yaml"))
|
||||||
|
noise_scheduler_kwargs = config.noise_scheduler_kwargs
|
||||||
|
|
||||||
|
cfg_random_null_text = True
|
||||||
|
cfg_random_null_text_ratio = 0
|
||||||
|
|
||||||
|
scale_lr = False
|
||||||
|
lr_warmup_steps = 0
|
||||||
|
lr_scheduler = "constant"
|
||||||
|
|
||||||
|
train_batch_size = 1
|
||||||
|
adam_beta1 = 0.9
|
||||||
|
adam_beta2 = 0.999
|
||||||
|
adam_weight_decay = 1e-2
|
||||||
|
gradient_accumulation_steps = 1
|
||||||
|
gradient_checkpointing = True
|
||||||
|
|
||||||
|
mixed_precision_training = True
|
||||||
|
|
||||||
|
global_seed = 33
|
||||||
|
|
||||||
|
is_debug = False
|
||||||
|
|
||||||
|
random_hflip_img = -1
|
||||||
|
use_motion_lora_format = True
|
||||||
|
single_spatial_lora = True
|
||||||
|
lora_rank = 32
|
||||||
|
lora_unet_dropout = 0.1
|
||||||
|
train_temporal_lora = True
|
||||||
|
target_spatial_modules = ["Transformer3DModel"]
|
||||||
|
target_temporal_modules = ["TemporalTransformerBlock"]
|
||||||
|
|
||||||
|
cache_latents = False
|
||||||
|
|
||||||
|
train_sample_validation = False
|
||||||
|
use_text_augmenter = False
|
||||||
|
use_lion_optim = True
|
||||||
|
use_offset_noise = False
|
||||||
|
|
||||||
|
validation_spatial_scale = 0.5
|
||||||
|
validation_seed = 44
|
||||||
|
validation_steps = 25
|
||||||
|
validation_steps_tuple = [2, 25]
|
||||||
|
|
||||||
|
|
||||||
|
# Initialize distributed training
|
||||||
|
num_processes = 1
|
||||||
|
seed = global_seed
|
||||||
|
torch.manual_seed(seed)
|
||||||
|
|
||||||
|
name = lora_name
|
||||||
|
|
||||||
|
date_calendar = datetime.datetime.now().strftime("%Y-%m-%d")
|
||||||
|
date_time = datetime.datetime.now().strftime("-%H-%M-%S")
|
||||||
|
folder_name = "debug" if is_debug else name + date_time
|
||||||
|
|
||||||
|
output_dir = os.path.join(script_directory, "outputs", date_calendar, folder_name)
|
||||||
|
|
||||||
|
if is_debug and os.path.exists(output_dir):
|
||||||
|
os.system(f"rm -rf {output_dir}")
|
||||||
|
|
||||||
|
*_, config = inspect.getargvalues(inspect.currentframe())
|
||||||
|
|
||||||
|
# Make one log on every process with the configuration for debugging.
|
||||||
|
logging.basicConfig(
|
||||||
|
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
|
||||||
|
datefmt="%m/%d/%Y %H:%M:%S",
|
||||||
|
level=logging.INFO,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Handle the output folder creation
|
||||||
|
lora_path = create_save_paths(output_dir)
|
||||||
|
#OmegaConf.save(config, os.path.join(output_dir, 'config.yaml'))
|
||||||
|
|
||||||
|
# Load scheduler, tokenizer and models.
|
||||||
|
noise_scheduler_kwargs.update({"steps_offset": 1})
|
||||||
|
noise_scheduler = DDIMScheduler(**OmegaConf.to_container(noise_scheduler_kwargs))
|
||||||
|
del noise_scheduler_kwargs["steps_offset"]
|
||||||
|
|
||||||
|
noise_scheduler_kwargs['beta_schedule'] = 'scaled_linear'
|
||||||
|
train_noise_scheduler_spatial = DDPMScheduler(**OmegaConf.to_container(noise_scheduler_kwargs))
|
||||||
|
|
||||||
|
# AnimateDiff uses a linear schedule for its temporal sampling
|
||||||
|
noise_scheduler_kwargs['beta_schedule'] = 'linear'
|
||||||
|
train_noise_scheduler = DDPMScheduler(**OmegaConf.to_container(noise_scheduler_kwargs))
|
||||||
|
|
||||||
|
# Freeze all models for LoRA training
|
||||||
|
unet.requires_grad_(False)
|
||||||
|
vae.requires_grad_(False)
|
||||||
|
text_encoder.requires_grad_(False)
|
||||||
|
|
||||||
|
if not use_lion_optim:
|
||||||
|
optimizer = torch.optim.AdamW
|
||||||
|
else:
|
||||||
|
optimizer = Lion
|
||||||
|
learning_rate, learning_rate_spatial = map(lambda lr: lr / 10, (learning_rate, learning_rate_spatial))
|
||||||
|
adam_weight_decay *= 10
|
||||||
|
|
||||||
|
if use_xformers:
|
||||||
|
unet.enable_xformers_memory_efficient_attention()
|
||||||
|
|
||||||
|
# Enable gradient checkpointing
|
||||||
|
if gradient_checkpointing:
|
||||||
|
unet.enable_gradient_checkpointing()
|
||||||
|
|
||||||
|
# Move models to GPU
|
||||||
|
vae.to(device)
|
||||||
|
text_encoder.to(device)
|
||||||
|
|
||||||
|
# Get the training iteration
|
||||||
|
if max_train_steps == -1:
|
||||||
|
assert max_train_epoch != -1
|
||||||
|
max_train_steps = max_train_epoch
|
||||||
|
|
||||||
|
if checkpointing_steps == -1:
|
||||||
|
assert checkpointing_epochs != -1
|
||||||
|
checkpointing_steps = checkpointing_epochs
|
||||||
|
|
||||||
|
if scale_lr:
|
||||||
|
learning_rate = (learning_rate * gradient_accumulation_steps * train_batch_size * num_processes)
|
||||||
|
|
||||||
|
# Validation pipeline
|
||||||
|
validation_pipeline = AnimationPipeline(
|
||||||
|
unet=unet, vae=vae, tokenizer=tokenizer, text_encoder=text_encoder, scheduler=noise_scheduler,
|
||||||
|
).to(device)
|
||||||
|
|
||||||
|
validation_pipeline = load_weights(
|
||||||
|
validation_pipeline,
|
||||||
|
motion_module_path=motion_module_path,
|
||||||
|
adapter_lora_path=domain_adapter_path,
|
||||||
|
dreambooth_model_path=unet_checkpoint_path
|
||||||
|
)
|
||||||
|
|
||||||
|
validation_pipeline.enable_vae_slicing()
|
||||||
|
validation_pipeline.to(device)
|
||||||
|
|
||||||
|
unet.to(device=device)
|
||||||
|
text_encoder.to(device=device)
|
||||||
|
|
||||||
|
# Temporal LoRA
|
||||||
|
if train_temporal_lora:
|
||||||
|
# one temporal lora
|
||||||
|
lora_manager_temporal = LoraHandler(use_unet_lora=True, unet_replace_modules=target_temporal_modules)
|
||||||
|
|
||||||
|
unet_lora_params_temporal, unet_negation_temporal = lora_manager_temporal.add_lora_to_model(
|
||||||
|
True, unet, lora_manager_temporal.unet_replace_modules, 0,
|
||||||
|
lora_path + '/temporal/', r=lora_rank)
|
||||||
|
|
||||||
|
optimizer_temporal = optimizer(
|
||||||
|
create_optimizer_params([param_optim(unet_lora_params_temporal, True, is_lora=True,
|
||||||
|
extra_params={**{"lr": learning_rate}}
|
||||||
|
)], learning_rate),
|
||||||
|
lr=learning_rate,
|
||||||
|
betas=(adam_beta1, adam_beta2),
|
||||||
|
weight_decay=adam_weight_decay
|
||||||
|
)
|
||||||
|
|
||||||
|
lr_scheduler_temporal = get_scheduler(
|
||||||
|
lr_scheduler,
|
||||||
|
optimizer=optimizer_temporal,
|
||||||
|
num_warmup_steps=lr_warmup_steps * gradient_accumulation_steps,
|
||||||
|
num_training_steps=max_train_steps * gradient_accumulation_steps,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
lora_manager_temporal = None
|
||||||
|
unet_lora_params_temporal, unet_negation_temporal = [], []
|
||||||
|
optimizer_temporal = None
|
||||||
|
lr_scheduler_temporal = None
|
||||||
|
|
||||||
|
# Spatial LoRAs
|
||||||
|
if single_spatial_lora:
|
||||||
|
spatial_lora_num = 1
|
||||||
|
|
||||||
|
lora_managers_spatial = []
|
||||||
|
unet_lora_params_spatial_list = []
|
||||||
|
optimizer_spatial_list = []
|
||||||
|
lr_scheduler_spatial_list = []
|
||||||
|
|
||||||
|
for i in range(spatial_lora_num):
|
||||||
|
lora_manager_spatial = LoraHandler(use_unet_lora=True, unet_replace_modules=target_spatial_modules)
|
||||||
|
lora_managers_spatial.append(lora_manager_spatial)
|
||||||
|
unet_lora_params_spatial, unet_negation_spatial = lora_manager_spatial.add_lora_to_model(
|
||||||
|
True, unet, lora_manager_spatial.unet_replace_modules, lora_unet_dropout,
|
||||||
|
lora_path + '/spatial/', r=lora_rank)
|
||||||
|
|
||||||
|
unet_lora_params_spatial_list.append(unet_lora_params_spatial)
|
||||||
|
|
||||||
|
optimizer_spatial = optimizer(
|
||||||
|
create_optimizer_params([param_optim(unet_lora_params_spatial, True, is_lora=True,
|
||||||
|
extra_params={**{"lr": learning_rate_spatial}}
|
||||||
|
)], learning_rate_spatial),
|
||||||
|
lr=learning_rate_spatial,
|
||||||
|
betas=(adam_beta1, adam_beta2),
|
||||||
|
weight_decay=adam_weight_decay
|
||||||
|
)
|
||||||
|
|
||||||
|
optimizer_spatial_list.append(optimizer_spatial)
|
||||||
|
|
||||||
|
# Scheduler
|
||||||
|
lr_scheduler_spatial = get_scheduler(
|
||||||
|
lr_scheduler,
|
||||||
|
optimizer=optimizer_spatial,
|
||||||
|
num_warmup_steps=lr_warmup_steps * gradient_accumulation_steps,
|
||||||
|
num_training_steps=max_train_steps * gradient_accumulation_steps,
|
||||||
|
)
|
||||||
|
lr_scheduler_spatial_list.append(lr_scheduler_spatial)
|
||||||
|
|
||||||
|
unet_negation_all = unet_negation_spatial + unet_negation_temporal
|
||||||
|
|
||||||
|
# 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(1) / gradient_accumulation_steps
|
||||||
|
|
||||||
|
# Afterwards we recalculate our number of training epochs
|
||||||
|
num_train_epochs = math.ceil(max_train_steps / num_update_steps_per_epoch)
|
||||||
|
# Train!
|
||||||
|
|
||||||
|
global_step = 0
|
||||||
|
first_epoch = 0
|
||||||
|
#num_train_epochs = 300
|
||||||
|
batch_size = 1
|
||||||
|
# Only show the progress bar once on each machine.
|
||||||
|
progress_bar = tqdm(range(global_step, max_train_steps))
|
||||||
|
progress_bar.set_description("Steps")
|
||||||
|
|
||||||
|
# Support mixed-precision training
|
||||||
|
scaler = torch.cuda.amp.GradScaler() if mixed_precision_training else None
|
||||||
|
|
||||||
|
pbar = comfy.utils.ProgressBar(batch_size * num_train_epochs)
|
||||||
|
|
||||||
|
### <<<< Training <<<< ###
|
||||||
|
for epoch in range(first_epoch, num_train_epochs):
|
||||||
|
unet.train()
|
||||||
|
|
||||||
|
for step in range(batch_size):
|
||||||
|
spatial_scheduler_lr = 0.0
|
||||||
|
temporal_scheduler_lr = 0.0
|
||||||
|
|
||||||
|
# Handle Lora Optimizers & Conditions
|
||||||
|
for optimizer_spatial in optimizer_spatial_list:
|
||||||
|
optimizer_spatial.zero_grad(set_to_none=True)
|
||||||
|
|
||||||
|
if optimizer_temporal is not None:
|
||||||
|
optimizer_temporal.zero_grad(set_to_none=True)
|
||||||
|
|
||||||
|
if train_temporal_lora:
|
||||||
|
mask_temporal_lora = False
|
||||||
|
else:
|
||||||
|
mask_temporal_lora = True
|
||||||
|
|
||||||
|
mask_spatial_lora = random.uniform(0, 1) < 0.2 and not mask_temporal_lora
|
||||||
|
|
||||||
|
if cfg_random_null_text:
|
||||||
|
text_prompt = [name if random.random() > cfg_random_null_text_ratio else "" for name in text_prompt]
|
||||||
|
|
||||||
|
if use_text_augmenter:
|
||||||
|
random.seed()
|
||||||
|
txt_idx = random.randint(0, len(augment_text_list) - 1)
|
||||||
|
augment_text = augment_text_list[txt_idx]
|
||||||
|
|
||||||
|
text_prompt = [
|
||||||
|
f"{augment_text} {prompt}" for prompt in text_prompt
|
||||||
|
]
|
||||||
|
|
||||||
|
#Data batch sanity check
|
||||||
|
# if epoch == first_epoch and step == 0:
|
||||||
|
# "DO SANITY CHECK"
|
||||||
|
# do_sanity_check(
|
||||||
|
# pixel_values,
|
||||||
|
# cache_latents,
|
||||||
|
# validation_pipeline,
|
||||||
|
# device,
|
||||||
|
# output_dir=output_dir,
|
||||||
|
# text_prompt=text_prompt
|
||||||
|
# )
|
||||||
|
|
||||||
|
# Convert videos to latent space
|
||||||
|
|
||||||
|
#torch.Size([1, 4, 16, 32, 48])
|
||||||
|
pixel_values = pixel_values.to(device)
|
||||||
|
|
||||||
|
video_length = pixel_values.shape[2]
|
||||||
|
bsz = pixel_values.shape[0]
|
||||||
|
|
||||||
|
# Sample a random timestep for each video
|
||||||
|
timesteps = torch.randint(0, train_noise_scheduler.config.num_train_timesteps, (bsz,), device=pixel_values.device)
|
||||||
|
timesteps = timesteps.long()
|
||||||
|
|
||||||
|
# Add noise to the latents according to the noise magnitude at each timestep
|
||||||
|
# (this is the forward diffusion process)
|
||||||
|
latents = tensor_to_vae_latent(pixel_values, vae) if not cache_latents else pixel_values
|
||||||
|
noise = sample_noise(latents, 0, use_offset_noise=use_offset_noise)
|
||||||
|
target = noise
|
||||||
|
|
||||||
|
# Get the text embedding for conditioning
|
||||||
|
with torch.no_grad():
|
||||||
|
prompt_ids = tokenizer(
|
||||||
|
text_prompt,
|
||||||
|
max_length=tokenizer.model_max_length,
|
||||||
|
padding="max_length",
|
||||||
|
truncation=True,
|
||||||
|
return_tensors="pt"
|
||||||
|
).input_ids.to(pixel_values.device)
|
||||||
|
encoder_hidden_states = text_encoder(prompt_ids)[0]
|
||||||
|
|
||||||
|
with torch.cuda.amp.autocast(enabled=mixed_precision_training):
|
||||||
|
if mask_spatial_lora:
|
||||||
|
loras = extract_lora_child_module(unet, target_replace_module=target_spatial_modules)
|
||||||
|
scale_loras(loras, 0.)
|
||||||
|
loss_spatial = None
|
||||||
|
else:
|
||||||
|
loras = extract_lora_child_module(unet, target_replace_module=target_spatial_modules)
|
||||||
|
if spatial_lora_num == 1:
|
||||||
|
scale_loras(loras, 1.0)
|
||||||
|
else:
|
||||||
|
scale_loras(loras, 0.)
|
||||||
|
scale_loras(loras, 1.0, step=step, spatial_lora_num=spatial_lora_num)
|
||||||
|
|
||||||
|
loras = extract_lora_child_module(unet, target_replace_module=target_temporal_modules)
|
||||||
|
if len(loras) > 0:
|
||||||
|
scale_loras(loras, 0.)
|
||||||
|
|
||||||
|
### >>>> Spatial LoRA Prediction >>>> ###
|
||||||
|
noisy_latents = train_noise_scheduler_spatial.add_noise(latents, noise, timesteps)
|
||||||
|
noisy_latents_input, target_spatial, use_hflip = get_spatial_latents(
|
||||||
|
pixel_values,
|
||||||
|
random_hflip_img,
|
||||||
|
cache_latents,
|
||||||
|
noisy_latents,
|
||||||
|
target,
|
||||||
|
timesteps,
|
||||||
|
train_noise_scheduler_spatial
|
||||||
|
)
|
||||||
|
|
||||||
|
if use_hflip:
|
||||||
|
model_pred_spatial = unet(noisy_latents_input, timesteps,
|
||||||
|
encoder_hidden_states=encoder_hidden_states).sample
|
||||||
|
model_pred_spatial.requires_grad_(True)
|
||||||
|
target_spatial.requires_grad_(True)
|
||||||
|
loss_spatial = F.mse_loss(model_pred_spatial[:, :, 0, :, :].float(),
|
||||||
|
target_spatial[:, :, 0, :, :].float(), reduction="mean")
|
||||||
|
else:
|
||||||
|
model_pred_spatial = unet(noisy_latents_input.unsqueeze(2), timesteps,
|
||||||
|
encoder_hidden_states=encoder_hidden_states).sample
|
||||||
|
model_pred_spatial.requires_grad_(True)
|
||||||
|
target_spatial.requires_grad_(True)
|
||||||
|
loss_spatial = F.mse_loss(model_pred_spatial[:, :, 0, :, :].float(),
|
||||||
|
target_spatial.float(), reduction="mean")
|
||||||
|
|
||||||
|
if mask_temporal_lora:
|
||||||
|
loras = extract_lora_child_module(unet, target_replace_module=target_temporal_modules)
|
||||||
|
scale_loras(loras, 0.)
|
||||||
|
loss_temporal = None
|
||||||
|
|
||||||
|
else:
|
||||||
|
loras = extract_lora_child_module(unet, target_replace_module=target_temporal_modules)
|
||||||
|
scale_loras(loras, 1.0)
|
||||||
|
|
||||||
|
### >>>> Temporal LoRA Prediction >>>> ###
|
||||||
|
noisy_latents = train_noise_scheduler.add_noise(latents, noise, timesteps)
|
||||||
|
model_pred = unet(noisy_latents, timesteps, encoder_hidden_states=encoder_hidden_states).sample
|
||||||
|
|
||||||
|
loss_temporal = F.mse_loss(model_pred.float(), target.float(), reduction="mean")
|
||||||
|
loss_temporal = create_ad_temporal_loss(model_pred, loss_temporal, target)
|
||||||
|
|
||||||
|
# Backpropagate
|
||||||
|
if not mask_spatial_lora:
|
||||||
|
scaler.scale(loss_spatial).backward(retain_graph=True)
|
||||||
|
if spatial_lora_num == 1:
|
||||||
|
scaler.step(optimizer_spatial_list[0])
|
||||||
|
|
||||||
|
else:
|
||||||
|
# https://github.com/nerfstudio-project/nerfstudio/pull/1919
|
||||||
|
if any(
|
||||||
|
any(p.grad is not None for p in g["params"]) for g in optimizer_spatial_list[step].param_groups
|
||||||
|
):
|
||||||
|
scaler.step(optimizer_spatial_list[step])
|
||||||
|
|
||||||
|
if not mask_temporal_lora and train_temporal_lora:
|
||||||
|
scaler.scale(loss_temporal).backward()
|
||||||
|
scaler.step(optimizer_temporal)
|
||||||
|
|
||||||
|
if spatial_lora_num == 1:
|
||||||
|
lr_scheduler_spatial_list[0].step()
|
||||||
|
spatial_scheduler_lr = lr_scheduler_spatial_list[0].get_lr()[0]
|
||||||
|
else:
|
||||||
|
lr_scheduler_spatial_list[step].step()
|
||||||
|
spatial_scheduler_lr = lr_scheduler_spatial_list[step].get_lr()[0]
|
||||||
|
|
||||||
|
if lr_scheduler_temporal is not None:
|
||||||
|
lr_scheduler_temporal.step()
|
||||||
|
temporal_scheduler_lr = lr_scheduler_temporal.get_lr()[0]
|
||||||
|
|
||||||
|
scaler.update()
|
||||||
|
progress_bar.update(1)
|
||||||
|
pbar.update(1)
|
||||||
|
global_step += 1
|
||||||
|
|
||||||
|
# Save checkpoint
|
||||||
|
if global_step % checkpointing_steps == 0:
|
||||||
|
import copy
|
||||||
|
|
||||||
|
# We do this to prevent VRAM spiking / increase from the new copy
|
||||||
|
validation_pipeline.to('cpu')
|
||||||
|
|
||||||
|
lora_manager_spatial.save_lora_weights(
|
||||||
|
model=copy.deepcopy(validation_pipeline),
|
||||||
|
save_path=lora_path+'/spatial',
|
||||||
|
step=global_step,
|
||||||
|
use_safetensors=True,
|
||||||
|
lora_rank=lora_rank,
|
||||||
|
lora_name=lora_name + "_spatial"
|
||||||
|
)
|
||||||
|
|
||||||
|
if lora_manager_temporal is not None:
|
||||||
|
lora_manager_temporal.save_lora_weights(
|
||||||
|
model=copy.deepcopy(validation_pipeline),
|
||||||
|
save_path=lora_path+'/temporal',
|
||||||
|
step=global_step,
|
||||||
|
use_safetensors=True,
|
||||||
|
lora_rank=lora_rank,
|
||||||
|
lora_name=lora_name + "_temporal",
|
||||||
|
use_motion_lora_format=use_motion_lora_format
|
||||||
|
)
|
||||||
|
|
||||||
|
validation_pipeline.to(device)
|
||||||
|
|
||||||
|
# Periodically validation
|
||||||
|
if (global_step % validation_steps == 0 or global_step in validation_steps_tuple):
|
||||||
|
samples = []
|
||||||
|
generator = torch.Generator(device=latents.device)
|
||||||
|
generator.manual_seed(global_seed if validation_seed == -1 else validation_seed)
|
||||||
|
|
||||||
|
if not train_sample_validation:
|
||||||
|
height, width = input_height, input_width
|
||||||
|
else:
|
||||||
|
height, width = [512] * 2
|
||||||
|
|
||||||
|
with torch.cuda.amp.autocast(enabled=True):
|
||||||
|
if gradient_checkpointing:
|
||||||
|
unet.disable_gradient_checkpointing()
|
||||||
|
|
||||||
|
loras = extract_lora_child_module(
|
||||||
|
unet,
|
||||||
|
target_replace_module=target_spatial_modules
|
||||||
|
)
|
||||||
|
scale_loras(loras, validation_spatial_scale)
|
||||||
|
|
||||||
|
with torch.inference_mode(False):
|
||||||
|
unet.eval()
|
||||||
|
|
||||||
|
if len(validation_prompt) == 0:
|
||||||
|
prompt = text_prompt
|
||||||
|
else:
|
||||||
|
prompt = validation_prompt
|
||||||
|
print(prompt)
|
||||||
|
sample = validation_pipeline(
|
||||||
|
prompt,
|
||||||
|
generator = generator,
|
||||||
|
video_length = video_length,
|
||||||
|
height = height,
|
||||||
|
width = width,
|
||||||
|
).videos
|
||||||
|
save_videos_grid(sample, f"{output_dir}/samples/sample-{global_step}.gif")
|
||||||
|
samples.append(sample)
|
||||||
|
|
||||||
|
unet.train()
|
||||||
|
|
||||||
|
samples = torch.concat(samples)
|
||||||
|
save_path = f"{output_dir}/samples/sample-{global_step}.gif"
|
||||||
|
save_videos_grid(samples, save_path)
|
||||||
|
|
||||||
|
logging.info(f"Saved samples to {save_path}")
|
||||||
|
|
||||||
|
logs = {
|
||||||
|
"Temporal Loss": loss_temporal.detach().item(),
|
||||||
|
"Temporal LR": temporal_scheduler_lr,
|
||||||
|
"Spatial Loss": loss_spatial.detach().item() if loss_spatial is not None else 0,
|
||||||
|
"Spatial LR": spatial_scheduler_lr
|
||||||
|
}
|
||||||
|
progress_bar.set_postfix(**logs)
|
||||||
|
|
||||||
|
if gradient_checkpointing:
|
||||||
|
unet.enable_gradient_checkpointing()
|
||||||
|
|
||||||
|
if global_step >= max_train_steps:
|
||||||
|
break
|
||||||
|
return samples,
|
||||||
|
|
||||||
|
import folder_paths
|
||||||
|
class DiffusersLoaderForTraining:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(cls):
|
||||||
|
paths = []
|
||||||
|
for search_path in folder_paths.get_folder_paths("diffusers"):
|
||||||
|
if os.path.exists(search_path):
|
||||||
|
for root, subdir, files in os.walk(search_path, followlinks=True):
|
||||||
|
if "model_index.json" in files:
|
||||||
|
paths.append(os.path.relpath(root, start=search_path))
|
||||||
|
|
||||||
|
return {"required":
|
||||||
|
{
|
||||||
|
"download_default": ("BOOLEAN", {"default": False},),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"model": (paths,),
|
||||||
|
}
|
||||||
|
|
||||||
|
}
|
||||||
|
RETURN_TYPES = ("MODEL", "CLIP", "TOKENIZER", "VAE")
|
||||||
|
FUNCTION = "load_checkpoint"
|
||||||
|
|
||||||
|
CATEGORY = "AD_MotionDirector"
|
||||||
|
|
||||||
|
def load_checkpoint(self, download_default, model=""):
|
||||||
|
with torch.inference_mode(False):
|
||||||
|
print(model)
|
||||||
|
if model != "":
|
||||||
|
model_path = model
|
||||||
|
else:
|
||||||
|
if download_default:
|
||||||
|
from huggingface_hub import snapshot_download
|
||||||
|
download_to = os.path.join(folder_paths.models_dir,'diffusers')
|
||||||
|
snapshot_download(repo_id="runwayml/stable-diffusion-v1-5", ignore_patterns=["*.safetensors","*.ckpt", "*.pt", "*.png", "*non_ema*", "*fp16*"],
|
||||||
|
local_dir=f"{download_to}/stable-diffusion-v1-5", local_dir_use_symlinks=False)
|
||||||
|
|
||||||
|
for search_path in folder_paths.get_folder_paths("diffusers"):
|
||||||
|
if os.path.exists(search_path):
|
||||||
|
path = os.path.join(search_path, model_path)
|
||||||
|
if os.path.exists(path):
|
||||||
|
model_path = path
|
||||||
|
break
|
||||||
|
|
||||||
|
config = OmegaConf.load(os.path.join(script_directory, f"configs/training/motion_director/training.yaml"))
|
||||||
|
vae = AutoencoderKL.from_pretrained(model_path, subfolder="vae")
|
||||||
|
tokenizer = CLIPTokenizer.from_pretrained(model_path, subfolder="tokenizer")
|
||||||
|
text_encoder = CLIPTextModel.from_pretrained(model_path, subfolder="text_encoder")
|
||||||
|
|
||||||
|
unet_additional_kwargs = config.unet_additional_kwargs
|
||||||
|
unet = UNet3DConditionModel.from_pretrained_2d(
|
||||||
|
model_path, subfolder="unet",
|
||||||
|
unet_additional_kwargs=unet_additional_kwargs
|
||||||
|
)
|
||||||
|
return (unet, text_encoder, tokenizer, vae,)
|
||||||
|
|
||||||
|
folder_paths.add_model_folder_path("animatediff_models", str(Path(__file__).parent.parent / "models"))
|
||||||
|
folder_paths.add_model_folder_path("animatediff_models", str(Path(folder_paths.models_dir) / "animatediff_models"))
|
||||||
|
|
||||||
|
class ValidationModelSelect:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"motion_module": (folder_paths.get_filename_list("animatediff_models"),),
|
||||||
|
"use_adapter_lora": ("BOOLEAN", {"default": True}),
|
||||||
|
"use_dreambooth_model": ("BOOLEAN", {"default": False}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"optional_adapter_lora": (folder_paths.get_filename_list("loras"),),
|
||||||
|
"optional_model": (folder_paths.get_filename_list("checkpoints"),),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
RETURN_TYPES = ("VALIDATION_MODELS",)
|
||||||
|
RETURN_NAMES = ("validation_models",)
|
||||||
|
FUNCTION = "select_models"
|
||||||
|
|
||||||
|
CATEGORY = "AD_MotionDirector"
|
||||||
|
|
||||||
|
def select_models(self, motion_module, use_adapter_lora, use_dreambooth_model, optional_adapter_lora="", optional_model=""):
|
||||||
|
validation_models = []
|
||||||
|
motion_module_path = folder_paths.get_full_path("animatediff_models", motion_module)
|
||||||
|
|
||||||
|
if use_adapter_lora:
|
||||||
|
adapter_lora_path = folder_paths.get_full_path("loras", optional_adapter_lora)
|
||||||
|
else:
|
||||||
|
adapter_lora_path = ""
|
||||||
|
if use_dreambooth_model:
|
||||||
|
model_path = folder_paths.get_full_path("checkpoints", optional_model)
|
||||||
|
else:
|
||||||
|
model_path = ""
|
||||||
|
|
||||||
|
validation_models.append(motion_module_path)
|
||||||
|
validation_models.append(adapter_lora_path)
|
||||||
|
validation_models.append(model_path)
|
||||||
|
return (validation_models,)
|
||||||
|
|
||||||
|
NODE_CLASS_MAPPINGS = {
|
||||||
|
"AD_MotionDirector_train": AD_MotionDirector_train,
|
||||||
|
"DiffusersLoaderForTraining": DiffusersLoaderForTraining,
|
||||||
|
"ValidationModelSelect": ValidationModelSelect
|
||||||
|
}
|
||||||
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
|
"AD_MotionDirector_train": "AD_MotionDirector_train",
|
||||||
|
"DiffusersLoaderForTraining": "DiffusersLoaderForTraining",
|
||||||
|
"ValidationModelSelect": "ValidationModelSelect"
|
||||||
|
}
|
||||||
@@ -0,0 +1,8 @@
|
|||||||
|
diffusers>=0.26.0
|
||||||
|
huggingface_hub>=0.20.3
|
||||||
|
transformers>=4.27.4
|
||||||
|
loralib
|
||||||
|
einops
|
||||||
|
omegaconf
|
||||||
|
lion-pytorch
|
||||||
|
peft
|
||||||
Reference in New Issue
Block a user