This commit is contained in:
kijai
2024-02-29 18:27:01 +02:00
parent ed4577cfbb
commit 41fa80e4b4
6 changed files with 66 additions and 61 deletions
+13 -25
View File
@@ -152,31 +152,19 @@ class SUPIRModel(DiffusionEngine):
samples = adaptive_instance_normalization(samples, x_stage1)
return samples
def init_tile_vae(self, encoder_tile_size=512, decoder_tile_size=64, reset=False):
if reset:
# Reset the models to their original forward methods
if hasattr(self.first_stage_model.denoise_encoder, 'original_forward'):
self.first_stage_model.denoise_encoder.forward = self.first_stage_model.denoise_encoder.original_forward
if hasattr(self.first_stage_model.encoder, 'original_forward'):
self.first_stage_model.encoder.forward = self.first_stage_model.encoder.original_forward
if hasattr(self.first_stage_model.decoder, 'original_forward'):
self.first_stage_model.decoder.forward = self.first_stage_model.decoder.original_forward
else:
# Save the original forward methods
self.first_stage_model.denoise_encoder.original_forward = self.first_stage_model.denoise_encoder.forward
self.first_stage_model.encoder.original_forward = self.first_stage_model.encoder.forward
self.first_stage_model.decoder.original_forward = self.first_stage_model.decoder.forward
# Apply the VAEHook to the models
self.first_stage_model.denoise_encoder.forward = VAEHook(
self.first_stage_model.denoise_encoder, encoder_tile_size, is_decoder=False, fast_decoder=False,
fast_encoder=False, color_fix=False, to_gpu=True)
self.first_stage_model.encoder.forward = VAEHook(
self.first_stage_model.encoder, encoder_tile_size, is_decoder=False, fast_decoder=False,
fast_encoder=False, color_fix=False, to_gpu=True)
self.first_stage_model.decoder.forward = VAEHook(
self.first_stage_model.decoder, decoder_tile_size, is_decoder=True, fast_decoder=False,
fast_encoder=False, color_fix=False, to_gpu=True)
def init_tile_vae(self, encoder_tile_size=512, decoder_tile_size=64):
self.first_stage_model.denoise_encoder.original_forward = self.first_stage_model.denoise_encoder.forward
self.first_stage_model.encoder.original_forward = self.first_stage_model.encoder.forward
self.first_stage_model.decoder.original_forward = self.first_stage_model.decoder.forward
self.first_stage_model.denoise_encoder.forward = VAEHook(
self.first_stage_model.denoise_encoder, encoder_tile_size, is_decoder=False, fast_decoder=False,
fast_encoder=False, color_fix=False, to_gpu=True)
self.first_stage_model.encoder.forward = VAEHook(
self.first_stage_model.encoder, encoder_tile_size, is_decoder=False, fast_decoder=False,
fast_encoder=False, color_fix=False, to_gpu=True)
self.first_stage_model.decoder.forward = VAEHook(
self.first_stage_model.decoder, decoder_tile_size, is_decoder=True, fast_decoder=False,
fast_encoder=False, color_fix=False, to_gpu=True)
if __name__ == '__main__':
+9 -3
View File
@@ -121,7 +121,13 @@ def test_for_nans(x, where):
raise NansException(message)
class Conv2d(torch.nn.Conv2d):
def reset_parameters(self):
return None
class Linear(torch.nn.Linear):
def reset_parameters(self):
return None
@lru_cache
def first_time_calculation():
"""
@@ -130,9 +136,9 @@ def first_time_calculation():
"""
x = torch.zeros((1, 1)).to(device, dtype)
linear = torch.nn.Linear(1, 1).to(device, dtype)
linear = Linear(1, 1).to(device, dtype)
linear(x)
x = torch.zeros((1, 1, 3, 3)).to(device, dtype)
conv2d = torch.nn.Conv2d(1, 1, (3, 3)).to(device, dtype)
conv2d = Conv2d(1, 1, (3, 3)).to(device, dtype)
conv2d(x)
+6 -5
View File
@@ -17,6 +17,7 @@ class SUPIR_Upscale:
self.current_sdxl_model = None
self.current_diffusion_dtype = None
self.current_encoder_dtype = None
self.tiled_vae_state = None
upscale_methods = ["nearest-exact", "bilinear", "area", "bicubic", "lanczos"]
@classmethod
def INPUT_TYPES(s):
@@ -123,7 +124,7 @@ class SUPIR_Upscale:
vae_dtype = encoder_dtype
print(f"Encoder using using {vae_dtype}")
if not hasattr(self, "model") or self.model is None or self.current_sdxl_model != sdxl_model or self.current_diffusion_dtype != diffusion_dtype or self.current_encoder_dtype != encoder_dtype:
if not hasattr(self, "model") or self.model is None or self.current_sdxl_model != sdxl_model or self.current_diffusion_dtype != diffusion_dtype or self.current_encoder_dtype != encoder_dtype or self.tiled_vae_state != use_tiled_vae:
self.current_diffusion_dtype = diffusion_dtype
self.current_encoder_dtype = encoder_dtype
self.current_sdxl_model = sdxl_model
@@ -137,10 +138,10 @@ class SUPIR_Upscale:
self.model.load_state_dict(sdxl_state_dict, strict=False)
self.model.to(device).to(dtype)
if use_tiled_vae:
self.model.init_tile_vae(encoder_tile_size=encoder_tile_size_pixels, decoder_tile_size=decoder_tile_size_latent, reset=False)
else:
self.model.init_tile_vae(encoder_tile_size=encoder_tile_size_pixels, decoder_tile_size=decoder_tile_size_latent, reset=True)
if use_tiled_vae:
self.tiled_vae_state = True
self.model.init_tile_vae(encoder_tile_size=encoder_tile_size_pixels, decoder_tile_size=decoder_tile_size_latent)
autocast_condition = dtype == torch.float16 or torch.bfloat16 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():
+6 -3
View File
@@ -14,7 +14,10 @@ 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
class AbstractAutoencoder(pl.LightningModule):
"""
This is the base class for all autoencoders, including image autoencoders, image autoencoders with discriminators,
@@ -294,8 +297,8 @@ class AutoencoderKL(AutoencodingEngine):
assert ddconfig["double_z"]
self.encoder = Encoder(**ddconfig)
self.decoder = Decoder(**ddconfig)
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 = Conv2d(2 * ddconfig["z_channels"], 2 * embed_dim, 1)
self.post_quant_conv = Conv2d(embed_dim, ddconfig["z_channels"], 1)
self.embed_dim = embed_dim
if ckpt_path is not None:
+8 -4
View File
@@ -10,6 +10,10 @@ from einops import rearrange, repeat
from packaging import version
from torch import nn
class Conv2d(torch.nn.Conv2d):
def reset_parameters(self):
return None
if version.parse(torch.__version__) >= version.parse("2.0.0"):
SDP_IS_AVAILABLE = True
from torch.backends.cuda import SDPBackend, sdp_kernel
@@ -154,16 +158,16 @@ class SpatialSelfAttention(nn.Module):
self.in_channels = in_channels
self.norm = Normalize(in_channels)
self.q = torch.nn.Conv2d(
self.q = Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.k = torch.nn.Conv2d(
self.k = Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.v = torch.nn.Conv2d(
self.v = Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.proj_out = torch.nn.Conv2d(
self.proj_out = Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
+24 -21
View File
@@ -19,7 +19,10 @@ except:
from ...modules.attention import LinearAttention, MemoryEfficientCrossAttention
class Conv2d(torch.nn.Conv2d):
def reset_parameters(self):
return None
def get_timestep_embedding(timesteps, embedding_dim):
"""
This matches the implementation in Denoising Diffusion Probabilistic Models:
@@ -57,7 +60,7 @@ class Upsample(nn.Module):
super().__init__()
self.with_conv = with_conv
if self.with_conv:
self.conv = torch.nn.Conv2d(
self.conv = Conv2d(
in_channels, in_channels, kernel_size=3, stride=1, padding=1
)
@@ -74,7 +77,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 = torch.nn.Conv2d(
self.conv = Conv2d(
in_channels, in_channels, kernel_size=3, stride=2, padding=0
)
@@ -105,23 +108,23 @@ class ResnetBlock(nn.Module):
self.use_conv_shortcut = conv_shortcut
self.norm1 = Normalize(in_channels)
self.conv1 = torch.nn.Conv2d(
self.conv1 = 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.norm2 = Normalize(out_channels)
self.dropout = torch.nn.Dropout(dropout)
self.conv2 = torch.nn.Conv2d(
self.conv2 = 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 = Conv2d(
in_channels, out_channels, kernel_size=3, stride=1, padding=1
)
else:
self.nin_shortcut = torch.nn.Conv2d(
self.nin_shortcut = Conv2d(
in_channels, out_channels, kernel_size=1, stride=1, padding=0
)
@@ -161,16 +164,16 @@ class AttnBlock(nn.Module):
self.in_channels = in_channels
self.norm = Normalize(in_channels)
self.q = torch.nn.Conv2d(
self.q = Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.k = torch.nn.Conv2d(
self.k = Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.v = torch.nn.Conv2d(
self.v = Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.proj_out = torch.nn.Conv2d(
self.proj_out = Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
@@ -211,16 +214,16 @@ class MemoryEfficientAttnBlock(nn.Module):
self.in_channels = in_channels
self.norm = Normalize(in_channels)
self.q = torch.nn.Conv2d(
self.q = Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.k = torch.nn.Conv2d(
self.k = Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.v = torch.nn.Conv2d(
self.v = Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.proj_out = torch.nn.Conv2d(
self.proj_out = Conv2d(
in_channels, in_channels, kernel_size=1, stride=1, padding=0
)
self.attention_op: Optional[Any] = None
@@ -343,7 +346,7 @@ class Model(nn.Module):
)
# downsampling
self.conv_in = torch.nn.Conv2d(
self.conv_in = Conv2d(
in_channels, self.ch, kernel_size=3, stride=1, padding=1
)
@@ -422,7 +425,7 @@ class Model(nn.Module):
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(
self.conv_out = Conv2d(
block_in, out_ch, kernel_size=3, stride=1, padding=1
)
@@ -509,7 +512,7 @@ class Encoder(nn.Module):
self.in_channels = in_channels
# downsampling
self.conv_in = torch.nn.Conv2d(
self.conv_in = Conv2d(
in_channels, self.ch, kernel_size=3, stride=1, padding=1
)
@@ -560,7 +563,7 @@ class Encoder(nn.Module):
# end
self.norm_out = Normalize(block_in)
self.conv_out = torch.nn.Conv2d(
self.conv_out = Conv2d(
block_in,
2 * z_channels if double_z else z_channels,
kernel_size=3,
@@ -643,7 +646,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 = Conv2d(
z_channels, block_in, kernel_size=3, stride=1, padding=1
)
@@ -702,7 +705,7 @@ class Decoder(nn.Module):
return ResnetBlock
def _make_conv(self) -> Callable:
return torch.nn.Conv2d
return Conv2d
def get_last_layer(self, **kwargs):
return self.conv_out.weight