Initial sketch controlnet support
This commit is contained in:
+17
-4
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
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])
|
||||
@@ -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:
|
||||
@@ -140,9 +141,70 @@ 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):
|
||||
device = mm.get_torch_device()
|
||||
mm.soft_empty_cache()
|
||||
custom_config = {
|
||||
'ckpt_name': model,
|
||||
}
|
||||
if not hasattr(self, 'model') or self.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": 1,
|
||||
"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
|
||||
}
|
||||
|
||||
cn_model = ControlNet(**cn_config)
|
||||
print("Loading ControlNet")
|
||||
cn_sd = comfy.utils.load_torch_file(cn_model_path)
|
||||
cn_model.load_state_dict(cn_sd, strict=True)
|
||||
print("ControlNet loaded")
|
||||
|
||||
return (cn_model,)
|
||||
|
||||
class DownloadAndLoadCLIPModel:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -633,6 +695,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 +704,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 +715,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 +806,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"
|
||||
@@ -820,6 +896,39 @@ class ToonCrafterInterpolation:
|
||||
|
||||
return (latent,)
|
||||
|
||||
class DynamiCrafterControlnetApply:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"model": ("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, model, images, control_scale):
|
||||
|
||||
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
|
||||
cn_tensor = cn_tensor[:, :1, :, :, :]
|
||||
print("control frame: ", cn_tensor.shape) # b c t h w
|
||||
|
||||
controlnet = {
|
||||
"model": model,
|
||||
"cn_videos": cn_tensor,
|
||||
}
|
||||
|
||||
return (controlnet,)
|
||||
|
||||
class ToonCrafterDecode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -1120,7 +1229,9 @@ NODE_CLASS_MAPPINGS = {
|
||||
"DownloadAndLoadDynamiCrafterModel": DownloadAndLoadDynamiCrafterModel,
|
||||
"DownloadAndLoadCLIPModel": DownloadAndLoadCLIPModel,
|
||||
"DownloadAndLoadCLIPVisionModel": DownloadAndLoadCLIPVisionModel,
|
||||
"DynamiCrafterLoadInitNoise": DynamiCrafterLoadInitNoise
|
||||
"DynamiCrafterLoadInitNoise": DynamiCrafterLoadInitNoise,
|
||||
"DownloadAndLoadDynamiCrafterCNModel": DownloadAndLoadDynamiCrafterCNModel,
|
||||
"DynamiCrafterControlnetApply": DynamiCrafterControlnetApply
|
||||
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -1129,8 +1240,10 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"DynamiCrafterBatchInterpolation": "DynamiCrafter BatchInterpolation",
|
||||
"ToonCrafterInterpolation": "ToonCrafter Interpolation",
|
||||
"ToonCrafterDecode": "ToonCrafter Decode",
|
||||
"DownloadAndLoadDynamiCrafterModel": "DownloadAndLoadDynamiCrafterModel",
|
||||
"DownloadAndLoadCLIPModel": "DownloadAndLoadCLIPModel",
|
||||
"DownloadAndLoadCLIPVisionModel": "DownloadAndLoadCLIPVisionModel",
|
||||
"DynamiCrafterLoadInitNoise": "DynamiCrafterLoadInitNoise"
|
||||
"DownloadAndLoadDynamiCrafterModel": "(Down)Load DynamiCrafterModel",
|
||||
"DownloadAndLoadCLIPModel": "(Down)Load CLIPModel",
|
||||
"DownloadAndLoadCLIPVisionModel": "(Down)Load CLIPVisionModel",
|
||||
"DynamiCrafterLoadInitNoise": "DynamiCrafter LoadInitNoise",
|
||||
"DownloadAndLoadDynamiCrafterCNModel": "(Down)Load DynamiCrafter CNModel",
|
||||
"DynamiCrafterControlnetApply": "DynamiCrafter ControlnetApply"
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user