Compare commits
8
Commits
freetraj_testing
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d312c62982 | ||
|
|
7670b09c28 | ||
|
|
282c69c5f3 | ||
|
|
18121e97a2 | ||
|
|
4a7dd402b4 | ||
|
|
2545a3afb7 | ||
|
|
0668ee6d0a | ||
|
|
69fe0e2e9f |
+17
-4
@@ -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)
|
||||||
|
|||||||
@@ -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])
|
||||||
@@ -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
@@ -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
@@ -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):
|
||||||
|
|||||||
Reference in New Issue
Block a user