use more comfy ops

This commit is contained in:
kijai
2024-07-06 12:20:21 +03:00
parent 2726b6ef34
commit a8780cd188
7 changed files with 84 additions and 94 deletions
+9 -6
View File
@@ -27,6 +27,9 @@ from functools import partial
import comfy.model_management
device = comfy.model_management.get_torch_device()
import comfy.ops
ops = comfy.ops.manual_cast
try:
import xformers
import xformers.ops
@@ -77,13 +80,13 @@ class ZeroSFT(nn.Module):
nhidden = 128
self.mlp_shared = nn.Sequential(
nn.Conv2d(label_nc, nhidden, kernel_size=ks, padding=pw),
ops.Conv2d(label_nc, nhidden, kernel_size=ks, padding=pw),
nn.SiLU()
)
self.zero_mul = zero_module(nn.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw))
self.zero_add = zero_module(nn.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw))
# self.zero_mul = nn.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw)
# self.zero_add = nn.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw)
self.zero_mul = zero_module(ops.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw))
self.zero_add = zero_module(ops.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw))
# self.zero_mul = ops.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw)
# self.zero_add = ops.Conv2d(nhidden, norm_nc + concat_channels, kernel_size=ks, padding=pw)
self.zero_conv = zero_module(conv_nd(2, label_nc, norm_nc, 1, 1, 0))
self.pre_concat = bool(concat_channels != 0)
@@ -296,7 +299,7 @@ class GLVControl(nn.Module):
self.label_emb = nn.Embedding(num_classes, time_embed_dim)
elif self.num_classes == "continuous":
print("setting up linear c_adm embedding layer")
self.label_emb = nn.Linear(1, time_embed_dim)
self.label_emb = ops.Linear(1, time_embed_dim)
elif self.num_classes == "timestep":
self.label_emb = checkpoint_wrapper_fn(
nn.Sequential(
+4 -5
View File
@@ -14,9 +14,8 @@ from ..modules.distributions.distributions import DiagonalGaussianDistribution
from ..modules.ema import LitEma
from ..util import default, get_obj_from_str, instantiate_from_config
class Conv2d(torch.nn.Conv2d):
def reset_parameters(self):
return None
import comfy.ops
ops = comfy.ops.manual_cast
class AbstractAutoencoder(pl.LightningModule):
"""
@@ -297,8 +296,8 @@ class AutoencoderKL(AutoencodingEngine):
assert ddconfig["double_z"]
self.encoder = Encoder(**ddconfig)
self.decoder = Decoder(**ddconfig)
self.quant_conv = Conv2d(2 * ddconfig["z_channels"], 2 * embed_dim, 1)
self.post_quant_conv = 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
if ckpt_path is not None:
+29 -34
View File
@@ -10,13 +10,8 @@ from einops import rearrange, repeat
from packaging import version
from torch import nn
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
if version.parse(torch.__version__) >= version.parse("2.0.0"):
SDP_IS_AVAILABLE = True
@@ -92,7 +87,7 @@ def init_(tensor):
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)
@@ -105,13 +100,13 @@ class FeedForward(nn.Module):
inner_dim = int(dim * mult)
dim_out = default(dim_out, dim)
project_in = (
nn.Sequential(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), Linear(inner_dim, dim_out)
project_in, nn.Dropout(dropout), ops.Linear(inner_dim, dim_out)
)
def forward(self, x):
@@ -128,7 +123,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
)
@@ -138,8 +133,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
@@ -162,16 +157,16 @@ class SpatialSelfAttention(nn.Module):
self.in_channels = in_channels
self.norm = Normalize(in_channels)
self.q = Conv2d(
self.q = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.k = Conv2d(
self.k = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.v = Conv2d(
self.v = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.proj_out = Conv2d(
self.proj_out = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
@@ -218,12 +213,12 @@ class CrossAttention(nn.Module):
self.scale = dim_head**-0.5
self.heads = heads
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)
ops.Linear(inner_dim, query_dim), nn.Dropout(dropout)
)
self.backend = backend
@@ -309,12 +304,12 @@ class MemoryEfficientCrossAttention(nn.Module):
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)
ops.Linear(inner_dim, query_dim), nn.Dropout(dropout)
)
self.attention_op: Optional[Any] = None
@@ -442,9 +437,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:
#print(f"{self.__class__.__name__} is using checkpointing")
@@ -523,8 +518,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):
@@ -588,11 +583,11 @@ class SpatialTransformer(nn.Module):
inner_dim = n_heads * d_head
self.norm = Normalize(in_channels)
if not use_linear:
self.proj_in = Conv2d(
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)
self.transformer_blocks = nn.ModuleList(
[
@@ -612,11 +607,11 @@ class SpatialTransformer(nn.Module):
)
if not use_linear:
self.proj_out = zero_module(
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(Linear(in_channels, inner_dim))
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
def forward(self, x, context=None):
+3 -1
View File
@@ -8,6 +8,8 @@ from torchvision import models
from ..util import get_ckpt_path
import comfy.ops
ops = comfy.ops.manual_cast
class LPIPS(nn.Module):
# Learned perceptual metric
@@ -91,7 +93,7 @@ class NetLinLayer(nn.Module):
else []
)
layers += [
nn.Conv2d(chn_in, chn_out, 1, stride=1, padding=0, bias=False),
ops.Conv2d(chn_in, chn_out, 1, stride=1, padding=0, bias=False),
]
self.model = nn.Sequential(*layers)
@@ -3,7 +3,8 @@ import functools
import torch.nn as nn
from ..util import ActNorm
import comfy.ops
ops = comfy.ops.manual_cast
def weights_init(m):
classname = m.__class__.__name__
@@ -42,7 +43,7 @@ class NLayerDiscriminator(nn.Module):
kw = 4
padw = 1
sequence = [
nn.Conv2d(input_nc, ndf, kernel_size=kw, stride=2, padding=padw),
ops.Conv2d(input_nc, ndf, kernel_size=kw, stride=2, padding=padw),
nn.LeakyReLU(0.2, True),
]
nf_mult = 1
@@ -51,7 +52,7 @@ class NLayerDiscriminator(nn.Module):
nf_mult_prev = nf_mult
nf_mult = min(2**n, 8)
sequence += [
nn.Conv2d(
ops.Conv2d(
ndf * nf_mult_prev,
ndf * nf_mult,
kernel_size=kw,
@@ -66,7 +67,7 @@ class NLayerDiscriminator(nn.Module):
nf_mult_prev = nf_mult
nf_mult = min(2**n_layers, 8)
sequence += [
nn.Conv2d(
ops.Conv2d(
ndf * nf_mult_prev,
ndf * nf_mult,
kernel_size=kw,
@@ -79,7 +80,7 @@ class NLayerDiscriminator(nn.Module):
]
sequence += [
nn.Conv2d(ndf * nf_mult, 1, kernel_size=kw, stride=1, padding=padw)
ops.Conv2d(ndf * nf_mult, 1, kernel_size=kw, stride=1, padding=padw)
] # output 1 channel prediction map
self.main = nn.Sequential(*sequence)
+26 -31
View File
@@ -19,13 +19,8 @@ except:
from ...modules.attention import LinearAttention, MemoryEfficientCrossAttention
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
def get_timestep_embedding(timesteps, embedding_dim):
"""
@@ -54,7 +49,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
)
@@ -64,7 +59,7 @@ class Upsample(nn.Module):
super().__init__()
self.with_conv = with_conv
if self.with_conv:
self.conv = Conv2d(
self.conv = ops.Conv2d(
in_channels, in_channels, kernel_size=3, stride=1, padding=1
)
@@ -91,7 +86,7 @@ class Downsample(nn.Module):
self.with_conv = with_conv
if self.with_conv:
# no asymmetric padding in torch conv, must do it ourselves
self.conv = Conv2d(
self.conv = ops.Conv2d(
in_channels, in_channels, kernel_size=3, stride=2, padding=0
)
@@ -122,23 +117,23 @@ class ResnetBlock(nn.Module):
self.use_conv_shortcut = conv_shortcut
self.norm1 = Normalize(in_channels)
self.conv1 = Conv2d(
self.conv1 = ops.Conv2d(
in_channels, out_channels, kernel_size=3, stride=1, padding=1
)
if temb_channels > 0:
self.temb_proj = 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 = 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 = Conv2d(
self.conv_shortcut = ops.Conv2d(
in_channels, out_channels, kernel_size=3, stride=1, padding=1
)
else:
self.nin_shortcut = Conv2d(
self.nin_shortcut = ops.Conv2d(
in_channels, out_channels, kernel_size=1, stride=1, padding=0
)
@@ -178,16 +173,16 @@ class AttnBlock(nn.Module):
self.in_channels = in_channels
self.norm = Normalize(in_channels)
self.q = Conv2d(
self.q = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.k = Conv2d(
self.k = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.v = Conv2d(
self.v = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.proj_out = Conv2d(
self.proj_out = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
@@ -228,16 +223,16 @@ class MemoryEfficientAttnBlock(nn.Module):
self.in_channels = in_channels
self.norm = Normalize(in_channels)
self.q = Conv2d(
self.q = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.k = Conv2d(
self.k = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.v = Conv2d(
self.v = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.proj_out = Conv2d(
self.proj_out = ops.Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.attention_op: Optional[Any] = None
@@ -354,13 +349,13 @@ class Model(nn.Module):
self.temb = nn.Module()
self.temb.dense = nn.ModuleList(
[
Linear(self.ch, self.temb_ch),
Linear(self.temb_ch, self.temb_ch),
ops.Linear(self.ch, self.temb_ch),
ops.Linear(self.temb_ch, self.temb_ch),
]
)
# downsampling
self.conv_in = Conv2d(
self.conv_in = ops.Conv2d(
in_channels, self.ch, kernel_size=3, stride=1, padding=1
)
@@ -439,7 +434,7 @@ class Model(nn.Module):
# end
self.norm_out = Normalize(block_in)
self.conv_out = Conv2d(
self.conv_out = ops.Conv2d(
block_in, out_ch, kernel_size=3, stride=1, padding=1
)
@@ -526,7 +521,7 @@ class Encoder(nn.Module):
self.in_channels = in_channels
# downsampling
self.conv_in = Conv2d(
self.conv_in = ops.Conv2d(
in_channels, self.ch, kernel_size=3, stride=1, padding=1
)
@@ -577,7 +572,7 @@ class Encoder(nn.Module):
# end
self.norm_out = Normalize(block_in)
self.conv_out = Conv2d(
self.conv_out = ops.Conv2d(
block_in,
2 * z_channels if double_z else z_channels,
kernel_size=3,
@@ -660,7 +655,7 @@ class Decoder(nn.Module):
make_resblock_cls = self._make_resblock()
make_conv_cls = self._make_conv()
# z to block_in
self.conv_in = Conv2d(
self.conv_in = ops.Conv2d(
z_channels, block_in, kernel_size=3, stride=1, padding=1
)
@@ -719,7 +714,7 @@ class Decoder(nn.Module):
return ResnetBlock
def _make_conv(self) -> Callable:
return Conv2d
return ops.Conv2d
def get_last_layer(self, **kwargs):
return self.conv_out.weight
+7 -12
View File
@@ -19,6 +19,9 @@ import comfy.model_management
device = comfy.model_management.get_torch_device()
from contextlib import nullcontext
import comfy.ops
ops = comfy.ops.manual_cast
def make_beta_schedule(
schedule,
n_timestep,
@@ -279,24 +282,16 @@ class GroupNorm32(nn.GroupNorm):
# return super().forward(x.float()).type(x.dtype)
return super().forward(x)
class Conv2d(torch.nn.Conv2d):
def reset_parameters(self):
return None
class Linear(torch.nn.Linear):
def reset_parameters(self):
return None
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 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}")
@@ -304,7 +299,7 @@ def linear(*args, **kwargs):
"""
Create a linear module.
"""
return Linear(*args, **kwargs)
return ops.Linear(*args, **kwargs)
def avg_pool_nd(dims, *args, **kwargs):