8 Commits
Author SHA1 Message Date
kijai d312c62982 Possible import error fix 2025-06-02 14:48:59 +03:00
kijai 7670b09c28 Fix other nodes that don't use controlnet (yet) 2024-07-20 19:22:46 +03:00
kijai 282c69c5f3 Update nodes.py 2024-07-20 00:27:22 +03:00
kijai 18121e97a2 Update nodes.py 2024-07-19 22:08:31 +03:00
kijai 4a7dd402b4 Support other SD 2.1 controlnets as well 2024-07-19 20:27:52 +03:00
kijai 2545a3afb7 Initial sketch controlnet support 2024-07-19 19:50:26 +03:00
kijai 0668ee6d0a fix DynamiCrafterBatchInterpolation too 2024-07-09 16:15:03 +03:00
kijai 69fe0e2e9f don't normalize image2 twice 2024-07-09 16:06:04 +03:00
5 changed files with 634 additions and 30 deletions
+17 -4
View File
@@ -387,6 +387,7 @@ class LatentDiffusion(DDPM):
logdir=None, logdir=None,
rand_cond_frame=False, rand_cond_frame=False,
en_and_decode_n_samples_a_time=None, en_and_decode_n_samples_a_time=None,
control_scale=1.0,
*args, **kwargs): *args, **kwargs):
self.num_timesteps_cond = default(num_timesteps_cond, 1) self.num_timesteps_cond = default(num_timesteps_cond, 1)
self.scale_by_std = scale_by_std self.scale_by_std = scale_by_std
@@ -404,6 +405,7 @@ class LatentDiffusion(DDPM):
self.loop_video = loop_video self.loop_video = loop_video
self.fps_condition_type = fps_condition_type self.fps_condition_type = fps_condition_type
self.perframe_ae = perframe_ae self.perframe_ae = perframe_ae
self.control_scale = control_scale
self.logdir = logdir self.logdir = logdir
self.rand_cond_frame = rand_cond_frame self.rand_cond_frame = rand_cond_frame
@@ -572,9 +574,20 @@ class LatentDiffusion(DDPM):
if not isinstance(cond, list): if not isinstance(cond, list):
cond = [cond] cond = [cond]
key = 'c_concat' if self.model.conditioning_key == 'concat' else 'c_crossattn' key = 'c_concat' if self.model.conditioning_key == 'concat' else 'c_crossattn'
cond = {key: cond} cond = {key: [cond[0]]}
x_recon = self.model(x_noisy, t, **cond, **kwargs) control_cond = cond["control_cond"]
if control_cond is not None:
control_cond = rearrange(control_cond, 'b c t h w-> (b t) c h w')
control_x = rearrange(x_noisy, 'b c t h w-> (b t) c h w')
control_context = repeat(cond["c_crossattn"][0], "b c l-> (repeat b) c l", repeat=16)
control = self.control_model(x=control_x, hint=control_cond, timesteps=t, context=control_context)
control = [c * self.control_model.control_scale for c in control]
else:
control = None
x_recon = self.model(x_noisy, t, c_crossattn=cond["c_crossattn"], c_concat=cond["c_concat"], control=control, **kwargs)
if isinstance(x_recon, tuple): if isinstance(x_recon, tuple):
return x_recon[0] return x_recon[0]
@@ -722,7 +735,7 @@ class DiffusionWrapper(pl.LightningModule):
self.diffusion_model = instantiate_from_config(diff_model_config) self.diffusion_model = instantiate_from_config(diff_model_config)
self.conditioning_key = conditioning_key self.conditioning_key = conditioning_key
def forward(self, x, t, c_concat: list = None, c_crossattn: list = None, def forward(self, x, t, c_concat: list = None, c_crossattn: list = None, control = None,
c_adm=None, s=None, mask=None, **kwargs): c_adm=None, s=None, mask=None, **kwargs):
# temporal_context = fps is foNone # temporal_context = fps is foNone
if self.conditioning_key is None: if self.conditioning_key is None:
@@ -737,7 +750,7 @@ class DiffusionWrapper(pl.LightningModule):
## it is just right [b,c,t,h,w]: concatenate in channel dim ## it is just right [b,c,t,h,w]: concatenate in channel dim
xc = torch.cat([x] + c_concat, dim=1) xc = torch.cat([x] + c_concat, dim=1)
cc = torch.cat(c_crossattn, 1) cc = torch.cat(c_crossattn, 1)
out = self.diffusion_model(xc, t, context=cc, **kwargs) out = self.diffusion_model(xc, t, context=cc, control=control, **kwargs)
elif self.conditioning_key == 'resblockcond': elif self.conditioning_key == 'resblockcond':
cc = c_crossattn[0] cc = c_crossattn[0]
out = self.diffusion_model(x, t, context=cc) out = self.diffusion_model(x, t, context=cc)
+404 -1
View File
@@ -2,7 +2,9 @@ from functools import partial
from abc import abstractmethod from abc import abstractmethod
import torch import torch
import torch.nn as nn import torch.nn as nn
import numpy as np
from einops import rearrange from einops import rearrange
import math
import torch.nn.functional as F import torch.nn.functional as F
from ....lvdm.models.utils_diffusion import timestep_embedding from ....lvdm.models.utils_diffusion import timestep_embedding
from ....lvdm.common import checkpoint from ....lvdm.common import checkpoint
@@ -18,6 +20,9 @@ from ....lvdm.modules.attention import SpatialTransformer, TemporalTransformer
import comfy.ops import comfy.ops
ops = comfy.ops.manual_cast ops = comfy.ops.manual_cast
def exists(x):
return x is not None
class TimestepBlock(nn.Module): class TimestepBlock(nn.Module):
""" """
Any module where forward() takes timestep embeddings as a second argument. Any module where forward() takes timestep embeddings as a second argument.
@@ -547,7 +552,7 @@ class UNetModel(nn.Module):
zero_module(conv_nd(dims, model_channels, out_channels, 3, padding=1)), zero_module(conv_nd(dims, model_channels, out_channels, 3, padding=1)),
) )
def forward(self, x, timesteps, context=None, features_adapter=None, fs=None, frame_window_size=None, frame_window_stride=None, **kwargs): def forward(self, x, timesteps, context=None, features_adapter=None, fs=None, frame_window_size=None, frame_window_stride=None, control=None, **kwargs):
b,_,t,_,_ = x.shape b,_,t,_,_ = x.shape
t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).type(x.dtype) t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).type(x.dtype)
emb = self.time_embed(t_emb) emb = self.time_embed(t_emb)
@@ -594,8 +599,15 @@ class UNetModel(nn.Module):
assert len(features_adapter)==adapter_idx, 'Wrong features_adapter' assert len(features_adapter)==adapter_idx, 'Wrong features_adapter'
h = self.middle_block(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride) h = self.middle_block(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride)
if control is not None:
h += control.pop()
for module in self.output_blocks: for module in self.output_blocks:
if control is None:
h = torch.cat([h, hs.pop()], dim=1) h = torch.cat([h, hs.pop()], dim=1)
else:
h = torch.cat([h, hs.pop() + control.pop()], dim=1)
h = module(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride) h = module(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride)
h = h.type(x.dtype) h = h.type(x.dtype)
y = self.out(h) y = self.out(h)
@@ -603,3 +615,394 @@ class UNetModel(nn.Module):
# reshape back to (b c t h w) # reshape back to (b c t h w)
y = rearrange(y, '(b t) c h w -> b c t h w', b=b) y = rearrange(y, '(b t) c h w -> b c t h w', b=b)
return y return y
class ControlNet(nn.Module):
def __init__(
self,
image_size,
in_channels,
model_channels,
hint_channels,
num_res_blocks,
attention_resolutions,
dropout=0,
channel_mult=(1, 2, 4, 8),
conv_resample=True,
dims=2,
use_checkpoint=False,
use_fp16=False,
num_heads=-1,
num_head_channels=-1,
num_heads_upsample=-1,
use_scale_shift_norm=False,
resblock_updown=False,
use_new_attention_order=False,
use_spatial_transformer=False, # custom transformer support
transformer_depth=1, # custom transformer support
context_dim=None, # custom transformer support
n_embed=None, # custom support for prediction of discrete ids into codebook of first stage vq model
legacy=True,
disable_self_attentions=None,
num_attention_blocks=None,
disable_middle_self_attn=False,
use_linear_in_transformer=False,
):
super().__init__()
if use_spatial_transformer:
assert context_dim is not None, 'Fool!! You forgot to include the dimension of your cross-attention conditioning...'
if context_dim is not None:
assert use_spatial_transformer, 'Fool!! You forgot to use the spatial transformer for your cross-attention conditioning...'
from omegaconf.listconfig import ListConfig
if type(context_dim) == ListConfig:
context_dim = list(context_dim)
if num_heads_upsample == -1:
num_heads_upsample = num_heads
if num_heads == -1:
assert num_head_channels != -1, 'Either num_heads or num_head_channels has to be set'
if num_head_channels == -1:
assert num_heads != -1, 'Either num_heads or num_head_channels has to be set'
self.dims = dims
self.image_size = image_size
self.in_channels = in_channels
self.model_channels = model_channels
if isinstance(num_res_blocks, int):
self.num_res_blocks = len(channel_mult) * [num_res_blocks]
else:
if len(num_res_blocks) != len(channel_mult):
raise ValueError("provide num_res_blocks either as an int (globally constant) or "
"as a list/tuple (per-level) with the same length as channel_mult")
self.num_res_blocks = num_res_blocks
if disable_self_attentions is not None:
# should be a list of booleans, indicating whether to disable self-attention in TransformerBlocks or not
assert len(disable_self_attentions) == len(channel_mult)
if num_attention_blocks is not None:
assert len(num_attention_blocks) == len(self.num_res_blocks)
assert all(map(lambda i: self.num_res_blocks[i] >= num_attention_blocks[i], range(len(num_attention_blocks))))
print(f"Constructor of UNetModel received num_attention_blocks={num_attention_blocks}. "
f"This option has LESS priority than attention_resolutions {attention_resolutions}, "
f"i.e., in cases where num_attention_blocks[i] > 0 but 2**i not in attention_resolutions, "
f"attention will still not be set.")
self.attention_resolutions = attention_resolutions
self.dropout = dropout
self.channel_mult = channel_mult
self.conv_resample = conv_resample
self.use_checkpoint = use_checkpoint
self.dtype = torch.float16 if use_fp16 else torch.float32
self.num_heads = num_heads
self.num_head_channels = num_head_channels
self.num_heads_upsample = num_heads_upsample
self.predict_codebook_ids = n_embed is not None
time_embed_dim = model_channels * 4
self.time_embed = nn.Sequential(
linear(model_channels, time_embed_dim),
nn.SiLU(),
linear(time_embed_dim, time_embed_dim),
)
self.input_blocks = nn.ModuleList(
[
TimestepEmbedSequential(
conv_nd(dims, in_channels, model_channels, 3, padding=1)
)
]
)
self.zero_convs = nn.ModuleList([self.make_zero_conv(model_channels)])
self.input_hint_block = TimestepEmbedSequential(
conv_nd(dims, hint_channels, 16, 3, padding=1),
nn.SiLU(),
conv_nd(dims, 16, 16, 3, padding=1),
nn.SiLU(),
conv_nd(dims, 16, 32, 3, padding=1, stride=2),
nn.SiLU(),
conv_nd(dims, 32, 32, 3, padding=1),
nn.SiLU(),
conv_nd(dims, 32, 96, 3, padding=1, stride=2),
nn.SiLU(),
conv_nd(dims, 96, 96, 3, padding=1),
nn.SiLU(),
conv_nd(dims, 96, 256, 3, padding=1, stride=2),
nn.SiLU(),
zero_module(conv_nd(dims, 256, model_channels, 3, padding=1))
)
self._feature_size = model_channels
input_block_chans = [model_channels]
ch = model_channels
ds = 1
for level, mult in enumerate(channel_mult):
for nr in range(self.num_res_blocks[level]):
layers = [
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=mult * model_channels,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
)
]
ch = mult * model_channels
if ds in attention_resolutions:
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
if legacy:
# num_heads = 1
dim_head = ch // num_heads if use_spatial_transformer else num_head_channels
if exists(disable_self_attentions):
disabled_sa = disable_self_attentions[level]
else:
disabled_sa = False
if not exists(num_attention_blocks) or nr < num_attention_blocks[level]:
layers.append(
AttentionBlock(
ch,
use_checkpoint=use_checkpoint,
num_heads=num_heads,
num_head_channels=dim_head,
use_new_attention_order=use_new_attention_order,
) if not use_spatial_transformer else SpatialTransformer(
ch, num_heads, dim_head, depth=transformer_depth, context_dim=context_dim,
disable_self_attn=disabled_sa, use_linear=use_linear_in_transformer,
use_checkpoint=use_checkpoint
)
)
self.input_blocks.append(TimestepEmbedSequential(*layers))
self.zero_convs.append(self.make_zero_conv(ch))
self._feature_size += ch
input_block_chans.append(ch)
if level != len(channel_mult) - 1:
out_ch = ch
self.input_blocks.append(
TimestepEmbedSequential(
ResBlock(
ch,
time_embed_dim,
dropout,
out_channels=out_ch,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
down=True,
)
if resblock_updown
else Downsample(
ch, conv_resample, dims=dims, out_channels=out_ch
)
)
)
ch = out_ch
input_block_chans.append(ch)
self.zero_convs.append(self.make_zero_conv(ch))
ds *= 2
self._feature_size += ch
if num_head_channels == -1:
dim_head = ch // num_heads
else:
num_heads = ch // num_head_channels
dim_head = num_head_channels
if legacy:
# num_heads = 1
dim_head = ch // num_heads if use_spatial_transformer else num_head_channels
self.middle_block = TimestepEmbedSequential(
ResBlock(
ch,
time_embed_dim,
dropout,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
),
AttentionBlock(
ch,
use_checkpoint=use_checkpoint,
num_heads=num_heads,
num_head_channels=dim_head,
use_new_attention_order=use_new_attention_order,
) if not use_spatial_transformer else SpatialTransformer( # always uses a self-attn
ch, num_heads, dim_head, depth=transformer_depth, context_dim=context_dim,
disable_self_attn=disable_middle_self_attn, use_linear=use_linear_in_transformer,
use_checkpoint=use_checkpoint
),
ResBlock(
ch,
time_embed_dim,
dropout,
dims=dims,
use_checkpoint=use_checkpoint,
use_scale_shift_norm=use_scale_shift_norm,
),
)
self.middle_block_out = self.make_zero_conv(ch)
self._feature_size += ch
def make_zero_conv(self, channels):
return TimestepEmbedSequential(zero_module(conv_nd(self.dims, channels, channels, 1, padding=0)))
def forward(self, x, hint, timesteps, context, **kwargs):
t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False)
emb = self.time_embed(t_emb)
guided_hint = self.input_hint_block(hint, emb, context)
outs = []
h = x.type(self.dtype)
for module, zero_conv in zip(self.input_blocks, self.zero_convs):
if guided_hint is not None:
h = module(h, emb, context)
h += guided_hint
guided_hint = None
else:
h = module(h, emb, context)
outs.append(zero_conv(h, emb, context, True))
h = self.middle_block(h, emb, context)
outs.append(self.middle_block_out(h, emb, context))
return outs
class AttentionBlock(nn.Module):
"""
An attention block that allows spatial positions to attend to each other.
Originally ported from here, but adapted to the N-d case.
https://github.com/hojonathanho/diffusion/blob/1e0dceb3b3495bbe19116a5e1b3596cd0706c543/diffusion_tf/models/unet.py#L66.
"""
def __init__(
self,
channels,
num_heads=1,
num_head_channels=-1,
use_checkpoint=False,
use_new_attention_order=False,
):
super().__init__()
self.channels = channels
if num_head_channels == -1:
self.num_heads = num_heads
else:
assert (
channels % num_head_channels == 0
), f"q,k,v channels {channels} is not divisible by num_head_channels {num_head_channels}"
self.num_heads = channels // num_head_channels
self.use_checkpoint = use_checkpoint
self.norm = normalization(channels)
self.qkv = conv_nd(1, channels, channels * 3, 1)
if use_new_attention_order:
# split qkv before split heads
self.attention = QKVAttention(self.num_heads)
else:
# split heads before split qkv
self.attention = QKVAttentionLegacy(self.num_heads)
self.proj_out = zero_module(conv_nd(1, channels, channels, 1))
def forward(self, x):
return checkpoint(self._forward, (x,), self.parameters(), True) # TODO: check checkpoint usage, is True # TODO: fix the .half call!!!
#return pt_checkpoint(self._forward, x) # pytorch
def _forward(self, x):
b, c, *spatial = x.shape
x = x.reshape(b, c, -1)
qkv = self.qkv(self.norm(x))
h = self.attention(qkv)
h = self.proj_out(h)
return (x + h).reshape(b, c, *spatial)
class QKVAttention(nn.Module):
"""
A module which performs QKV attention and splits in a different order.
"""
def __init__(self, n_heads):
super().__init__()
self.n_heads = n_heads
def forward(self, qkv):
"""
Apply QKV attention.
:param qkv: an [N x (3 * H * C) x T] tensor of Qs, Ks, and Vs.
:return: an [N x (H * C) x T] tensor after attention.
"""
bs, width, length = qkv.shape
assert width % (3 * self.n_heads) == 0
ch = width // (3 * self.n_heads)
q, k, v = qkv.chunk(3, dim=1)
scale = 1 / math.sqrt(math.sqrt(ch))
weight = torch.einsum(
"bct,bcs->bts",
(q * scale).view(bs * self.n_heads, ch, length),
(k * scale).view(bs * self.n_heads, ch, length),
) # More stable with f16 than dividing afterwards
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
a = torch.einsum("bts,bcs->bct", weight, v.reshape(bs * self.n_heads, ch, length))
return a.reshape(bs, -1, length)
@staticmethod
def count_flops(model, _x, y):
return count_flops_attn(model, _x, y)
class QKVAttentionLegacy(nn.Module):
"""
A module which performs QKV attention. Matches legacy QKVAttention + input/ouput heads shaping
"""
def __init__(self, n_heads):
super().__init__()
self.n_heads = n_heads
def forward(self, qkv):
"""
Apply QKV attention.
:param qkv: an [N x (H * 3 * C) x T] tensor of Qs, Ks, and Vs.
:return: an [N x (H * C) x T] tensor after attention.
"""
bs, width, length = qkv.shape
assert width % (3 * self.n_heads) == 0
ch = width // (3 * self.n_heads)
q, k, v = qkv.reshape(bs * self.n_heads, ch * 3, length).split(ch, dim=1)
scale = 1 / math.sqrt(math.sqrt(ch))
weight = torch.einsum(
"bct,bcs->bts", q * scale, k * scale
) # More stable with f16 than dividing afterwards
weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype)
a = torch.einsum("bts,bcs->bct", weight, v)
return a.reshape(bs, -1, length)
@staticmethod
def count_flops(model, _x, y):
return count_flops_attn(model, _x, y)
def count_flops_attn(model, _x, y):
"""
A counter for the `thop` package to count the operations in an
attention operation.
Meant to be used like:
macs, params = thop.profile(
model,
inputs=(inputs, timestamps),
custom_ops={QKVAttention: QKVAttention.count_flops},
)
"""
b, c, *spatial = y[0].shape
num_spatial = int(np.prod(spatial))
# We perform two matmuls with the same number of ops.
# The first computes the weight matrix, the second computes
# the combination of the value vectors.
matmul_ops = 2 * b * (num_spatial ** 2) * c
model.total_ops += torch.DoubleTensor([matmul_ops])
+206 -21
View File
@@ -10,6 +10,7 @@ import comfy.model_management as mm
import comfy.utils import comfy.utils
from contextlib import nullcontext from contextlib import nullcontext
from .lvdm.models.samplers.ddim import DDIMSampler from .lvdm.models.samplers.ddim import DDIMSampler
from .lvdm.modules.networks.openaimodel3d import ControlNet
from contextlib import nullcontext from contextlib import nullcontext
try: try:
@@ -80,7 +81,7 @@ class DownloadAndLoadDynamiCrafterModel:
'ckpt_name': model, 'ckpt_name': model,
'fp8_unet': fp8_unet 'fp8_unet': fp8_unet
} }
if not hasattr(self, 'model') or self.model == None or custom_config != self.current_config: if not hasattr(self, 'model') or self.model is None or custom_config != self.current_config:
self.current_config = custom_config self.current_config = custom_config
download_path = os.path.join(folder_paths.models_dir, "checkpoints", "dynamicrafter") download_path = os.path.join(folder_paths.models_dir, "checkpoints", "dynamicrafter")
model_path = os.path.join(download_path, model) model_path = os.path.join(download_path, model)
@@ -140,9 +141,131 @@ class DownloadAndLoadDynamiCrafterModel:
dcmodel = { dcmodel = {
'model': self.model, 'model': self.model,
'model_name': model, 'model_name': model,
'dtype': precision
} }
return (dcmodel,) return (dcmodel,)
class DownloadAndLoadDynamiCrafterCNModel:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"model": (
[
'sketch_encoder-fp16.safetensors',
],
),
},
}
RETURN_TYPES = ("DC_CN_MODEL",)
RETURN_NAMES = ("DynCraft_CN_model",)
FUNCTION = "loadmodel"
CATEGORY = "DynamiCrafterWrapper"
def loadmodel(self, model):
custom_config = {
'ckpt_name': model,
}
if not hasattr(self, 'cn_model') or self.cn_model == None or custom_config != self.current_config:
download_path = os.path.join(folder_paths.models_dir, "checkpoints", "dynamicrafter", "controlnet")
cn_model_path = os.path.join(download_path, model)
if not os.path.exists(cn_model_path):
print(f"Downloading model to: {cn_model_path}")
from huggingface_hub import snapshot_download
snapshot_download(repo_id="Kijai/DynamiCrafter_pruned",
allow_patterns=[f"*{model}*"],
local_dir=download_path,
local_dir_use_symlinks=False)
cn_config = {
"use_checkpoint": True,
"image_size": 32, # unused
"in_channels": 4,
"hint_channels": 3,
"model_channels": 320,
"attention_resolutions": [4, 2, 1],
"num_res_blocks": 2,
"channel_mult": [1, 2, 4, 4],
"num_head_channels": 64, # need to fix for flash-attn
"use_spatial_transformer": True,
"use_linear_in_transformer": True,
"transformer_depth": 1,
"context_dim": 1024,
"legacy": False
}
if "sketch_encoder" in model:
cn_config["hint_channels"] = 1
self.cn_model = ControlNet(**cn_config)
print("Loading ControlNet")
cn_sd = comfy.utils.load_torch_file(cn_model_path)
self.cn_model.load_state_dict(cn_sd, strict=True)
print("ControlNet loaded")
controlnet = {
'model': self.cn_model,
'config': cn_config,
}
return (controlnet,)
class DynamiCrafterCNLoader:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"ckpt_name": (folder_paths.get_filename_list("controlnet"), ),
},
}
RETURN_TYPES = ("DC_CN_MODEL",)
RETURN_NAMES = ("DynCraft_CN_model",)
FUNCTION = "loadmodel"
CATEGORY = "DynamiCrafterWrapper"
def loadmodel(self, ckpt_name):
custom_config = {
'ckpt_name': ckpt_name,
}
if not hasattr(self, 'cn_model') or self.cn_model == None or custom_config != self.current_config:
self.current_config = custom_config
model_path = folder_paths.get_full_path("controlnet", ckpt_name)
print(f"Loading ControlNet from: {model_path}")
cn_config = {
"use_checkpoint": True,
"image_size": 32, # unused
"in_channels": 4,
"hint_channels": 3,
"model_channels": 320,
"attention_resolutions": [4, 2, 1],
"num_res_blocks": 2,
"channel_mult": [1, 2, 4, 4],
"num_head_channels": 64, # need to fix for flash-attn
"use_spatial_transformer": True,
"use_linear_in_transformer": True,
"transformer_depth": 1,
"context_dim": 1024,
"legacy": False
}
if "sketch_encoder" in ckpt_name:
cn_config["hint_channels"] = 1
self.cn_model = ControlNet(**cn_config)
print("Loading ControlNet")
cn_sd = comfy.utils.load_torch_file(model_path)
self.cn_model.load_state_dict(cn_sd, strict=True)
del cn_sd
print("ControlNet loaded")
controlnet = {
'model': self.cn_model,
'config': cn_config,
}
return (controlnet,)
class DownloadAndLoadCLIPModel: class DownloadAndLoadCLIPModel:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -389,6 +512,7 @@ class DynamiCrafterI2V:
self.model.to(device) self.model.to(device)
autocast_condition = (dtype != torch.float32) and not comfy.model_management.is_device_mps(device) autocast_condition = (dtype != torch.float32) and not comfy.model_management.is_device_mps(device)
with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext(): with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
image = image.permute(0, 3, 1, 2).to(dtype).to(device) image = image.permute(0, 3, 1, 2).to(dtype).to(device)
if augmentation_level > 0: if augmentation_level > 0:
image += torch.randn_like(image) * augmentation_level image += torch.randn_like(image) * augmentation_level
@@ -400,7 +524,7 @@ class DynamiCrafterI2V:
if H % 64 != 0: if H % 64 != 0:
H = H - (H % 64) H = H - (H % 64)
if orig_H % 64 != 0 or orig_W % 64 != 0: if orig_H % 64 != 0 or orig_W % 64 != 0:
image = F.interpolate(image, size=(H, W), mode="bicubic") image = F.interpolate(image, size=(H, W), mode="bilinear")
B, C, H, W = image.shape B, C, H, W = image.shape
noise_shape = [B, self.model.model.diffusion_model.out_channels, frames, H // 8, W // 8] noise_shape = [B, self.model.model.diffusion_model.out_channels, frames, H // 8, W // 8]
@@ -410,14 +534,13 @@ class DynamiCrafterI2V:
z = get_latent_z(self.model, encode_pixels) #bc,1,hw z = get_latent_z(self.model, encode_pixels) #bc,1,hw
if image2 is not None: if image2 is not None:
image2 = image2 * 2 - 1
image2 = image2.permute(0, 3, 1, 2).to(dtype).to(device) image2 = image2.permute(0, 3, 1, 2).to(dtype).to(device)
if augmentation_level > 0: if augmentation_level > 0:
image2 += torch.randn_like(image2) * augmentation_level image2 += torch.randn_like(image2) * augmentation_level
if image2.shape != image.shape: if image2.shape != image.shape:
image2 = F.interpolate(image, size=(H, W), mode="bicubic") image2 = F.interpolate(image, size=(H, W), mode="bilinear")
encode_pixels = image2.unsqueeze(2) * 2 - 1 encode_pixels = image2.unsqueeze(2) * 2 - 1
z2 = get_latent_z(self.model, encode_pixels) #bc,1,hw z2 = get_latent_z(self.model, encode_pixels) #bc,1,hw
@@ -441,7 +564,9 @@ class DynamiCrafterI2V:
del cond_images, img_emb, text_emb, encode_pixels del cond_images, img_emb, text_emb, encode_pixels
fs = torch.tensor([fs], dtype=torch.long, device=self.model.device) fs = torch.tensor([fs], dtype=torch.long, device=self.model.device)
cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat]} cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat], "control_cond": None}
self.model.control_model = None
if noise_shape[-1] == 32: if noise_shape[-1] == 32:
timestep_spacing = "uniform" timestep_spacing = "uniform"
@@ -633,6 +758,7 @@ class ToonCrafterInterpolation:
"augmentation_level": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.0001}), "augmentation_level": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.0001}),
"optional_latents": ("LATENT",), "optional_latents": ("LATENT",),
"ddpm_from": ("INT", {"default": 1000, "min": 1, "max": 1000, "step": 1}), "ddpm_from": ("INT", {"default": 1000, "min": 1, "max": 1000, "step": 1}),
"controlnet": ("DC_CONTROL",),
} }
} }
@@ -641,7 +767,8 @@ class ToonCrafterInterpolation:
FUNCTION = "process" FUNCTION = "process"
CATEGORY = "DynamiCrafterWrapper" CATEGORY = "DynamiCrafterWrapper"
def process(self, model, clip_vision, images, positive, negative, cfg, steps, eta, seed, fs, frames, vae_dtype, image_embed_ratio=1.0, augmentation_level=0, optional_latents=None, ddpm_from=1000): def process(self, model, clip_vision, images, positive, negative, cfg, steps, eta, seed, fs, frames,
vae_dtype, image_embed_ratio=1.0, augmentation_level=0, optional_latents=None, ddpm_from=1000, controlnet=None):
device = mm.get_torch_device() device = mm.get_torch_device()
offload_device = mm.unet_offload_device() offload_device = mm.unet_offload_device()
mm.unload_all_models() mm.unload_all_models()
@@ -651,6 +778,11 @@ class ToonCrafterInterpolation:
self.model = model['model'] self.model = model['model']
if controlnet is not None:
self.model.control_model = controlnet["model"]
else:
self.model.control_model = None
dtype = self.model.dtype dtype = self.model.dtype
if vae_dtype == "auto": if vae_dtype == "auto":
try: try:
@@ -737,7 +869,14 @@ class ToonCrafterInterpolation:
fs = torch.tensor([fs], dtype=torch.float32, device=self.model.device) fs = torch.tensor([fs], dtype=torch.float32, device=self.model.device)
else: else:
fs = torch.tensor([fs], dtype=torch.float64, device=self.model.device) fs = torch.tensor([fs], dtype=torch.float64, device=self.model.device)
cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat]}
if controlnet is not None:
cn_videos = controlnet["cn_videos"]
cn_videos = cn_videos.to(dtype).to(device)
else:
cn_videos = None
cond = {"c_crossattn": [imtext_cond], "fs": fs, "c_concat": [img_tensor_repeat], "control_cond": cn_videos}
if noise_shape[-1] == 32: if noise_shape[-1] == 32:
timestep_spacing = "uniform" timestep_spacing = "uniform"
@@ -820,6 +959,41 @@ class ToonCrafterInterpolation:
return (latent,) return (latent,)
class DynamiCrafterControlnetApply:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"controlnet": ("DC_CN_MODEL",),
"images": ("IMAGE",),
"control_scale": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 1.0, "step": 0.01}),
}
}
RETURN_TYPES = ("DC_CONTROL",)
RETURN_NAMES = ("controlnet",)
FUNCTION = "process"
CATEGORY = "DynamiCrafterWrapper"
def process(self, controlnet, images, control_scale):
controlnet['model'].control_scale = control_scale
#images = images * 2.0 - 1.0
cn_tensor = images.permute(3, 0, 1, 2).unsqueeze(0)
print("control frame: ", cn_tensor.shape) # b c t h w
print(controlnet["config"])
if controlnet["config"]["hint_channels"] == 1:
cn_tensor = cn_tensor[:, :1, :, :, :]
print("control frame: ", cn_tensor.shape) # b c t h w
controlnet = {
"model": controlnet['model'],
"cn_videos": cn_tensor,
}
return (controlnet,)
class ToonCrafterDecode: class ToonCrafterDecode:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
@@ -952,8 +1126,9 @@ class DynamiCrafterBatchInterpolation:
mm.soft_empty_cache() mm.soft_empty_cache()
torch.manual_seed(seed) torch.manual_seed(seed)
dtype = model.dtype
self.model = model['model'] self.model = model['model']
dtype = self.model.dtype
if vae_dtype == "auto": if vae_dtype == "auto":
try: try:
@@ -968,7 +1143,6 @@ class DynamiCrafterBatchInterpolation:
print(f"VAE using dtype: {self.model.first_stage_model.dtype}") print(f"VAE using dtype: {self.model.first_stage_model.dtype}")
self.model.to(device) self.model.to(device)
images = images * 2 - 1
images = images.permute(0, 3, 1, 2).to(dtype).to(device) images = images.permute(0, 3, 1, 2).to(dtype).to(device)
B, C, H, W = images.shape B, C, H, W = images.shape
orig_H, orig_W = H, W orig_H, orig_W = H, W
@@ -991,8 +1165,11 @@ class DynamiCrafterBatchInterpolation:
self.model.first_stage_model.to(device) self.model.first_stage_model.to(device)
z = get_latent_z(self.model, image.unsqueeze(2)) #bc,1,hw encode_pixels1 = image * 2 - 1
z2 = get_latent_z(self.model, image2.unsqueeze(2)) #bc,1,hw encode_pixels2 = image2 * 2 - 1
z = get_latent_z(self.model, encode_pixels1.unsqueeze(2)) #bc,1,hw
z2 = get_latent_z(self.model, encode_pixels2.unsqueeze(2)) #bc,1,hw
img_tensor_repeat = repeat(z, 'b c t h w -> b c (repeat t) h w', repeat=frames) img_tensor_repeat = repeat(z, 'b c t h w -> b c (repeat t) h w', repeat=frames)
img_tensor_repeat = torch.zeros_like(img_tensor_repeat) img_tensor_repeat = torch.zeros_like(img_tensor_repeat)
img_tensor_repeat[:,:,:1,:,:] = z img_tensor_repeat[:,:,:1,:,:] = z
@@ -1009,7 +1186,9 @@ class DynamiCrafterBatchInterpolation:
imtext_cond = torch.cat([text_emb, img_emb], dim=1) imtext_cond = torch.cat([text_emb, img_emb], dim=1)
fs = torch.tensor([fs], dtype=torch.long, device=self.model.device) fs = torch.tensor([fs], dtype=torch.long, device=self.model.device)
cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat]} cond = {"c_crossattn": [imtext_cond], "c_concat": [img_tensor_repeat], "control_cond": None}
self.model.control_model = None
if noise_shape[-1] == 32: if noise_shape[-1] == 32:
timestep_spacing = "uniform" timestep_spacing = "uniform"
@@ -1117,17 +1296,23 @@ NODE_CLASS_MAPPINGS = {
"DownloadAndLoadDynamiCrafterModel": DownloadAndLoadDynamiCrafterModel, "DownloadAndLoadDynamiCrafterModel": DownloadAndLoadDynamiCrafterModel,
"DownloadAndLoadCLIPModel": DownloadAndLoadCLIPModel, "DownloadAndLoadCLIPModel": DownloadAndLoadCLIPModel,
"DownloadAndLoadCLIPVisionModel": DownloadAndLoadCLIPVisionModel, "DownloadAndLoadCLIPVisionModel": DownloadAndLoadCLIPVisionModel,
"DynamiCrafterLoadInitNoise": DynamiCrafterLoadInitNoise "DynamiCrafterLoadInitNoise": DynamiCrafterLoadInitNoise,
"DownloadAndLoadDynamiCrafterCNModel": DownloadAndLoadDynamiCrafterCNModel,
"DynamiCrafterControlnetApply": DynamiCrafterControlnetApply,
"DynamiCrafterCNLoader": DynamiCrafterCNLoader
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"DynamiCrafterI2V": "DynamiCrafterI2V", "DynamiCrafterI2V": "DynamiCrafterI2V",
"DynamiCrafterModelLoader": "DynamiCrafterModelLoader", "DynamiCrafterModelLoader": "DynamiCrafter ModelLoader",
"DynamiCrafterBatchInterpolation": "DynamiCrafterBatchInterpolation", "DynamiCrafterBatchInterpolation": "DynamiCrafter BatchInterpolation",
"ToonCrafterInterpolation": "ToonCrafterInterpolation", "ToonCrafterInterpolation": "ToonCrafter Interpolation",
"ToonCrafterDecode": "ToonCrafterDecode", "ToonCrafterDecode": "ToonCrafter Decode",
"DownloadAndLoadDynamiCrafterModel": "DownloadAndLoadDynamiCrafterModel", "DownloadAndLoadDynamiCrafterModel": "(Down)Load DynamiCrafterModel",
"DownloadAndLoadCLIPModel": "DownloadAndLoadCLIPModel", "DownloadAndLoadCLIPModel": "(Down)Load CLIPModel",
"DownloadAndLoadCLIPVisionModel": "DownloadAndLoadCLIPVisionModel", "DownloadAndLoadCLIPVisionModel": "(Down)Load CLIPVisionModel",
"DynamiCrafterLoadInitNoise": "DynamiCrafterLoadInitNoise" "DynamiCrafterLoadInitNoise": "DynamiCrafter LoadInitNoise",
"DownloadAndLoadDynamiCrafterCNModel": "(Down)Load DynamiCrafter CNModel",
"DynamiCrafterControlnetApply": "DynamiCrafter ControlnetApply",
"DynamiCrafterCNLoader": "DynamiCrafter CNLoader"
} }
+1 -1
View File
@@ -1,7 +1,7 @@
[project] [project]
name = "comfyui-dynamicrafterwrapper" name = "comfyui-dynamicrafterwrapper"
description = "Wrapper nodes to use Dynami/ToonCrafter image2video and frame interpolation models in ComfyUI" description = "Wrapper nodes to use Dynami/ToonCrafter image2video and frame interpolation models in ComfyUI"
version = "1.0.2" version = "1.0.3"
license = "Apache-2.0" license = "Apache-2.0"
dependencies = ["einops>=0.3.0", "numpy>=1.24.2", "omegaconf>=2.1.1", "pytorch_lightning>=2.2.1", "tqdm>=4.65.0", "transformers>=4.25.1", "timm"] dependencies = ["einops>=0.3.0", "numpy>=1.24.2", "omegaconf>=2.1.1", "pytorch_lightning>=2.2.1", "tqdm>=4.65.0", "transformers>=4.25.1", "timm"]
+5 -2
View File
@@ -34,12 +34,15 @@ def instantiate_from_config(config):
def get_obj_from_str(string, reload=False): def get_obj_from_str(string, reload=False):
package_directory_name = os.path.basename(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
module, cls = string.rsplit(".", 1) module, cls = string.rsplit(".", 1)
if reload: if reload:
module_imp = importlib.import_module(module) module_imp = importlib.import_module(module)
importlib.reload(module_imp) importlib.reload(module_imp)
return getattr(importlib.import_module(module, package=package_directory_name), cls) try:
obj = getattr(importlib.import_module(module, package=os.path.basename(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))), cls)
except:
obj = getattr(importlib.import_module(module, package=os.path.dirname(os.path.dirname(os.path.abspath( __file__ )))), cls)
return obj
def load_npz_from_dir(data_dir): def load_npz_from_dir(data_dir):