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