fixex
This commit is contained in:
+13
-25
@@ -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__':
|
||||
|
||||
@@ -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)
|
||||
@@ -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():
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user