From a8780cd1880066d9b9172f526b7ecda11f2be86b Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 6 Jul 2024 12:20:21 +0300 Subject: [PATCH] use more comfy ops --- SUPIR/modules/SUPIR_v0.py | 15 +++-- sgm/models/autoencoder.py | 9 ++- sgm/modules/attention.py | 63 +++++++++---------- sgm/modules/autoencoding/lpips/loss/lpips.py | 4 +- sgm/modules/autoencoding/lpips/model/model.py | 11 ++-- sgm/modules/diffusionmodules/model.py | 57 ++++++++--------- sgm/modules/diffusionmodules/util.py | 19 +++--- 7 files changed, 84 insertions(+), 94 deletions(-) diff --git a/SUPIR/modules/SUPIR_v0.py b/SUPIR/modules/SUPIR_v0.py index bae7dcb..6d91f33 100644 --- a/SUPIR/modules/SUPIR_v0.py +++ b/SUPIR/modules/SUPIR_v0.py @@ -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( diff --git a/sgm/models/autoencoder.py b/sgm/models/autoencoder.py index 4426c35..eaa2df7 100644 --- a/sgm/models/autoencoder.py +++ b/sgm/models/autoencoder.py @@ -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: diff --git a/sgm/modules/attention.py b/sgm/modules/attention.py index adf8f32..bb84628 100644 --- a/sgm/modules/attention.py +++ b/sgm/modules/attention.py @@ -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): diff --git a/sgm/modules/autoencoding/lpips/loss/lpips.py b/sgm/modules/autoencoding/lpips/loss/lpips.py index 3e34f3d..d145eac 100644 --- a/sgm/modules/autoencoding/lpips/loss/lpips.py +++ b/sgm/modules/autoencoding/lpips/loss/lpips.py @@ -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) diff --git a/sgm/modules/autoencoding/lpips/model/model.py b/sgm/modules/autoencoding/lpips/model/model.py index 66357d4..1c20151 100644 --- a/sgm/modules/autoencoding/lpips/model/model.py +++ b/sgm/modules/autoencoding/lpips/model/model.py @@ -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) diff --git a/sgm/modules/diffusionmodules/model.py b/sgm/modules/diffusionmodules/model.py index 2d0c99a..c78a7c8 100644 --- a/sgm/modules/diffusionmodules/model.py +++ b/sgm/modules/diffusionmodules/model.py @@ -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 diff --git a/sgm/modules/diffusionmodules/util.py b/sgm/modules/diffusionmodules/util.py index 65ccd0a..bc35757 100644 --- a/sgm/modules/diffusionmodules/util.py +++ b/sgm/modules/diffusionmodules/util.py @@ -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, @@ -278,25 +281,17 @@ class GroupNorm32(nn.GroupNorm): def forward(self, x): # 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):