From 38f5a74926f24687415cc5018bba4a93bd955c00 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 16 Jan 2024 19:58:39 +0200 Subject: [PATCH] force reset vae when changing methods --- ldm/models/diffusion/ddpm_ccsr_stage2.py | 8 +++++++- nodes.py | 2 ++ 2 files changed, 9 insertions(+), 1 deletion(-) diff --git a/ldm/models/diffusion/ddpm_ccsr_stage2.py b/ldm/models/diffusion/ddpm_ccsr_stage2.py index 80e0d04..895dd06 100644 --- a/ldm/models/diffusion/ddpm_ccsr_stage2.py +++ b/ldm/models/diffusion/ddpm_ccsr_stage2.py @@ -737,7 +737,13 @@ class LatentDiffusion(DDPM): self.first_stage_model.decoder.forward = VAEHook( decoder, decoder_tile_size, is_decoder=True, fast_decoder=fast_decoder, fast_encoder=fast_encoder, color_fix=color_fix, to_gpu=vae_to_gpu) - + + def reset_encoder_decoder(self): + # Restore the original forward methods + 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 def instantiate_first_stage(self, config): model = instantiate_from_config(config) self.first_stage_model = model.eval() diff --git a/nodes.py b/nodes.py index e90a263..c2e29f1 100644 --- a/nodes.py +++ b/nodes.py @@ -115,6 +115,7 @@ class CCSR_Upscale: for i in range(batch_size): img = resized_image[i].unsqueeze(0).to(device) if sampling_method == 'ccsr_tiled_mixdiff': + self.model.reset_encoder_decoder() print("Using tiled mixdiff") samples = sampler.sample_with_mixdiff_ccsr( empty_text_embed, tile_size=tile_size, tile_stride=tile_stride, @@ -134,6 +135,7 @@ class CCSR_Upscale: color_fix_type=color_fix_type ) else: + self.model.reset_encoder_decoder() print("no tiling") samples = sampler.sample_ccsr( empty_text_embed, steps=steps, t_max=t_max, t_min=t_min, shape=shape, cond_img=img,