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,
rand_cond_frame=False,
en_and_decode_n_samples_a_time=None,
control_scale=1.0,
*args, **kwargs):
self.num_timesteps_cond = default(num_timesteps_cond, 1)
self.scale_by_std = scale_by_std
@@ -404,6 +405,7 @@ class LatentDiffusion(DDPM):
self.loop_video = loop_video
self.fps_condition_type = fps_condition_type
self.perframe_ae = perframe_ae
self.control_scale = control_scale
self.logdir = logdir
self.rand_cond_frame = rand_cond_frame
@@ -572,9 +574,20 @@ class LatentDiffusion(DDPM):
if not isinstance(cond, list):
cond = [cond]
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):
return x_recon[0]
@@ -722,7 +735,7 @@ class DiffusionWrapper(pl.LightningModule):
self.diffusion_model = instantiate_from_config(diff_model_config)
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):
# temporal_context = fps is foNone
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
xc = torch.cat([x] + c_concat, dim=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':
cc = c_crossattn[0]
out = self.diffusion_model(x, t, context=cc)
+405 -2
View File
@@ -2,7 +2,9 @@ from functools import partial
from abc import abstractmethod
import torch
import torch.nn as nn
import numpy as np
from einops import rearrange
import math
import torch.nn.functional as F
from ....lvdm.models.utils_diffusion import timestep_embedding
from ....lvdm.common import checkpoint
@@ -18,6 +20,9 @@ from ....lvdm.modules.attention import SpatialTransformer, TemporalTransformer
import comfy.ops
ops = comfy.ops.manual_cast
def exists(x):
return x is not None
class TimestepBlock(nn.Module):
"""
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)),
)
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
t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).type(x.dtype)
emb = self.time_embed(t_emb)
@@ -594,8 +599,15 @@ class UNetModel(nn.Module):
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)
if control is not None:
h += control.pop()
for module in self.output_blocks:
h = torch.cat([h, hs.pop()], dim=1)
if control is None:
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 = h.type(x.dtype)
y = self.out(h)
@@ -603,3 +615,394 @@ class UNetModel(nn.Module):
# reshape back to (b c t h w)
y = rearrange(y, '(b t) c h w -> b c t h w', b=b)
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
from contextlib import nullcontext
from .lvdm.models.samplers.ddim import DDIMSampler
from .lvdm.modules.networks.openaimodel3d import ControlNet
from contextlib import nullcontext
try:
@@ -80,7 +81,7 @@ class DownloadAndLoadDynamiCrafterModel:
'ckpt_name': model,
'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
download_path = os.path.join(folder_paths.models_dir, "checkpoints", "dynamicrafter")
model_path = os.path.join(download_path, model)
@@ -140,8 +141,130 @@ class DownloadAndLoadDynamiCrafterModel:
dcmodel = {
'model': self.model,
'model_name': model,
'dtype': precision
}
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:
@classmethod
@@ -389,6 +512,7 @@ class DynamiCrafterI2V:
self.model.to(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():
image = image.permute(0, 3, 1, 2).to(dtype).to(device)
if augmentation_level > 0:
image += torch.randn_like(image) * augmentation_level
@@ -400,7 +524,7 @@ class DynamiCrafterI2V:
if H % 64 != 0:
H = H - (H % 64)
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
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
if image2 is not None:
image2 = image2 * 2 - 1
image2 = image2.permute(0, 3, 1, 2).to(dtype).to(device)
if augmentation_level > 0:
image2 += torch.randn_like(image2) * augmentation_level
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
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
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:
timestep_spacing = "uniform"
@@ -633,6 +758,7 @@ class ToonCrafterInterpolation:
"augmentation_level": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.0001}),
"optional_latents": ("LATENT",),
"ddpm_from": ("INT", {"default": 1000, "min": 1, "max": 1000, "step": 1}),
"controlnet": ("DC_CONTROL",),
}
}
@@ -641,7 +767,8 @@ class ToonCrafterInterpolation:
FUNCTION = "process"
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()
offload_device = mm.unet_offload_device()
mm.unload_all_models()
@@ -651,6 +778,11 @@ class ToonCrafterInterpolation:
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
if vae_dtype == "auto":
try:
@@ -737,7 +869,14 @@ class ToonCrafterInterpolation:
fs = torch.tensor([fs], dtype=torch.float32, device=self.model.device)
else:
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:
timestep_spacing = "uniform"
@@ -819,6 +958,41 @@ class ToonCrafterInterpolation:
}
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:
@classmethod
@@ -952,8 +1126,9 @@ class DynamiCrafterBatchInterpolation:
mm.soft_empty_cache()
torch.manual_seed(seed)
dtype = model.dtype
self.model = model['model']
dtype = self.model.dtype
if vae_dtype == "auto":
try:
@@ -968,7 +1143,6 @@ class DynamiCrafterBatchInterpolation:
print(f"VAE using dtype: {self.model.first_stage_model.dtype}")
self.model.to(device)
images = images * 2 - 1
images = images.permute(0, 3, 1, 2).to(dtype).to(device)
B, C, H, W = images.shape
orig_H, orig_W = H, W
@@ -991,8 +1165,11 @@ class DynamiCrafterBatchInterpolation:
self.model.first_stage_model.to(device)
z = get_latent_z(self.model, image.unsqueeze(2)) #bc,1,hw
z2 = get_latent_z(self.model, image2.unsqueeze(2)) #bc,1,hw
encode_pixels1 = image * 2 - 1
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 = torch.zeros_like(img_tensor_repeat)
img_tensor_repeat[:,:,:1,:,:] = z
@@ -1009,7 +1186,9 @@ class DynamiCrafterBatchInterpolation:
imtext_cond = torch.cat([text_emb, img_emb], dim=1)
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:
timestep_spacing = "uniform"
@@ -1117,17 +1296,23 @@ NODE_CLASS_MAPPINGS = {
"DownloadAndLoadDynamiCrafterModel": DownloadAndLoadDynamiCrafterModel,
"DownloadAndLoadCLIPModel": DownloadAndLoadCLIPModel,
"DownloadAndLoadCLIPVisionModel": DownloadAndLoadCLIPVisionModel,
"DynamiCrafterLoadInitNoise": DynamiCrafterLoadInitNoise
"DynamiCrafterLoadInitNoise": DynamiCrafterLoadInitNoise,
"DownloadAndLoadDynamiCrafterCNModel": DownloadAndLoadDynamiCrafterCNModel,
"DynamiCrafterControlnetApply": DynamiCrafterControlnetApply,
"DynamiCrafterCNLoader": DynamiCrafterCNLoader
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DynamiCrafterI2V": "DynamiCrafterI2V",
"DynamiCrafterModelLoader": "DynamiCrafterModelLoader",
"DynamiCrafterBatchInterpolation": "DynamiCrafterBatchInterpolation",
"ToonCrafterInterpolation": "ToonCrafterInterpolation",
"ToonCrafterDecode": "ToonCrafterDecode",
"DownloadAndLoadDynamiCrafterModel": "DownloadAndLoadDynamiCrafterModel",
"DownloadAndLoadCLIPModel": "DownloadAndLoadCLIPModel",
"DownloadAndLoadCLIPVisionModel": "DownloadAndLoadCLIPVisionModel",
"DynamiCrafterLoadInitNoise": "DynamiCrafterLoadInitNoise"
"DynamiCrafterModelLoader": "DynamiCrafter ModelLoader",
"DynamiCrafterBatchInterpolation": "DynamiCrafter BatchInterpolation",
"ToonCrafterInterpolation": "ToonCrafter Interpolation",
"ToonCrafterDecode": "ToonCrafter Decode",
"DownloadAndLoadDynamiCrafterModel": "(Down)Load DynamiCrafterModel",
"DownloadAndLoadCLIPModel": "(Down)Load CLIPModel",
"DownloadAndLoadCLIPVisionModel": "(Down)Load CLIPVisionModel",
"DynamiCrafterLoadInitNoise": "DynamiCrafter LoadInitNoise",
"DownloadAndLoadDynamiCrafterCNModel": "(Down)Load DynamiCrafter CNModel",
"DynamiCrafterControlnetApply": "DynamiCrafter ControlnetApply",
"DynamiCrafterCNLoader": "DynamiCrafter CNLoader"
}
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-dynamicrafterwrapper"
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"
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):
package_directory_name = os.path.basename(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
module, cls = string.rsplit(".", 1)
if reload:
module_imp = importlib.import_module(module)
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):