Merge branch 'main' into develop
This commit is contained in:
@@ -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 }}
|
||||
@@ -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
@@ -431,7 +431,7 @@
|
||||
"Node name for S&R": "DownloadAndLoadDynamiCrafterModel"
|
||||
},
|
||||
"widgets_values": [
|
||||
"dynamicrafter_1024_v1_bf16.safetensors",
|
||||
"dynamicrafter_1024_fp16_pruned.safetensors",
|
||||
"auto",
|
||||
true
|
||||
]
|
||||
|
||||
@@ -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
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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 = []
|
||||
|
||||
@@ -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])
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
@@ -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 = []
|
||||
|
||||
Reference in New Issue
Block a user