Merge pull request #177 from hj-desperado/main
Add dtype selection in decoder, aligning with encoder
This commit is contained in:
+24
-9
@@ -180,6 +180,14 @@ class SUPIR_decode:
|
|||||||
"latents": ("LATENT",),
|
"latents": ("LATENT",),
|
||||||
"use_tiled_vae": ("BOOLEAN", {"default": True}),
|
"use_tiled_vae": ("BOOLEAN", {"default": True}),
|
||||||
"decoder_tile_size": ("INT", {"default": 512, "min": 64, "max": 8192, "step": 64}),
|
"decoder_tile_size": ("INT", {"default": 512, "min": 64, "max": 8192, "step": 64}),
|
||||||
|
"decoder_dtype": (
|
||||||
|
[
|
||||||
|
'bf16',
|
||||||
|
'fp32',
|
||||||
|
'auto'
|
||||||
|
], {
|
||||||
|
"default": 'auto'
|
||||||
|
}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -188,7 +196,7 @@ class SUPIR_decode:
|
|||||||
FUNCTION = "decode"
|
FUNCTION = "decode"
|
||||||
CATEGORY = "SUPIR"
|
CATEGORY = "SUPIR"
|
||||||
|
|
||||||
def decode(self, SUPIR_VAE, latents, use_tiled_vae, decoder_tile_size):
|
def decode(self, SUPIR_VAE, latents, use_tiled_vae, decoder_tile_size, decoder_dtype):
|
||||||
device = mm.get_torch_device()
|
device = mm.get_torch_device()
|
||||||
mm.unload_all_models()
|
mm.unload_all_models()
|
||||||
samples = latents["samples"]
|
samples = latents["samples"]
|
||||||
@@ -196,15 +204,22 @@ class SUPIR_decode:
|
|||||||
B, H, W, C = samples.shape
|
B, H, W, C = samples.shape
|
||||||
|
|
||||||
pbar = comfy.utils.ProgressBar(B)
|
pbar = comfy.utils.ProgressBar(B)
|
||||||
|
|
||||||
if mm.should_use_bf16():
|
if decoder_dtype == 'auto':
|
||||||
print("Decoder using bf16")
|
try:
|
||||||
dtype = torch.bfloat16
|
if mm.should_use_bf16():
|
||||||
|
print("Decoder using bf16")
|
||||||
|
vae_dtype = 'bf16'
|
||||||
|
else:
|
||||||
|
print("Decoder using fp32")
|
||||||
|
vae_dtype = 'fp32'
|
||||||
|
except:
|
||||||
|
raise AttributeError("ComfyUI version too old, can't autodetect properly. Set your dtypes manually.")
|
||||||
else:
|
else:
|
||||||
print("Decoder using fp32")
|
vae_dtype = decoder_dtype
|
||||||
dtype = torch.float32
|
print(f"Decoder using {vae_dtype}")
|
||||||
print("SUPIR decoder using", dtype)
|
|
||||||
|
dtype = convert_dtype(vae_dtype)
|
||||||
SUPIR_VAE.to(dtype).to(device)
|
SUPIR_VAE.to(dtype).to(device)
|
||||||
samples = samples.to(device)
|
samples = samples.to(device)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user