Fix for latest ComfyUI

This commit is contained in:
kijai
2026-03-15 17:02:02 +02:00
parent 0613a9239f
commit fe0d660f7c
2 changed files with 10 additions and 25 deletions
+9 -24
View File
@@ -510,7 +510,8 @@ SUPIR Tiles -node for preview to understand how the image is tiled.
_samples = self.sampler(denoiser, noised_z, cond=positive[i], uc=negative[i], x_center=sample.unsqueeze(0), control_scale=control_scale_end,
use_linear_control_scale=use_linear_control_scale, control_scale_start=control_scale_start)
except torch.cuda.OutOfMemoryError as e:
mm.free_memory(mm.get_total_memory(mm.get_torch_device()), mm.get_torch_device())
SUPIR_model = None
@@ -912,11 +913,7 @@ high_vram: uses Accelerate to load weights to GPU, slightly faster model loading
sdxl_state_dict = model.model.state_dict_for_saving(model.model.diffusion_model.state_dict(), vae_state_dict=vae.get_sd())
except:
sdxl_state_dict = model.model.state_dict_for_saving(None, vae.get_sd(), None)
if is_accelerate_available:
for key in sdxl_state_dict:
set_module_tensor_to_device(self.model, key, device=device, dtype=dtype, value=sdxl_state_dict[key])
else:
self.model.load_state_dict(sdxl_state_dict, strict=False)
self.model.load_state_dict(sdxl_state_dict, strict=False)
if fp8_unet:
self.model.model.to(torch.float8_e4m3fn)
else:
@@ -924,7 +921,7 @@ high_vram: uses Accelerate to load weights to GPU, slightly faster model loading
del sdxl_state_dict
pbar.update(1)
except:
raise Exception("Failed to load SDXL model")
raise Exception("Failed to load SDXL model")
gc.collect()
mm.soft_empty_cache()
#first clip model from SDXL checkpoint
@@ -976,11 +973,7 @@ high_vram: uses Accelerate to load weights to GPU, slightly faster model loading
try:
print(f'Attempting to load SUPIR model: [{SUPIR_MODEL_PATH}]')
supir_state_dict = load_state_dict(SUPIR_MODEL_PATH)
if "Q" not in supir_model or not is_accelerate_available: #I don't know why this doesn't work with the Q model.
for key in supir_state_dict:
set_module_tensor_to_device(self.model, key, device=device, dtype=dtype, value=supir_state_dict[key])
else:
self.model.load_state_dict(supir_state_dict, strict=False)
self.model.load_state_dict(supir_state_dict, strict=False)
if fp8_unet:
self.model.model.to(torch.float8_e4m3fn)
else:
@@ -992,7 +985,7 @@ high_vram: uses Accelerate to load weights to GPU, slightly faster model loading
mm.soft_empty_cache()
return (self.model, self.model.first_stage_model,)
class SUPIR_model_loader_v2_clip:
@classmethod
def INPUT_TYPES(s):
@@ -1098,11 +1091,7 @@ high_vram: uses Accelerate to load weights to GPU, slightly faster model loading
sdxl_state_dict = model.model.state_dict_for_saving(model.model.diffusion_model.state_dict(), vae_state_dict=vae.get_sd())
except:
sdxl_state_dict = model.model.state_dict_for_saving(None, vae.get_sd(), None)
if is_accelerate_available:
for key in sdxl_state_dict:
set_module_tensor_to_device(self.model, key, device=device, dtype=dtype, value=sdxl_state_dict[key])
else:
self.model.load_state_dict(sdxl_state_dict, strict=False)
self.model.load_state_dict(sdxl_state_dict, strict=False)
if fp8_unet:
self.model.model.to(torch.float8_e4m3fn)
else:
@@ -1110,7 +1099,7 @@ high_vram: uses Accelerate to load weights to GPU, slightly faster model loading
del sdxl_state_dict
pbar.update(1)
except:
raise Exception("Failed to load SDXL model")
raise Exception("Failed to load SDXL model")
gc.collect()
mm.soft_empty_cache()
#first clip model from SDXL checkpoint
@@ -1168,11 +1157,7 @@ high_vram: uses Accelerate to load weights to GPU, slightly faster model loading
try:
print(f'Attempting to load SUPIR model: [{SUPIR_MODEL_PATH}]')
supir_state_dict = load_state_dict(SUPIR_MODEL_PATH)
if "Q" not in supir_model or not is_accelerate_available: #I don't know why this doesn't work with the Q model.
for key in supir_state_dict:
set_module_tensor_to_device(self.model, key, device=device, dtype=dtype, value=supir_state_dict[key])
else:
self.model.load_state_dict(supir_state_dict, strict=False)
self.model.load_state_dict(supir_state_dict, strict=False)
if fp8_unet:
self.model.model.to(torch.float8_e4m3fn)
else:
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-supir"
description = "Wrapper nodes to use SUPIR upscaling process in ComfyUI"
version = "1.0.3"
version = "1.0.4"
license = { file = "LICENSE" }
dependencies = ["transformers>=4.28.1", "fsspec>=2023.4.0", "kornia>=0.6.9", "open-clip-torch>=2.24.0", "Pillow>=9.4.0", "pytorch-lightning>=2.2.1", "omegaconf", "accelerate"]