Merge branch 'main' into develop

This commit is contained in:
kijai
2024-08-16 00:15:38 +03:00
26 changed files with 2093 additions and 290 deletions
+21
View File
@@ -0,0 +1,21 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
paths:
- "pyproject.toml"
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+6
View File
@@ -1,4 +1,10 @@
## DynamiCrafter wrapper nodes for ComfyUI
## Update2: Refactor
Changed lots of things to better integrate this to ComfyUI, you can (and have to) use clip_vision and clip models, but memory usage is much better and I was able to do 512x320 under 10GB VRAM.
New example workflows are included, all old workflows will have to be updated.
## Update: ToonCrafter
Initial ToonCrafter support with it's own node.
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -431,7 +431,7 @@
"Node name for S&R": "DownloadAndLoadDynamiCrafterModel"
},
"widgets_values": [
"dynamicrafter_1024_v1_bf16.safetensors",
"dynamicrafter_1024_fp16_pruned.safetensors",
"auto",
true
]
+1 -1
View File
@@ -617,7 +617,7 @@
"Node name for S&R": "DownloadAndLoadDynamiCrafterModel"
},
"widgets_values": [
"tooncrafter_512_interp-fp16.safetensors",
"tooncrafter_512_interp-pruned-fp16.safetensors",
"auto",
false
]
@@ -549,7 +549,7 @@
"Node name for S&R": "DownloadAndLoadDynamiCrafterModel"
},
"widgets_values": [
"tooncrafter_512_interp-fp16.safetensors",
"tooncrafter_512_interp-pruned-fp16.safetensors",
"auto",
false
]
Binary file not shown.
Binary file not shown.
+6 -5
View File
@@ -9,7 +9,8 @@
import torch.nn as nn
from ..utils.utils import instantiate_from_config
import comfy.ops
ops = comfy.ops.manual_cast
def disabled_train(self, mode=True):
"""Overwrite model.train with this function to make sure train/eval mode
@@ -38,11 +39,11 @@ def conv_nd(dims, *args, **kwargs):
Create a 1D, 2D, or 3D convolution module.
"""
if dims == 1:
return nn.Conv1d(*args, **kwargs)
return ops.Conv1d(*args, **kwargs)
elif dims == 2:
return nn.Conv2d(*args, **kwargs)
return ops.Conv2d(*args, **kwargs)
elif dims == 3:
return nn.Conv3d(*args, **kwargs)
return ops.Conv3d(*args, **kwargs)
raise ValueError(f"unsupported dimensions: {dims}")
@@ -50,7 +51,7 @@ def linear(*args, **kwargs):
"""
Create a linear module.
"""
return nn.Linear(*args, **kwargs)
return ops.Linear(*args, **kwargs)
def avg_pool_nd(dims, *args, **kwargs):
+4 -3
View File
@@ -8,7 +8,8 @@ import pytorch_lightning as pl
from ...lvdm.modules.networks.ae_modules import Encoder, Decoder
from ...lvdm.distributions import DiagonalGaussianDistribution
from ...utils.utils import instantiate_from_config
import comfy.ops
ops = comfy.ops.manual_cast
TIMESTEPS=16
class AutoencoderKL(pl.LightningModule):
def __init__(self,
@@ -34,8 +35,8 @@ class AutoencoderKL(pl.LightningModule):
self.decoder = Decoder(**ddconfig)
self.loss = instantiate_from_config(lossconfig)
assert ddconfig["double_z"]
self.quant_conv = torch.nn.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1)
self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1)
self.quant_conv = ops.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1)
self.post_quant_conv = ops.Conv2d(embed_dim, ddconfig["z_channels"], 1)
self.embed_dim = embed_dim
self.input_dim = input_dim
self.test = test
+129 -36
View File
@@ -9,17 +9,23 @@ import torch.nn as nn
from packaging import version
logpy = logging.getLogger(__name__)
try:
import xformers
import xformers.ops
import comfy.model_management
if comfy.model_management.XFORMERS_IS_AVAILABLE:
try:
import xformers
import xformers.ops
XFORMERS_IS_AVAILABLE = True
except:
XFORMERS_IS_AVAILABLE = True
except:
XFORMERS_IS_AVAILABLE = False
logpy.warning("no module 'xformers'. Processing without...")
else:
XFORMERS_IS_AVAILABLE = False
logpy.warning("no module 'xformers'. Processing without...")
from ...lvdm.modules.attention_svd import LinearAttention, MemoryEfficientCrossAttention
from ...lvdm.modules.attention_svd import LinearAttention, MemoryEfficientCrossAttention, CrossAttention
import comfy.ops
ops = comfy.ops.manual_cast
def nonlinearity(x):
# swish
@@ -27,7 +33,7 @@ def nonlinearity(x):
def Normalize(in_channels, num_groups=32):
return torch.nn.GroupNorm(
return ops.GroupNorm(
num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True
)
@@ -49,23 +55,23 @@ class ResnetBlock(nn.Module):
self.use_conv_shortcut = conv_shortcut
self.norm1 = Normalize(in_channels)
self.conv1 = torch.nn.Conv2d(
self.conv1 = ops.Conv2d(
in_channels, out_channels, kernel_size=3, stride=1, padding=1
)
if temb_channels > 0:
self.temb_proj = torch.nn.Linear(temb_channels, out_channels)
self.temb_proj = ops.Linear(temb_channels, out_channels)
self.norm2 = Normalize(out_channels)
self.dropout = torch.nn.Dropout(dropout)
self.conv2 = torch.nn.Conv2d(
self.conv2 = ops.Conv2d(
out_channels, out_channels, kernel_size=3, stride=1, padding=1
)
if self.in_channels != self.out_channels:
if self.use_conv_shortcut:
self.conv_shortcut = torch.nn.Conv2d(
self.conv_shortcut = ops.Conv2d(
in_channels, out_channels, kernel_size=3, stride=1, padding=1
)
else:
self.nin_shortcut = torch.nn.Conv2d(
self.nin_shortcut = ops.Conv2d(
in_channels, out_channels, kernel_size=1, stride=1, padding=0
)
@@ -105,16 +111,16 @@ class AttnBlock(nn.Module):
self.in_channels = in_channels
self.norm = Normalize(in_channels)
self.q = torch.nn.Conv2d(
self.q = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.k = torch.nn.Conv2d(
self.k = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.v = torch.nn.Conv2d(
self.v = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.proj_out = torch.nn.Conv2d(
self.proj_out = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
@@ -155,16 +161,16 @@ class MemoryEfficientAttnBlock(nn.Module):
self.in_channels = in_channels
self.norm = Normalize(in_channels)
self.q = torch.nn.Conv2d(
self.q = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.k = torch.nn.Conv2d(
self.k = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.v = torch.nn.Conv2d(
self.v = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.proj_out = torch.nn.Conv2d(
self.proj_out = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.attention_op: Optional[Any] = None
@@ -206,6 +212,14 @@ class MemoryEfficientAttnBlock(nn.Module):
return x + h_
class CrossAttentionWrapper(CrossAttention):
def forward(self, x, context=None, mask=None, **unused_kwargs):
b, c, h, w = x.shape
x = rearrange(x, "b c h w -> b 1 (h w) c").contiguous()
out = super().forward(x, context=context, mask=mask)
out = rearrange(out, "b 1 (h w) c -> b c h w", h=h, w=w, c=c, b=b)
return x + out
class MemoryEfficientCrossAttentionWrapper(MemoryEfficientCrossAttention):
def forward(self, x, context=None, mask=None, **unused_kwargs):
b, c, h, w = x.shape
@@ -220,9 +234,11 @@ def make_attn(in_channels, attn_type="vanilla", attn_kwargs=None):
assert attn_type in [
"vanilla",
"vanilla-xformers",
"cross-attn",
"memory-efficient-cross-attn",
"linear",
"none",
"cross-attn-fusion",
"memory-efficient-cross-attn-fusion",
], f"attn_type {attn_type} unknown"
if (
@@ -243,9 +259,15 @@ def make_attn(in_channels, attn_type="vanilla", attn_kwargs=None):
f"building MemoryEfficientAttnBlock with {in_channels} in_channels..."
)
return MemoryEfficientAttnBlock(in_channels)
elif attn_type == "cross-attn":
attn_kwargs["query_dim"] = in_channels
return CrossAttentionWrapper(**attn_kwargs)
elif attn_type == "memory-efficient-cross-attn":
attn_kwargs["query_dim"] = in_channels
return MemoryEfficientCrossAttentionWrapper(**attn_kwargs)
elif attn_type == "cross-attn-fusion":
attn_kwargs["query_dim"] = in_channels
return CrossAttentionWrapperFusion(**attn_kwargs)
elif attn_type == "memory-efficient-cross-attn-fusion":
attn_kwargs["query_dim"] = in_channels
return MemoryEfficientCrossAttentionWrapperFusion(**attn_kwargs)
@@ -254,6 +276,76 @@ def make_attn(in_channels, attn_type="vanilla", attn_kwargs=None):
else:
return LinAttnBlock(in_channels)
class CrossAttentionWrapperFusion(CrossAttention):
def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0, **kwargs):
super().__init__(query_dim, context_dim, heads, dim_head, dropout, **kwargs)
self.dim_head = dim_head
self.norm = Normalize(query_dim)
nn.init.zeros_(self.to_out[0].weight)
nn.init.zeros_(self.to_out[0].bias)
def forward(self, x, context=None, mask=None):
if self.training:
return checkpoint(self._forward, x, context, mask, use_reentrant=False)
else:
return self._forward(x, context, mask)
def _forward(
self,
x,
context=None,
mask=None,
):
bt, c, h, w = x.shape
h_ = self.norm(x)
h_ = rearrange(h_, "b c h w -> b (h w) c")
q = self.to_q(h_)
b, c, l, h, w = context.shape
context = rearrange(context, "b c l h w -> (b l) (h w) c")
k = self.to_k(context)
v = self.to_v(context)
k = rearrange(k, "(b l) d c -> b l d c", l=l)
k = torch.cat([k[:, [0] * (bt // b)], k[:, [1] * (bt // b)]], dim=2)
k = rearrange(k, "b l d c -> (b l) d c")
v = rearrange(v, "(b l) d c -> b l d c", l=l)
v = torch.cat([v[:, [0] * (bt // b)], v[:, [1] * (bt // b)]], dim=2)
v = rearrange(v, "b l d c -> (b l) d c")
b, _, _ = q.shape # actually bt
q, k, v = map(
lambda t: t.unsqueeze(3)
.reshape(b, t.shape[1], self.heads, self.dim_head)
.permute(0, 2, 1, 3)
.reshape(b * self.heads, t.shape[1], self.dim_head)
.contiguous(),
(q, k, v),
)
sdpa = torch.nn.functional.scaled_dot_product_attention
def slow_sdpa(q, k, v):
out_list = []
step = 10
for i in range(0, q.shape[0], step):
out_i = sdpa(q[i:i + step], k[i:i + step], v[i:i + step])
out_list.append(out_i)
return torch.cat(out_list, dim=0)
out = slow_sdpa(q, k, v)
out = (
out.unsqueeze(0)
.reshape(b, self.heads, out.shape[1], self.dim_head)
.permute(0, 2, 1, 3)
.reshape(b, out.shape[1], self.heads * self.dim_head)
)
out = self.to_out(out)
out = rearrange(out, "bt (h w) c -> bt c h w", h=h, w=w, c=c)
return x + out
class MemoryEfficientCrossAttentionWrapperFusion(MemoryEfficientCrossAttention):
# print('x.shape: ',x.shape, 'context.shape: ',context.shape) ##torch.Size([8, 128, 256, 256]) torch.Size([1, 128, 2, 256, 256])
def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0, **kwargs):
@@ -344,7 +436,7 @@ class MemoryEfficientCrossAttentionWrapperFusion(MemoryEfficientCrossAttention):
class Combiner(nn.Module):
def __init__(self, ch) -> None:
super().__init__()
self.conv = nn.Conv2d(ch,ch,1,padding=0)
self.conv = ops.Conv2d(ch,ch,1,padding=0)
nn.init.zeros_(self.conv.weight)
nn.init.zeros_(self.conv.bias)
@@ -417,7 +509,7 @@ class Decoder(nn.Module):
make_resblock_cls = self._make_resblock()
make_conv_cls = self._make_conv()
# z to block_in
self.conv_in = torch.nn.Conv2d(
self.conv_in = ops.Conv2d(
z_channels, block_in, kernel_size=3, stride=1, padding=1
)
@@ -465,7 +557,8 @@ class Decoder(nn.Module):
self.up.insert(0, up) # prepend to get consistent order
if i_level in self.attn_level:
self.attn_refinement.insert(0, make_attn_cls(block_in, attn_type='memory-efficient-cross-attn-fusion', attn_kwargs={}))
_attn_type = 'memory-efficient-cross-attn-fusion' if XFORMERS_IS_AVAILABLE else 'cross-attn-fusion'
self.attn_refinement.insert(0, make_attn_cls(block_in, attn_type=_attn_type, attn_kwargs={}))
else:
self.attn_refinement.insert(0, Combiner(block_in))
# end
@@ -482,7 +575,7 @@ class Decoder(nn.Module):
return ResnetBlock
def _make_conv(self) -> Callable:
return torch.nn.Conv2d
return ops.Conv2d
def get_last_layer(self, **kwargs):
return self.conv_out.weight
@@ -737,7 +830,7 @@ class VideoTransformerBlock(nn.Module):
self.is_res = inner_dim == dim
if self.ff_in:
self.norm_in = nn.LayerNorm(dim)
self.norm_in = ops.LayerNorm(dim)
self.ff_in = FeedForward(
dim, dim_out=inner_dim, dropout=dropout, glu=gated_ff
)
@@ -765,7 +858,7 @@ class VideoTransformerBlock(nn.Module):
else:
self.attn2 = None
else:
self.norm2 = nn.LayerNorm(inner_dim)
self.norm2 = ops.LayerNorm(inner_dim)
if switch_temporal_ca_to_sa:
self.attn2 = attn_cls(
query_dim=inner_dim, heads=n_heads, dim_head=d_head, dropout=dropout
@@ -779,8 +872,8 @@ class VideoTransformerBlock(nn.Module):
dropout=dropout,
) # is self-attn if context is none
self.norm1 = nn.LayerNorm(inner_dim)
self.norm3 = nn.LayerNorm(inner_dim)
self.norm1 = ops.LayerNorm(inner_dim)
self.norm3 = ops.LayerNorm(inner_dim)
self.switch_temporal_ca_to_sa = switch_temporal_ca_to_sa
self.checkpoint = checkpoint
@@ -912,7 +1005,7 @@ class VideoResBlock(ResnetBlock):
return x
class AE3DConv(torch.nn.Conv2d):
class AE3DConv(ops.Conv2d):
def __init__(self, in_channels, out_channels, video_kernel_size=3, *args, **kwargs):
super().__init__(in_channels, out_channels, *args, **kwargs)
if isinstance(video_kernel_size, Iterable):
@@ -920,7 +1013,7 @@ class AE3DConv(torch.nn.Conv2d):
else:
padding = int(video_kernel_size // 2)
self.time_mix_conv = torch.nn.Conv3d(
self.time_mix_conv = ops.Conv3d(
in_channels=out_channels,
out_channels=out_channels,
kernel_size=video_kernel_size,
@@ -953,9 +1046,9 @@ class VideoBlock(AttnBlock):
time_embed_dim = self.in_channels * 4
self.video_time_embed = torch.nn.Sequential(
torch.nn.Linear(self.in_channels, time_embed_dim),
ops.Linear(self.in_channels, time_embed_dim),
torch.nn.SiLU(),
torch.nn.Linear(time_embed_dim, self.in_channels),
ops.Linear(time_embed_dim, self.in_channels),
)
self.merge_strategy = merge_strategy
@@ -1023,9 +1116,9 @@ class MemoryEfficientVideoBlock(MemoryEfficientAttnBlock):
time_embed_dim = self.in_channels * 4
self.video_time_embed = torch.nn.Sequential(
torch.nn.Linear(self.in_channels, time_embed_dim),
ops.Linear(self.in_channels, time_embed_dim),
torch.nn.SiLU(),
torch.nn.Linear(time_embed_dim, self.in_channels),
ops.Linear(time_embed_dim, self.in_channels),
)
self.merge_strategy = merge_strategy
@@ -1114,7 +1207,7 @@ def make_time_attn(
return NotImplementedError()
class Conv2DWrapper(torch.nn.Conv2d):
class Conv2DWrapper(ops.Conv2d):
def forward(self, input: torch.Tensor, **kwargs) -> torch.Tensor:
return super().forward(input)
+5 -2
View File
@@ -9,6 +9,9 @@ from ...lvdm.modules.networks.ae_modules import Encoder, Decoder
from ...lvdm.distributions import DiagonalGaussianDistribution
from ...utils.utils import instantiate_from_config
import comfy.ops
ops = comfy.ops.manual_cast
TIMESTEPS=16
class AutoencoderKL(pl.LightningModule):
def __init__(self,
@@ -31,8 +34,8 @@ class AutoencoderKL(pl.LightningModule):
self.decoder = Decoder(**ddconfig)
self.loss = instantiate_from_config(lossconfig)
assert ddconfig["double_z"]
self.quant_conv = torch.nn.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1)
self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1)
self.quant_conv = ops.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1)
self.post_quant_conv = ops.Conv2d(embed_dim, ddconfig["z_channels"], 1)
self.embed_dim = embed_dim
self.input_dim = input_dim
self.test = test
+35 -17
View File
@@ -33,6 +33,8 @@ from ...lvdm.models.autoencoder_dualref import VideoDecoder
__conditioning_keys__ = {'concat': 'c_concat',
'crossattn': 'c_crossattn',
'adm': 'y'}
import comfy.model_management as mm
device = mm.get_torch_device()
class DDPM(pl.LightningModule):
# classic DDPM with Gaussian diffusion, in image space
@@ -220,6 +222,9 @@ class DDPM(pl.LightningModule):
variance = extract_into_tensor(1.0 - self.alphas_cumprod, t, x_start.shape)
log_variance = extract_into_tensor(self.log_one_minus_alphas_cumprod, t, x_start.shape)
return mean, variance, log_variance
def get_sqrt_alpha_t_bar(self,x_start,t):
return extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape)
def predict_start_from_noise(self, x_t, t, noise):
return (
@@ -382,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
@@ -399,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
@@ -528,17 +535,17 @@ class LatentDiffusion(DDPM):
n_samples = default(self.en_and_decode_n_samples_a_time, self.temporal_length)
n_rounds = math.ceil(z.shape[0] / n_samples)
with torch.autocast("cuda", enabled=True):
for n in range(n_rounds):
if isinstance(self.first_stage_model.decoder, VideoDecoder):
kwargs.update({"timesteps": len(z[n * n_samples : (n + 1) * n_samples])})
else:
kwargs = {}
out = self.first_stage_model.decode(
z[n * n_samples : (n + 1) * n_samples], **kwargs
)
results.append(out)
#with torch.autocast(mm.get_autocast_device(device), enabled=True):
for n in range(n_rounds):
if isinstance(self.first_stage_model.decoder, VideoDecoder):
kwargs.update({"timesteps": len(z[n * n_samples : (n + 1) * n_samples])})
else:
kwargs = {}
out = self.first_stage_model.decode(
z[n * n_samples : (n + 1) * n_samples], **kwargs
)
results.append(out)
results = torch.cat(results, dim=0)
if reshape_back:
@@ -567,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]
@@ -699,9 +717,9 @@ class LatentDiffusion(DDPM):
class LatentVisualDiffusion(LatentDiffusion):
def __init__(self, img_cond_stage_config, image_proj_stage_config, freeze_embedder=True, *args, **kwargs):
super().__init__(*args, **kwargs)
#self._init_embedder(img_cond_stage_config, freeze_embedder)
self._init_embedder(img_cond_stage_config, freeze_embedder)
self.image_proj_model = instantiate_from_config(image_proj_stage_config)
self.embedder = None
def _init_embedder(self, config, freeze=True):
embedder = instantiate_from_config(config)
if freeze:
@@ -717,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:
@@ -732,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)
+20 -12
View File
@@ -4,8 +4,10 @@ import torch
from ....lvdm.models.utils_diffusion import make_ddim_sampling_parameters, make_ddim_timesteps, rescale_noise_cfg
from ....lvdm.common import noise_like
from ....lvdm.common import extract_into_tensor
import copy
import comfy.utils
import comfy.model_management as mm
device = mm.get_torch_device()
class DDIMSampler(object):
def __init__(self, model, schedule="linear", **kwargs):
@@ -17,13 +19,16 @@ class DDIMSampler(object):
def register_buffer(self, name, attr):
if type(attr) == torch.Tensor:
if attr.device != torch.device("cuda"):
attr = attr.to(torch.device("cuda"))
if attr.device != torch.device(device):
if mm.is_device_mps(device):
attr = attr.to(torch.device(device), torch.float32)
else:
attr = attr.to(torch.device(device))
setattr(self, name, attr)
def make_schedule(self, ddim_num_steps, ddim_discretize="uniform", ddim_eta=0., verbose=True):
def make_schedule(self, ddim_num_steps, ddim_discretize="uniform", ddim_eta=0., ddpm_from=1000, verbose=True):
self.ddim_timesteps = make_ddim_timesteps(ddim_discr_method=ddim_discretize, num_ddim_timesteps=ddim_num_steps,
num_ddpm_timesteps=self.ddpm_num_timesteps,verbose=verbose)
num_ddpm_timesteps=ddpm_from,verbose=verbose)
alphas_cumprod = self.model.alphas_cumprod
assert alphas_cumprod.shape[0] == self.ddpm_num_timesteps, 'alphas have to be defined for each timestep'
to_torch = lambda x: x.clone().detach().to(torch.float32).to(self.model.device)
@@ -83,6 +88,7 @@ class DDIMSampler(object):
fs=None,
timestep_spacing='uniform', #uniform_trailing for starting from last timestep
guidance_rescale=0.0,
ddpm_from=1000,
**kwargs
):
@@ -100,7 +106,7 @@ class DDIMSampler(object):
if conditioning.shape[0] != batch_size:
print(f"Warning: Got {conditioning.shape[0]} conditionings but batch-size is {batch_size}")
self.make_schedule(ddim_num_steps=S, ddim_discretize=timestep_spacing, ddim_eta=eta, verbose=schedule_verbose)
self.make_schedule(ddim_num_steps=S, ddim_discretize=timestep_spacing, ddim_eta=eta, ddpm_from=ddpm_from, verbose=schedule_verbose)
# make shape
if len(shape) == 3:
@@ -142,8 +148,10 @@ class DDIMSampler(object):
device = self.model.betas.device
b = shape[0]
if x_T is None:
print("Using random noise")
img = torch.randn(shape, device=device)
else:
print("Using input noise")
img = x_T
if precision is not None:
if precision == 16:
@@ -165,6 +173,9 @@ class DDIMSampler(object):
clean_cond = kwargs.pop("clean_cond", False)
sigmas = self.ddim_sigmas_for_original_num_steps if ddim_use_original_steps else self.ddim_sigmas
print("Sigmas:", sigmas)
# cond_copy, unconditional_conditioning_copy = copy.deepcopy(cond), copy.deepcopy(unconditional_conditioning)
pbar = comfy.utils.ProgressBar(total_steps)
for i, step in enumerate(iterator):
@@ -180,10 +191,7 @@ class DDIMSampler(object):
img_orig = self.model.q_sample(x0, ts) # TODO: deterministic forward pass? <ddim inversion>
img = img_orig * mask + (1. - mask) * img # keep original & modify use img
outs = self.p_sample_ddim(img, cond, ts, index=index, use_original_steps=ddim_use_original_steps,
outs = self.p_sample_ddim(img, cond, ts, sigmas, index=index, use_original_steps=ddim_use_original_steps,
quantize_denoised=quantize_denoised, temperature=temperature,
noise_dropout=noise_dropout, score_corrector=score_corrector,
corrector_kwargs=corrector_kwargs,
@@ -204,7 +212,7 @@ class DDIMSampler(object):
return img, intermediates
@torch.no_grad()
def p_sample_ddim(self, x, c, t, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False,
def p_sample_ddim(self, x, c, t, sigmas, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False,
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None,
unconditional_guidance_scale=1., unconditional_conditioning=None,
uc_type=None, conditional_guidance_scale_temporal=None,mask=None,x0=None,guidance_rescale=0.0,**kwargs):
@@ -242,7 +250,7 @@ class DDIMSampler(object):
alphas_prev = self.model.alphas_cumprod_prev if use_original_steps else self.ddim_alphas_prev
sqrt_one_minus_alphas = self.model.sqrt_one_minus_alphas_cumprod if use_original_steps else self.ddim_sqrt_one_minus_alphas
# sigmas = self.model.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas
sigmas = self.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas
#sigmas = self.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas
# select parameters corresponding to the currently considered timestep
if is_video:
+11 -4
View File
@@ -4,6 +4,8 @@ import torch
import torch.nn.functional as F
from einops import repeat
import comfy.model_management as mm
device = mm.get_torch_device()
def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False):
"""
@@ -29,14 +31,19 @@ def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False):
def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3):
if mm.is_device_mps(device):
dtype = torch.float32
else:
dtype = torch.float64
if schedule == "linear":
betas = (
torch.linspace(linear_start ** 0.5, linear_end ** 0.5, n_timestep, dtype=torch.float64) ** 2
torch.linspace(linear_start ** 0.5, linear_end ** 0.5, n_timestep, dtype=dtype) ** 2
)
elif schedule == "cosine":
timesteps = (
torch.arange(n_timestep + 1, dtype=torch.float64) / n_timestep + cosine_s
torch.arange(n_timestep + 1, dtype=dtype) / n_timestep + cosine_s
)
alphas = timesteps / (1 + cosine_s) * np.pi / 2
alphas = torch.cos(alphas).pow(2)
@@ -45,9 +52,9 @@ def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2,
betas = np.clip(betas, a_min=0, a_max=0.999)
elif schedule == "sqrt_linear":
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64)
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=dtype)
elif schedule == "sqrt":
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64) ** 0.5
betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=dtype) ** 0.5
else:
raise ValueError(f"schedule '{schedule}' unknown.")
return betas.numpy()
+32 -37
View File
@@ -16,13 +16,8 @@ from ...lvdm.common import (
)
from ...lvdm.basics import zero_module
class Conv2d(torch.nn.Conv2d):
def reset_parameters(self):
return None
class Linear(torch.nn.Linear):
def reset_parameters(self):
return None
import comfy.ops
ops = comfy.ops.manual_cast
class RelativePosition(nn.Module):
""" https://github.com/evelinehong/Transformer_Relative_Position_PyTorch/blob/master/relative_position.py """
@@ -57,11 +52,11 @@ class CrossAttention(nn.Module):
self.scale = dim_head**-0.5
self.heads = heads
self.dim_head = dim_head
self.to_q = Linear(query_dim, inner_dim, bias=False)
self.to_k = Linear(context_dim, inner_dim, bias=False)
self.to_v = Linear(context_dim, inner_dim, bias=False)
self.to_q = ops.Linear(query_dim, inner_dim, bias=False)
self.to_k = ops.Linear(context_dim, inner_dim, bias=False)
self.to_v = ops.Linear(context_dim, inner_dim, bias=False)
self.to_out = nn.Sequential(Linear(inner_dim, query_dim), nn.Dropout(dropout))
self.to_out = nn.Sequential(ops.Linear(inner_dim, query_dim), nn.Dropout(dropout))
self.relative_position = relative_position
if self.relative_position:
@@ -79,8 +74,8 @@ class CrossAttention(nn.Module):
self.text_context_len = text_context_len
self.image_cross_attention_scale_learnable = image_cross_attention_scale_learnable
if self.image_cross_attention:
self.to_k_ip = Linear(context_dim, inner_dim, bias=False)
self.to_v_ip = Linear(context_dim, inner_dim, bias=False)
self.to_k_ip = ops.Linear(context_dim, inner_dim, bias=False)
self.to_v_ip = ops.Linear(context_dim, inner_dim, bias=False)
if image_cross_attention_scale_learnable:
self.register_parameter('alpha', nn.Parameter(torch.tensor(0.)) )
@@ -229,9 +224,9 @@ class BasicTransformerBlock(nn.Module):
self.attn2 = attn_cls(query_dim=dim, context_dim=context_dim, heads=n_heads, dim_head=d_head, dropout=dropout, video_length=video_length, image_cross_attention=image_cross_attention, image_cross_attention_scale=image_cross_attention_scale, image_cross_attention_scale_learnable=image_cross_attention_scale_learnable,text_context_len=text_context_len)
self.image_cross_attention = image_cross_attention
self.norm1 = nn.LayerNorm(dim)
self.norm2 = nn.LayerNorm(dim)
self.norm3 = nn.LayerNorm(dim)
self.norm1 = ops.LayerNorm(dim)
self.norm2 = ops.LayerNorm(dim)
self.norm3 = ops.LayerNorm(dim)
self.checkpoint = checkpoint
@@ -269,11 +264,11 @@ class SpatialTransformer(nn.Module):
super().__init__()
self.in_channels = in_channels
inner_dim = n_heads * d_head
self.norm = torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
self.norm = ops.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
if not use_linear:
self.proj_in = Conv2d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0)
self.proj_in = ops.Conv2d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0)
else:
self.proj_in = Linear(in_channels, inner_dim)
self.proj_in = ops.Linear(in_channels, inner_dim)
attention_cls = None
self.transformer_blocks = nn.ModuleList([
@@ -292,9 +287,9 @@ class SpatialTransformer(nn.Module):
) for d in range(depth)
])
if not use_linear:
self.proj_out = zero_module(Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0))
self.proj_out = zero_module(ops.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0))
else:
self.proj_out = zero_module(Linear(inner_dim, in_channels))
self.proj_out = zero_module(ops.Linear(inner_dim, in_channels))
self.use_linear = use_linear
@@ -335,12 +330,12 @@ class TemporalTransformer(nn.Module):
self.in_channels = in_channels
inner_dim = n_heads * d_head
self.norm = torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
self.proj_in = nn.Conv1d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0)
self.norm = ops.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
self.proj_in = ops.Conv1d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0)
if not use_linear:
self.proj_in = nn.Conv1d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0)
self.proj_in = ops.Conv1d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0)
else:
self.proj_in = Linear(in_channels, inner_dim)
self.proj_in = ops.Linear(in_channels, inner_dim)
if relative_position:
assert(temporal_length is not None)
@@ -364,9 +359,9 @@ class TemporalTransformer(nn.Module):
checkpoint=use_checkpoint) for d in range(depth)
])
if not use_linear:
self.proj_out = zero_module(nn.Conv1d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0))
self.proj_out = zero_module(ops.Conv1d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0))
else:
self.proj_out = zero_module(Linear(inner_dim, in_channels))
self.proj_out = zero_module(ops.Linear(inner_dim, in_channels))
self.use_linear = use_linear
@@ -444,7 +439,7 @@ class TemporalTransformer(nn.Module):
class GEGLU(nn.Module):
def __init__(self, dim_in, dim_out):
super().__init__()
self.proj = Linear(dim_in, dim_out * 2)
self.proj = ops.Linear(dim_in, dim_out * 2)
def forward(self, x):
x, gate = self.proj(x).chunk(2, dim=-1)
@@ -457,14 +452,14 @@ class FeedForward(nn.Module):
inner_dim = int(dim * mult)
dim_out = default(dim_out, dim)
project_in = nn.Sequential(
Linear(dim, inner_dim),
ops.Linear(dim, inner_dim),
nn.GELU()
) if not glu else GEGLU(dim, inner_dim)
self.net = nn.Sequential(
project_in,
nn.Dropout(dropout),
Linear(inner_dim, dim_out)
ops.Linear(inner_dim, dim_out)
)
def forward(self, x):
@@ -476,8 +471,8 @@ class LinearAttention(nn.Module):
super().__init__()
self.heads = heads
hidden_dim = dim_head * heads
self.to_qkv = Conv2d(dim, hidden_dim * 3, 1, bias = False)
self.to_out = Conv2d(hidden_dim, dim, 1)
self.to_qkv = ops.Conv2d(dim, hidden_dim * 3, 1, bias = False)
self.to_out = ops.Conv2d(hidden_dim, dim, 1)
def forward(self, x):
b, c, h, w = x.shape
@@ -495,23 +490,23 @@ class SpatialSelfAttention(nn.Module):
super().__init__()
self.in_channels = in_channels
self.norm = torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
self.q = torch.Conv2d(in_channels,
self.norm = ops.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True)
self.q = ops.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.k = torch.Conv2d(in_channels,
self.k = ops.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.v = torch.Conv2d(in_channels,
self.v = ops.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.proj_out = torch.Conv2d(in_channels,
self.proj_out = ops.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
+32 -29
View File
@@ -12,6 +12,9 @@ from torch.utils.checkpoint import checkpoint
logpy = logging.getLogger(__name__)
import comfy.ops
ops = comfy.ops.manual_cast
if version.parse(torch.__version__) >= version.parse("2.0.0"):
SDP_IS_AVAILABLE = True
from torch.backends.cuda import SDPBackend, sdp_kernel
@@ -87,7 +90,7 @@ def init_(tensor):
class GEGLU(nn.Module):
def __init__(self, dim_in, dim_out):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out * 2)
self.proj = ops.Linear(dim_in, dim_out * 2)
def forward(self, x):
x, gate = self.proj(x).chunk(2, dim=-1)
@@ -100,13 +103,13 @@ class FeedForward(nn.Module):
inner_dim = int(dim * mult)
dim_out = default(dim_out, dim)
project_in = (
nn.Sequential(nn.Linear(dim, inner_dim), nn.GELU())
nn.Sequential(ops.Linear(dim, inner_dim), nn.GELU())
if not glu
else GEGLU(dim, inner_dim)
)
self.net = nn.Sequential(
project_in, nn.Dropout(dropout), nn.Linear(inner_dim, dim_out)
project_in, nn.Dropout(dropout), ops.Linear(inner_dim, dim_out)
)
def forward(self, x):
@@ -123,7 +126,7 @@ def zero_module(module):
def Normalize(in_channels):
return torch.nn.GroupNorm(
return ops.GroupNorm(
num_groups=32, num_channels=in_channels, eps=1e-6, affine=True
)
@@ -133,8 +136,8 @@ class LinearAttention(nn.Module):
super().__init__()
self.heads = heads
hidden_dim = dim_head * heads
self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias=False)
self.to_out = nn.Conv2d(hidden_dim, dim, 1)
self.to_qkv = ops.Conv2d(dim, hidden_dim * 3, 1, bias=False)
self.to_out = ops.Conv2d(hidden_dim, dim, 1)
def forward(self, x):
b, c, h, w = x.shape
@@ -169,9 +172,9 @@ class SelfAttention(nn.Module):
head_dim = dim // num_heads
self.scale = qk_scale or head_dim**-0.5
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
self.qkv = ops.Linear(dim, dim * 3, bias=qkv_bias)
self.attn_drop = nn.Dropout(attn_drop)
self.proj = nn.Linear(dim, dim)
self.proj = ops.Linear(dim, dim)
self.proj_drop = nn.Dropout(proj_drop)
assert attn_mode in self.ATTENTION_MODES
self.attn_mode = attn_mode
@@ -213,16 +216,16 @@ class SpatialSelfAttention(nn.Module):
self.in_channels = in_channels
self.norm = Normalize(in_channels)
self.q = torch.nn.Conv2d(
self.q = torch.ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.k = torch.nn.Conv2d(
self.k = torch.ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.v = torch.nn.Conv2d(
self.v = torch.ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.proj_out = torch.nn.Conv2d(
self.proj_out = torch.ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
@@ -269,12 +272,12 @@ class CrossAttention(nn.Module):
self.scale = dim_head**-0.5
self.heads = heads
self.to_q = nn.Linear(query_dim, inner_dim, bias=False)
self.to_k = nn.Linear(context_dim, inner_dim, bias=False)
self.to_v = nn.Linear(context_dim, inner_dim, bias=False)
self.to_q = ops.Linear(query_dim, inner_dim, bias=False)
self.to_k = ops.Linear(context_dim, inner_dim, bias=False)
self.to_v = ops.Linear(context_dim, inner_dim, bias=False)
self.to_out = nn.Sequential(
nn.Linear(inner_dim, query_dim), nn.Dropout(dropout)
ops.Linear(inner_dim, query_dim), nn.Dropout(dropout)
)
self.backend = backend
@@ -361,12 +364,12 @@ class MemoryEfficientCrossAttention(nn.Module):
self.heads = heads
self.dim_head = dim_head
self.to_q = nn.Linear(query_dim, inner_dim, bias=False)
self.to_k = nn.Linear(context_dim, inner_dim, bias=False)
self.to_v = nn.Linear(context_dim, inner_dim, bias=False)
self.to_q = ops.Linear(query_dim, inner_dim, bias=False)
self.to_k = ops.Linear(context_dim, inner_dim, bias=False)
self.to_v = ops.Linear(context_dim, inner_dim, bias=False)
self.to_out = nn.Sequential(
nn.Linear(inner_dim, query_dim), nn.Dropout(dropout)
ops.Linear(inner_dim, query_dim), nn.Dropout(dropout)
)
self.attention_op: Optional[Any] = None
@@ -517,9 +520,9 @@ class BasicTransformerBlock(nn.Module):
dropout=dropout,
backend=sdp_backend,
) # is self-attn if context is none
self.norm1 = nn.LayerNorm(dim)
self.norm2 = nn.LayerNorm(dim)
self.norm3 = nn.LayerNorm(dim)
self.norm1 = ops.LayerNorm(dim)
self.norm2 = ops.LayerNorm(dim)
self.norm3 = ops.LayerNorm(dim)
self.checkpoint = checkpoint
if self.checkpoint:
logpy.debug(f"{self.__class__.__name__} is using checkpointing")
@@ -601,8 +604,8 @@ class BasicTransformerSingleLayerBlock(nn.Module):
context_dim=context_dim,
)
self.ff = FeedForward(dim, dropout=dropout, glu=gated_ff)
self.norm1 = nn.LayerNorm(dim)
self.norm2 = nn.LayerNorm(dim)
self.norm1 = ops.LayerNorm(dim)
self.norm2 = ops.LayerNorm(dim)
self.checkpoint = checkpoint
def forward(self, x, context=None):
@@ -668,11 +671,11 @@ class SpatialTransformer(nn.Module):
inner_dim = n_heads * d_head
self.norm = Normalize(in_channels)
if not use_linear:
self.proj_in = nn.Conv2d(
self.proj_in = ops.Conv2d(
in_channels, inner_dim, kernel_size=1, stride=1, padding=0
)
else:
self.proj_in = nn.Linear(in_channels, inner_dim)
self.proj_in = ops.Linear(in_channels, inner_dim)
self.transformer_blocks = nn.ModuleList(
[
@@ -692,11 +695,11 @@ class SpatialTransformer(nn.Module):
)
if not use_linear:
self.proj_out = zero_module(
nn.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0)
ops.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0)
)
else:
# self.proj_out = zero_module(nn.Linear(in_channels, inner_dim))
self.proj_out = zero_module(nn.Linear(inner_dim, in_channels))
self.proj_out = zero_module(ops.Linear(inner_dim, in_channels))
self.use_linear = use_linear
def forward(self, x, context=None):
+14 -12
View File
@@ -1,12 +1,14 @@
import torch
import torch.nn as nn
import kornia
import open_clip
from torch.utils.checkpoint import checkpoint
from transformers import T5Tokenizer, T5EncoderModel, CLIPTokenizer, CLIPTextModel
#import kornia
#import open_clip
#from torch.utils.checkpoint import checkpoint
#from transformers import T5Tokenizer, T5EncoderModel, CLIPTokenizer, CLIPTextModel
from ....lvdm.common import autocast
from ....utils.utils import count_params
import comfy.model_management as mm
device = mm.get_torch_device()
class AbstractEncoder(nn.Module):
def __init__(self):
@@ -41,7 +43,7 @@ class ClassEmbedder(nn.Module):
c = self.embedding(c)
return c
def get_unconditional_conditioning(self, bs, device="cuda"):
def get_unconditional_conditioning(self, bs, device=device):
uc_class = self.n_classes - 1 # 1000 classes --> 0 ... 999, one extra class for ucg (class 1000)
uc = torch.ones((bs,), device=device) * uc_class
uc = {self.key: uc}
@@ -57,7 +59,7 @@ def disabled_train(self, mode=True):
class FrozenT5Embedder(AbstractEncoder):
"""Uses the T5 transformer encoder for text"""
def __init__(self, version="google/t5-v1_1-large", device="cuda", max_length=77,
def __init__(self, version="google/t5-v1_1-large", device=device, max_length=77,
freeze=True): # others are google/t5-v1_1-xl and google/t5-v1_1-xxl
super().__init__()
self.tokenizer = T5Tokenizer.from_pretrained(version)
@@ -94,7 +96,7 @@ class FrozenCLIPEmbedder(AbstractEncoder):
"hidden"
]
def __init__(self, version="openai/clip-vit-large-patch14", device="cuda", max_length=77,
def __init__(self, version="openai/clip-vit-large-patch14", device=device, max_length=77,
freeze=True, layer="last", layer_idx=None): # clip-vit-base-patch32
super().__init__()
assert layer in self.LAYERS
@@ -138,7 +140,7 @@ class ClipImageEmbedder(nn.Module):
self,
model,
jit=False,
device='cuda' if torch.cuda.is_available() else 'cpu',
device=device,
antialias=True,
ucg_rate=0.
):
@@ -181,7 +183,7 @@ class FrozenOpenCLIPEmbedder(AbstractEncoder):
"penultimate"
]
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", max_length=77,
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device=device, max_length=77,
freeze=True, layer="last"):
super().__init__()
assert layer in self.LAYERS
@@ -239,7 +241,7 @@ class FrozenOpenCLIPImageEmbedder(AbstractEncoder):
Uses the OpenCLIP vision transformer encoder for images
"""
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", max_length=77,
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device=device, max_length=77,
freeze=True, layer="pooled", antialias=True, ucg_rate=0.):
super().__init__()
model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'),
@@ -297,7 +299,7 @@ class FrozenOpenCLIPImageEmbedderV2(AbstractEncoder):
Uses the OpenCLIP vision transformer encoder for images
"""
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda",
def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device=device,
freeze=True, layer="pooled", antialias=True):
super().__init__()
return
@@ -373,7 +375,7 @@ class FrozenOpenCLIPImageEmbedderV2(AbstractEncoder):
return x
class FrozenCLIPT5Encoder(AbstractEncoder):
def __init__(self, clip_version="openai/clip-vit-large-patch14", t5_version="google/t5-v1_1-xl", device="cuda",
def __init__(self, clip_version="openai/clip-vit-large-patch14", t5_version="google/t5-v1_1-xl", device=device,
clip_max_length=77, t5_max_length=77):
super().__init__()
self.clip_encoder = FrozenCLIPEmbedder(clip_version, device, max_length=clip_max_length)
+15 -13
View File
@@ -5,6 +5,8 @@ import math
import torch
import torch.nn as nn
import comfy.ops
ops = comfy.ops.manual_cast
class ImageProjModel(nn.Module):
"""Projection Model"""
@@ -12,8 +14,8 @@ class ImageProjModel(nn.Module):
super().__init__()
self.cross_attention_dim = cross_attention_dim
self.clip_extra_context_tokens = clip_extra_context_tokens
self.proj = nn.Linear(clip_embeddings_dim, self.clip_extra_context_tokens * cross_attention_dim)
self.norm = nn.LayerNorm(cross_attention_dim)
self.proj = ops.Linear(clip_embeddings_dim, self.clip_extra_context_tokens * cross_attention_dim)
self.norm = ops.LayerNorm(cross_attention_dim)
def forward(self, image_embeds):
#embeds = image_embeds
@@ -27,10 +29,10 @@ class ImageProjModel(nn.Module):
def FeedForward(dim, mult=4):
inner_dim = int(dim * mult)
return nn.Sequential(
nn.LayerNorm(dim),
nn.Linear(dim, inner_dim, bias=False),
ops.LayerNorm(dim),
ops.Linear(dim, inner_dim, bias=False),
nn.GELU(),
nn.Linear(inner_dim, dim, bias=False),
ops.Linear(inner_dim, dim, bias=False),
)
@@ -53,12 +55,12 @@ class PerceiverAttention(nn.Module):
self.heads = heads
inner_dim = dim_head * heads
self.norm1 = nn.LayerNorm(dim)
self.norm2 = nn.LayerNorm(dim)
self.norm1 = ops.LayerNorm(dim)
self.norm2 = ops.LayerNorm(dim)
self.to_q = nn.Linear(dim, inner_dim, bias=False)
self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False)
self.to_out = nn.Linear(inner_dim, dim, bias=False)
self.to_q = ops.Linear(dim, inner_dim, bias=False)
self.to_kv = ops.Linear(dim, inner_dim * 2, bias=False)
self.to_out = ops.Linear(inner_dim, dim, bias=False)
def forward(self, x, latents):
@@ -116,9 +118,9 @@ class Resampler(nn.Module):
num_queries = num_queries * video_length
self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5)
self.proj_in = nn.Linear(embedding_dim, dim)
self.proj_out = nn.Linear(dim, output_dim)
self.norm_out = nn.LayerNorm(output_dim)
self.proj_in = ops.Linear(embedding_dim, dim)
self.proj_out = ops.Linear(dim, output_dim)
self.norm_out = ops.LayerNorm(output_dim)
self.layers = nn.ModuleList([])
for _ in range(depth):
+31 -28
View File
@@ -7,13 +7,16 @@ from einops import rearrange
from ....utils.utils import instantiate_from_config
from ....lvdm.modules.attention import LinearAttention
import comfy.ops
ops = comfy.ops.manual_cast
def nonlinearity(x):
# swish
return x*torch.sigmoid(x)
def Normalize(in_channels, num_groups=32):
return torch.nn.GroupNorm(num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True)
return ops.GroupNorm(num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True)
@@ -29,22 +32,22 @@ class AttnBlock(nn.Module):
self.in_channels = in_channels
self.norm = Normalize(in_channels)
self.q = torch.nn.Conv2d(in_channels,
self.q = ops.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.k = torch.nn.Conv2d(in_channels,
self.k = ops.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.v = torch.nn.Conv2d(in_channels,
self.v = ops.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
padding=0)
self.proj_out = torch.nn.Conv2d(in_channels,
self.proj_out = ops.Conv2d(in_channels,
in_channels,
kernel_size=1,
stride=1,
@@ -94,7 +97,7 @@ class Downsample(nn.Module):
self.in_channels = in_channels
if self.with_conv:
# no asymmetric padding in torch conv, must do it ourselves
self.conv = torch.nn.Conv2d(in_channels,
self.conv = ops.Conv2d(in_channels,
in_channels,
kernel_size=3,
stride=2,
@@ -114,7 +117,7 @@ class Upsample(nn.Module):
self.with_conv = with_conv
self.in_channels = in_channels
if self.with_conv:
self.conv = torch.nn.Conv2d(in_channels,
self.conv = ops.Conv2d(in_channels,
in_channels,
kernel_size=3,
stride=1,
@@ -158,30 +161,30 @@ class ResnetBlock(nn.Module):
self.use_conv_shortcut = conv_shortcut
self.norm1 = Normalize(in_channels)
self.conv1 = torch.nn.Conv2d(in_channels,
self.conv1 = ops.Conv2d(in_channels,
out_channels,
kernel_size=3,
stride=1,
padding=1)
if temb_channels > 0:
self.temb_proj = torch.nn.Linear(temb_channels,
self.temb_proj = ops.Linear(temb_channels,
out_channels)
self.norm2 = Normalize(out_channels)
self.dropout = torch.nn.Dropout(dropout)
self.conv2 = torch.nn.Conv2d(out_channels,
self.conv2 = ops.Conv2d(out_channels,
out_channels,
kernel_size=3,
stride=1,
padding=1)
if self.in_channels != self.out_channels:
if self.use_conv_shortcut:
self.conv_shortcut = torch.nn.Conv2d(in_channels,
self.conv_shortcut = ops.Conv2d(in_channels,
out_channels,
kernel_size=3,
stride=1,
padding=1)
else:
self.nin_shortcut = torch.nn.Conv2d(in_channels,
self.nin_shortcut = ops.Conv2d(in_channels,
out_channels,
kernel_size=1,
stride=1,
@@ -227,14 +230,14 @@ class Model(nn.Module):
# timestep embedding
self.temb = nn.Module()
self.temb.dense = nn.ModuleList([
torch.nn.Linear(self.ch,
ops.Linear(self.ch,
self.temb_ch),
torch.nn.Linear(self.temb_ch,
ops.Linear(self.temb_ch,
self.temb_ch),
])
# downsampling
self.conv_in = torch.nn.Conv2d(in_channels,
self.conv_in = ops.Conv2d(in_channels,
self.ch,
kernel_size=3,
stride=1,
@@ -303,7 +306,7 @@ class Model(nn.Module):
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(block_in,
self.conv_out = ops.Conv2d(block_in,
out_ch,
kernel_size=3,
stride=1,
@@ -376,7 +379,7 @@ class Encoder(nn.Module):
self.in_channels = in_channels
# downsampling
self.conv_in = torch.nn.Conv2d(in_channels,
self.conv_in = ops.Conv2d(in_channels,
self.ch,
kernel_size=3,
stride=1,
@@ -421,7 +424,7 @@ class Encoder(nn.Module):
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(block_in,
self.conv_out = ops.Conv2d(block_in,
2*z_channels if double_z else z_channels,
kernel_size=3,
stride=1,
@@ -498,7 +501,7 @@ class Decoder(nn.Module):
self.z_shape, np.prod(self.z_shape)))
# z to block_in
self.conv_in = torch.nn.Conv2d(z_channels,
self.conv_in = ops.Conv2d(z_channels,
block_in,
kernel_size=3,
stride=1,
@@ -540,7 +543,7 @@ class Decoder(nn.Module):
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(block_in,
self.conv_out = ops.Conv2d(block_in,
out_ch,
kernel_size=3,
stride=1,
@@ -591,7 +594,7 @@ class Decoder(nn.Module):
class SimpleDecoder(nn.Module):
def __init__(self, in_channels, out_channels, *args, **kwargs):
super().__init__()
self.model = nn.ModuleList([nn.Conv2d(in_channels, in_channels, 1),
self.model = nn.ModuleList([ops.Conv2d(in_channels, in_channels, 1),
ResnetBlock(in_channels=in_channels,
out_channels=2 * in_channels,
temb_channels=0, dropout=0.0),
@@ -601,11 +604,11 @@ class SimpleDecoder(nn.Module):
ResnetBlock(in_channels=4 * in_channels,
out_channels=2 * in_channels,
temb_channels=0, dropout=0.0),
nn.Conv2d(2*in_channels, in_channels, 1),
ops.Conv2d(2*in_channels, in_channels, 1),
Upsample(in_channels, with_conv=True)])
# end
self.norm_out = Normalize(in_channels)
self.conv_out = torch.nn.Conv2d(in_channels,
self.conv_out = ops.Conv2d(in_channels,
out_channels,
kernel_size=3,
stride=1,
@@ -652,7 +655,7 @@ class UpsampleDecoder(nn.Module):
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(block_in,
self.conv_out = ops.Conv2d(block_in,
out_channels,
kernel_size=3,
stride=1,
@@ -677,7 +680,7 @@ class LatentRescaler(nn.Module):
super().__init__()
# residual block, interpolate, residual block
self.factor = factor
self.conv_in = nn.Conv2d(in_channels,
self.conv_in = ops.Conv2d(in_channels,
mid_channels,
kernel_size=3,
stride=1,
@@ -692,7 +695,7 @@ class LatentRescaler(nn.Module):
temb_channels=0,
dropout=0.0) for _ in range(depth)])
self.conv_out = nn.Conv2d(mid_channels,
self.conv_out = ops.Conv2d(mid_channels,
out_channels,
kernel_size=1,
)
@@ -774,7 +777,7 @@ class Resize(nn.Module):
raise NotImplementedError()
assert in_channels is not None
# no asymmetric padding in torch conv, must do it ourselves
self.conv = torch.nn.Conv2d(in_channels,
self.conv = ops.Conv2d(in_channels,
in_channels,
kernel_size=4,
stride=2,
@@ -809,7 +812,7 @@ class FirstStagePostProcessor(nn.Module):
n_channels = self.pretrained_model.encoder.ch
self.proj_norm = Normalize(in_channels,num_groups=in_channels//2)
self.proj = nn.Conv2d(in_channels,n_channels,kernel_size=3,
self.proj = ops.Conv2d(in_channels,n_channels,kernel_size=3,
stride=1,padding=1)
blocks = []
+417 -12
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
@@ -15,6 +17,11 @@ from ....lvdm.basics import (
)
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):
"""
@@ -167,7 +174,7 @@ class ResBlock(TimestepBlock):
self.emb_layers = nn.Sequential(
nn.SiLU(),
nn.Linear(
ops.Linear(
emb_channels,
2 * self.out_channels if use_scale_shift_norm else self.out_channels,
),
@@ -176,7 +183,7 @@ class ResBlock(TimestepBlock):
normalization(self.out_channels),
nn.SiLU(),
nn.Dropout(p=dropout),
zero_module(nn.Conv2d(self.out_channels, self.out_channels, 3, padding=1)),
zero_module(ops.Conv2d(self.out_channels, self.out_channels, 3, padding=1)),
)
if self.out_channels == channels:
@@ -253,17 +260,17 @@ class TemporalConvBlock(nn.Module):
# conv layers
self.conv1 = nn.Sequential(
nn.GroupNorm(32, in_channels), nn.SiLU(),
nn.Conv3d(in_channels, out_channels, th_kernel_shape, padding=th_padding_shape))
ops.GroupNorm(32, in_channels), nn.SiLU(),
ops.Conv3d(in_channels, out_channels, th_kernel_shape, padding=th_padding_shape))
self.conv2 = nn.Sequential(
nn.GroupNorm(32, out_channels), nn.SiLU(), nn.Dropout(dropout),
nn.Conv3d(out_channels, in_channels, tw_kernel_shape, padding=tw_padding_shape))
ops.GroupNorm(32, out_channels), nn.SiLU(), nn.Dropout(dropout),
ops.Conv3d(out_channels, in_channels, tw_kernel_shape, padding=tw_padding_shape))
self.conv3 = nn.Sequential(
nn.GroupNorm(32, out_channels), nn.SiLU(), nn.Dropout(dropout),
nn.Conv3d(out_channels, in_channels, th_kernel_shape, padding=th_padding_shape))
ops.GroupNorm(32, out_channels), nn.SiLU(), nn.Dropout(dropout),
ops.Conv3d(out_channels, in_channels, th_kernel_shape, padding=th_padding_shape))
self.conv4 = nn.Sequential(
nn.GroupNorm(32, out_channels), nn.SiLU(), nn.Dropout(dropout),
nn.Conv3d(out_channels, in_channels, tw_kernel_shape, padding=tw_padding_shape))
ops.GroupNorm(32, out_channels), nn.SiLU(), nn.Dropout(dropout),
ops.Conv3d(out_channels, in_channels, tw_kernel_shape, padding=tw_padding_shape))
# zero out the last layer params,so the conv block is identity
nn.init.zeros_(self.conv4[-1].weight)
@@ -545,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)
@@ -592,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)
@@ -601,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])
+14 -11
View File
@@ -7,6 +7,9 @@ import torch
from torch import nn, einsum
import torch.nn.functional as F
import comfy.ops
ops = comfy.ops.manual_cast
# constants
DEFAULT_DIM_HEAD = 64
@@ -183,7 +186,7 @@ class GRUGating(nn.Module):
class GEGLU(nn.Module):
def __init__(self, dim_in, dim_out):
super().__init__()
self.proj = nn.Linear(dim_in, dim_out * 2)
self.proj = ops.Linear(dim_in, dim_out * 2)
def forward(self, x):
x, gate = self.proj(x).chunk(2, dim=-1)
@@ -196,14 +199,14 @@ class FeedForward(nn.Module):
inner_dim = int(dim * mult)
dim_out = default(dim_out, dim)
project_in = nn.Sequential(
nn.Linear(dim, inner_dim),
ops.Linear(dim, inner_dim),
nn.GELU()
) if not glu else GEGLU(dim, inner_dim)
self.net = nn.Sequential(
project_in,
nn.Dropout(dropout),
nn.Linear(inner_dim, dim_out)
ops.Linear(inner_dim, dim_out)
)
def forward(self, x):
@@ -236,9 +239,9 @@ class Attention(nn.Module):
inner_dim = dim_head * heads
self.to_q = nn.Linear(dim, inner_dim, bias=False)
self.to_k = nn.Linear(dim, inner_dim, bias=False)
self.to_v = nn.Linear(dim, inner_dim, bias=False)
self.to_q = ops.Linear(dim, inner_dim, bias=False)
self.to_k = ops.Linear(dim, inner_dim, bias=False)
self.to_v = ops.Linear(dim, inner_dim, bias=False)
self.dropout = nn.Dropout(dropout)
# talking heads
@@ -262,7 +265,7 @@ class Attention(nn.Module):
# attention on attention
self.attn_on_attn = on_attn
self.to_out = nn.Sequential(nn.Linear(inner_dim, dim * 2), nn.GLU()) if on_attn else nn.Linear(inner_dim, dim)
self.to_out = nn.Sequential(ops.Linear(inner_dim, dim * 2), nn.GLU()) if on_attn else ops.Linear(inner_dim, dim)
def forward(
self,
@@ -413,7 +416,7 @@ class AttentionLayers(nn.Module):
self.residual_attn = residual_attn
self.cross_residual_attn = cross_residual_attn
norm_class = ScaleNorm if use_scalenorm else nn.LayerNorm
norm_class = ScaleNorm if use_scalenorm else ops.LayerNorm
norm_class = RMSNorm if use_rmsnorm else norm_class
norm_fn = partial(norm_class, dim)
@@ -573,13 +576,13 @@ class TransformerWrapper(nn.Module):
use_pos_emb and not attn_layers.has_pos_emb) else always(0)
self.emb_dropout = nn.Dropout(emb_dropout)
self.project_emb = nn.Linear(emb_dim, dim) if emb_dim != dim else nn.Identity()
self.project_emb = ops.Linear(emb_dim, dim) if emb_dim != dim else nn.Identity()
self.attn_layers = attn_layers
self.norm = nn.LayerNorm(dim)
self.norm = ops.LayerNorm(dim)
self.init_()
self.to_logits = nn.Linear(dim, num_tokens) if not tie_embedding else lambda t: t @ self.token_emb.weight.t()
self.to_logits = ops.Linear(dim, num_tokens) if not tie_embedding else lambda t: t @ self.token_emb.weight.t()
# memory tokens (like [cls]) from Memory Transformers paper
num_memory_tokens = default(num_memory_tokens, 0)
+16 -25
View File
@@ -557,8 +557,6 @@ class DynamiCrafterI2V:
text_emb = positive[0][0].to(device)
cond_images = clip_vision.encode_image(image.permute(0, 2, 3, 1))['last_hidden_state'].to(device)
cond_images = torch.sum(cond_images, dim=0).unsqueeze(0)
cond_images = torch.mean(cond_images, dim=0).unsqueeze(0)
img_emb = self.model.image_proj_model(cond_images)
@@ -816,12 +814,11 @@ class ToonCrafterInterpolation:
pbar = comfy.utils.ProgressBar(len(images) - 1)
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():
for i in range(len(images) - 1) if len(images) > 1 else range(len(images)):
for i in range(len(images) - 1):
videos, videos2 = None, None
mm.soft_empty_cache()
image = images[i].unsqueeze(0)
if len(images) !=1:
image2 = images[i+1].unsqueeze(0)
image2 = images[i+1].unsqueeze(0)
B, C, H, W = image.shape
noise_shape = [B, self.model.model.diffusion_model.out_channels, frames, H // 8, W // 8]
@@ -833,16 +830,12 @@ class ToonCrafterInterpolation:
image2 += torch.randn_like(image) * augmentation_level
encode_pixels = image.unsqueeze(2) * 2 - 1
videos = encode_pixels # bc1hw
videos = repeat(videos, 'b c t h w -> b c (repeat t) h w', repeat=frames // 2)
if len(images) == 1:
videos = torch.cat([videos, videos], dim=2)
else:
encode_pixels = image2.unsqueeze(2) * 2 - 1
videos2 = encode_pixels # bc1hw
videos2 = repeat(videos2, 'b c t h w -> b c (repeat t) h w', repeat=frames // 2)
videos = torch.cat([videos, videos2], dim=2)
videos = encode_pixels # bc1hw
videos = repeat(videos, 'b c t h w -> b c (repeat t) h w', repeat=frames//2)
encode_pixels = image2.unsqueeze(2) * 2 - 1
videos2 = encode_pixels # bc1hw
videos2 = repeat(videos2, 'b c t h w -> b c (repeat t) h w', repeat=frames//2)
videos = torch.cat([videos, videos2], dim=2)
try:
z, hs = get_latent_z_with_hidden_states(self.model, videos)
@@ -854,25 +847,23 @@ class ToonCrafterInterpolation:
img_tensor_repeat = torch.zeros_like(z)
img_tensor_repeat[:,:,:1,:,:] = z[:,:,:1,:,:]
if len(images) !=1:
img_tensor_repeat[:,:,-1:,:,:] = z[:,:,-1:,:,:]
img_tensor_repeat[:,:,-1:,:,:] = z[:,:,-1:,:,:]
self.model.first_stage_model.to(offload_device)
text_emb = positive[0][0].to(device)
self.model.image_proj_model.to(device)
cond_images = clip_vision.encode_image(image.permute(0, 2, 3, 1))["last_hidden_state"].to(device)
cond_images2 = clip_vision.encode_image(image2.permute(0, 2, 3, 1))["last_hidden_state"].to(device)
self.model.image_proj_model.to(device)
img_emb = self.model.image_proj_model(cond_images)
if len(images) !=1:
cond_images2 = clip_vision.encode_image(image2.permute(0, 2, 3, 1))["last_hidden_state"].to(device)
img_emb2 = self.model.image_proj_model(cond_images2)
img_embeds = img_emb * image_embed_ratio + img_emb2 * (1.0 - image_embed_ratio)
else:
img_embeds = img_emb
img_emb2 = self.model.image_proj_model(cond_images2)
img_embeds = img_emb * image_embed_ratio + img_emb2 * (1.0 - image_embed_ratio)
imtext_cond = torch.cat([text_emb, img_embeds], dim=1)
del cond_images, img_emb, text_emb
del cond_images, img_emb, img_emb2, text_emb
if comfy.model_management.is_device_mps(device):
fs = torch.tensor([fs], dtype=torch.float32, device=self.model.device)
+15
View File
@@ -0,0 +1,15 @@
[project]
name = "comfyui-dynamicrafterwrapper"
description = "Wrapper nodes to use Dynami/ToonCrafter image2video and frame interpolation models in ComfyUI"
version = "1.0.2"
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"]
[project.urls]
Repository = "https://github.com/kijai/ComfyUI-DynamiCrafterWrapper"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = "kijai"
DisplayName = "ComfyUI-DynamiCrafterWrapper"
Icon = ""
+1 -3
View File
@@ -1,10 +1,8 @@
einops>=0.3.0
numpy>=1.24.2
omegaconf>=2.1.1
Pillow>=9.5.0
pytorch_lightning>=2.2.1
tqdm>=4.65.0
transformers>=4.25.1
timm
open_clip_torch>=2.23.0
kornia
accelerate
+26 -38
View File
@@ -5,51 +5,39 @@ import torch
from einops import rearrange
from safetensors.torch import load_file
def load_model_checkpoint(model, ckpt):
def load_checkpoint(model, ckpt, full_strict):
if "safetensors" in ckpt:
try:
state_dict = load_file(ckpt)
except:
state_dict = torch.load(ckpt, map_location="cpu")
else:
state_dict = torch.load(ckpt, map_location="cpu")
if "state_dict" in list(state_dict.keys()):
state_dict = state_dict["state_dict"]
from contextlib import nullcontext
try:
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
is_accelerate_available = True
except:
pass
filtered_state_dict = {
k: v
for k, v in state_dict.items()
if not (k.startswith("cond_stage_model") or k.startswith("embedder"))
#if not (k.startswith("cond_stage_model"))
} # Filter out keys starting with "cond_stage_model" and "embedder"
def load_model_checkpoint(model, file_path, dtype, device):
if "safetensors" in file_path:
try:
model.load_state_dict(filtered_state_dict, strict=full_strict)
state_dict = load_file(file_path)
except:
## rename the keys for 256x256 model
new_pl_sd = OrderedDict()
for k,v in state_dict.items():
new_pl_sd[k] = v
state_dict = torch.load(file_path, map_location="cpu")
else:
state_dict = torch.load(file_path, map_location="cpu")
if "state_dict" in list(state_dict.keys()):
state_dict = state_dict["state_dict"]
for k in list(new_pl_sd.keys()):
if "framestride_embed" in k:
new_key = k.replace("framestride_embed", "fps_embedding")
new_pl_sd[new_key] = new_pl_sd[k]
del new_pl_sd[k]
model.load_state_dict(new_pl_sd, strict=full_strict)
# else:
# ## deepspeed
# new_pl_sd = OrderedDict()
# for key in state_dict['module'].keys():
# new_pl_sd[key[16:]]=state_dict['module'][key]
# model.load_state_dict(new_pl_sd, strict=full_strict)
filtered_state_dict = {
k: v
for k, v in state_dict.items()
if not (k.startswith("cond_stage_model") or k.startswith("embedder"))
#if not (k.startswith("cond_stage_model"))
} # Filter out keys starting with "cond_stage_model" and "embedder"
if is_accelerate_available:
for key in filtered_state_dict:
set_module_tensor_to_device(model, key, dtype=dtype, device=device, value=filtered_state_dict[key])
else:
model.load_state_dict(filtered_state_dict, strict=True)
return model
load_checkpoint(model, ckpt, full_strict=False)
print('>>> model checkpoint loaded.')
return model
def load_prompts(prompt_file):
f = open(prompt_file, 'r')
prompt_list = []