Support SDXL latents directly

Bypassing need for SUPIR VAE
This commit is contained in:
kijai
2024-06-27 16:40:45 +03:00
parent c257cce555
commit 006754633a
2 changed files with 41 additions and 16 deletions
+1 -1
View File
@@ -21,7 +21,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"SUPIR_encode": "SUPIR Encode",
"SUPIR_decode": "SUPIR Decode",
"SUPIR_conditioner": "SUPIR Conditioner",
"SUPIR_tiles": "SUPIR Tiles",
"SUPIR_tiles": "SUPIR Tiles Preview",
"SUPIR_model_loader_v2": "SUPIR Model Loader (v2)",
"SUPIR_model_loader_v2_clip": "SUPIR Model Loader (v2) (Clip)"
}
+35 -10
View File
@@ -192,14 +192,20 @@ class SUPIR_decode:
device = mm.get_torch_device()
mm.unload_all_models()
samples = latents["samples"]
dtype = SUPIR_VAE.dtype
orig_H, orig_W = latents["original_size"]
B, H, W, C = samples.shape
pbar = comfy.utils.ProgressBar(B)
SUPIR_VAE.to(device)
if mm.should_use_bf16():
print("Decoder using bf16")
dtype = torch.bfloat16
else:
print("Decoder using fp32")
dtype = torch.float32
print("SUPIR decoder using", dtype)
SUPIR_VAE.to(dtype).to(device)
samples = samples.to(device)
if use_tiled_vae:
@@ -220,11 +226,14 @@ class SUPIR_decode:
autocast_condition = (dtype != torch.float32) 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():
sample = 1.0 / 0.13025 * sample
decoded_image = SUPIR_VAE.decode(sample.unsqueeze(0)).float()
decoded_image = SUPIR_VAE.decode(sample.unsqueeze(0))
out.append(decoded_image)
pbar.update(1)
decoded_out= torch.cat(out, dim=0)
decoded_out= torch.cat(out, dim=0).float()
if "original_size" in latents and latents["original_size"] is not None:
orig_H, orig_W = latents["original_size"]
if decoded_out.shape[2] != orig_H or decoded_out.shape[3] != orig_W:
print("Restoring original dimensions: ", orig_W,"x",orig_H)
decoded_out = F.interpolate(decoded_out, size=(orig_H, orig_W), mode="bicubic")
@@ -360,9 +369,9 @@ class SUPIR_sample:
"EDM_s_churn": ("INT", {"default": 5, "min": 0, "max": 40, "step": 1}),
"s_noise": ("FLOAT", {"default": 1.003, "min": 1.0, "max": 1.1, "step": 0.001}),
"DPMPP_eta": ("FLOAT", {"default": 1.0, "min": 0, "max": 10.0, "step": 0.01}),
"control_scale_start": ("FLOAT", {"default": 1.0, "min": 0, "max": 10.0, "step": 0.05}),
"control_scale_end": ("FLOAT", {"default": 1.0, "min": 0, "max": 10.0, "step": 0.05}),
"restore_cfg": ("FLOAT", {"default": -1.0, "min": -1.0, "max": 20.0, "step": 0.05}),
"control_scale_start": ("FLOAT", {"default": 1.0, "min": 0, "max": 10.0, "step": 0.01}),
"control_scale_end": ("FLOAT", {"default": 1.0, "min": 0, "max": 10.0, "step": 0.01}),
"restore_cfg": ("FLOAT", {"default": -1.0, "min": -1.0, "max": 20.0, "step": 0.01}),
"keep_model_loaded": ("BOOLEAN", {"default": False}),
"sampler": (
[
@@ -483,7 +492,8 @@ SUPIR Tiles -node for preview to understand how the image is tiled.
noised_z = torch.randn_like(sample.unsqueeze(0), device=samples.device)
else:
print("Using latent from input")
noised_z = sample.unsqueeze(0) * 0.13025
noised_z = torch.randn_like(sample.unsqueeze(0), device=samples.device)
noised_z += sample.unsqueeze(0)
if len(positive) != len(samples):
print("Tiled sampling")
_samples = self.sampler(denoiser, noised_z, cond=positive, uc=negative, x_center=sample.unsqueeze(0), control_scale=control_scale_end,
@@ -518,6 +528,9 @@ SUPIR Tiles -node for preview to understand how the image is tiled.
else:
samples_out_stacked = torch.stack(out, dim=0)
if original_size is None:
samples_out_stacked = samples_out_stacked / 0.13025
return ({"samples":samples_out_stacked, "original_size": original_size},)
class SUPIR_conditioner:
@@ -555,7 +568,14 @@ If a list of captions is given and it matches the incoming image batch, each ima
device = mm.get_torch_device()
mm.soft_empty_cache()
if "original_size" in latents:
original_size = latents["original_size"]
samples = latents["samples"]
else:
original_size = None
samples = latents["samples"] * 0.13025
N, H, W, C = samples.shape
import copy
@@ -623,7 +643,12 @@ If a list of captions is given and it matches the incoming image batch, each ima
SUPIR_model.conditioner.to('cpu')
return ({"cond": c, "original_size":latents["original_size"]}, {"uncond": uc},)
if "original_size" in latents:
original_size = latents["original_size"]
else:
original_size = None
return ({"cond": c, "original_size":original_size}, {"uncond": uc},)
class SUPIR_model_loader:
@classmethod