diff --git a/SUPIR/models/SUPIR_model.py b/SUPIR/models/SUPIR_model.py index df35fe3..e0669fb 100644 --- a/SUPIR/models/SUPIR_model.py +++ b/SUPIR/models/SUPIR_model.py @@ -7,6 +7,7 @@ import random from ...SUPIR.utils.colorfix import wavelet_reconstruction, adaptive_instance_normalization from pytorch_lightning import seed_everything from ...SUPIR.utils.tilevae import VAEHook +from ...SUPIR.util import convert_dtype from contextlib import nullcontext import comfy.model_management @@ -20,23 +21,8 @@ class SUPIRModel(DiffusionEngine): self.first_stage_model.denoise_encoder = copy.deepcopy(self.first_stage_model.encoder) self.sampler_config = kwargs['sampler_config'] - assert (ae_dtype in ['fp32', 'fp16', 'bf16']) and (diffusion_dtype in ['fp32', 'fp16', 'bf16']) - if ae_dtype == 'fp32': - ae_dtype = torch.float32 - elif ae_dtype == 'fp16': - raise RuntimeError('fp16 cause NaN in AE') - elif ae_dtype == 'bf16': - ae_dtype = torch.bfloat16 - - if diffusion_dtype == 'fp32': - diffusion_dtype = torch.float32 - elif diffusion_dtype == 'fp16': - diffusion_dtype = torch.float16 - elif diffusion_dtype == 'bf16': - diffusion_dtype = torch.bfloat16 - - self.ae_dtype = ae_dtype - self.model.dtype = diffusion_dtype + self.ae_dtype = convert_dtype(ae_dtype) + self.model.dtype = convert_dtype(diffusion_dtype) self.p_p = p_p self.n_p = n_p @@ -72,7 +58,6 @@ class SUPIRModel(DiffusionEngine): @torch.no_grad() def decode_first_stage(self, z): z = 1.0 / self.scale_factor * z - #with torch.autocast(device, dtype=self.ae_dtype): autocast_condition = (self.ae_dtype == torch.float16 or self.ae_dtype == torch.bfloat16) and not comfy.model_management.is_device_mps(device) with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=self.ae_dtype) if autocast_condition else nullcontext(): out = self.first_stage_model.decode(z) diff --git a/nodes.py b/nodes.py index 6b021ef..f2179be 100644 --- a/nodes.py +++ b/nodes.py @@ -121,6 +121,7 @@ class SUPIR_Upscale: "captions": ("STRING", {"forceInput": True, "multiline": False, "default": "", }), "diffusion_dtype": ( [ + 'float8_e4m3fn', 'fp16', 'bf16', 'fp32', @@ -140,6 +141,8 @@ class SUPIR_Upscale: "use_tiled_sampling": ("BOOLEAN", {"default": False}), "sampler_tile_size": ("INT", {"default": 1024, "min": 64, "max": 4096, "step": 32}), "sampler_tile_stride": ("INT", {"default": 512, "min": 32, "max": 2048, "step": 32}), + "fp8_unet": ("BOOLEAN", {"default": False}), + "fp8_vae": ("BOOLEAN", {"default": False}), } } @@ -153,7 +156,7 @@ class SUPIR_Upscale: encoder_tile_size_pixels, decoder_tile_size_latent, control_scale, cfg_scale_start, control_scale_start, restoration_scale, keep_model_loaded, a_prompt, n_prompt, sdxl_model, supir_model, use_tiled_vae, use_tiled_sampling=False, sampler_tile_size=128, sampler_tile_stride=64, captions="", diffusion_dtype="auto", - encoder_dtype="auto", batch_size=1): + encoder_dtype="auto", batch_size=1, fp8_unet=False, fp8_vae=False): device = mm.get_torch_device() mm.unload_all_models() @@ -172,6 +175,7 @@ class SUPIR_Upscale: 'use_tiled_vae': use_tiled_vae, 'supir_model': supir_model, 'use_tiled_sampling': use_tiled_sampling, + 'fp8_unet': fp8_unet } if diffusion_dtype == 'auto': @@ -286,6 +290,10 @@ class SUPIR_Upscale: try: self.model.to(dtype) + if fp8_unet: + self.model.model.to(torch.float8_e4m3fn) + if fp8_vae: + self.model.first_stage_model.to(torch.float8_e4m3fn) self.model.to(device) except Exception as e: print("Failed to move model to device")