Fix model changing

This commit is contained in:
kijai
2024-02-29 15:26:11 +02:00
parent 07906ead59
commit 437beaeafb
2 changed files with 32 additions and 17 deletions
+25 -13
View File
@@ -152,19 +152,31 @@ 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):
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)
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)
if __name__ == '__main__':
+7 -4
View File
@@ -65,7 +65,8 @@ class SUPIR_Upscale:
device = comfy.model_management.get_torch_device()
image = image.to(device)
batch_size = image.shape[0]
self.current_sdxl_model = None
SUPIR_MODEL_PATH = folder_paths.get_full_path("checkpoints", supir_model)
SDXL_MODEL_PATH = folder_paths.get_full_path("checkpoints", sdxl_model)
@@ -87,8 +88,8 @@ class SUPIR_Upscale:
vae_dtype = 'fp32'
model_dtype = 'fp32'
if not hasattr(self, "model") or self.model is None:
if not hasattr(self, "model") or self.model is None or self.current_sdxl_model != sdxl_model:
self.current_sdxl_model = sdxl_model
config = OmegaConf.load(config_path)
config.model.params.ae_dtype = vae_dtype
config.model.params.diffusion_dtype = model_dtype
@@ -101,7 +102,9 @@ 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)
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)
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():