force reset vae when changing methods
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user