force reset vae when changing methods

This commit is contained in:
kijai
2024-01-16 19:58:39 +02:00
parent cbb9ad8f0b
commit 38f5a74926
2 changed files with 9 additions and 1 deletions
+7 -1
View File
@@ -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()
+2
View File
@@ -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,