From 8df063060e48055a3beacc424a1dcfc6e640e228 Mon Sep 17 00:00:00 2001 From: asagi4 <130366179+asagi4@users.noreply.github.com> Date: Sun, 8 Dec 2024 17:20:42 +0200 Subject: [PATCH] Use folder_paths to find npnet models --- README.md | 2 +- __init__.py | 15 ++++++++++++--- 2 files changed, 13 insertions(+), 4 deletions(-) diff --git a/README.md b/README.md index c4f034f..22aa01b 100644 --- a/README.md +++ b/README.md @@ -3,7 +3,7 @@ A very barebones mostly-copypaste implementation of https://github.com/xie-lab-ml/Golden-Noise-for-Diffusion-Models ## Requirements -You need the pre-trained weights for your model. +You need the pre-trained weights for your model. Download and place them under `models/npnet` in your ComfyUI folder, or add an extra path in `extra_model_paths.yaml` for the `npnet` type. You can find safetensors-converted weights at https://huggingface.co/asagi4/NPNet diff --git a/__init__.py b/__init__.py index 4d4f812..65ef6ad 100644 --- a/__init__.py +++ b/__init__.py @@ -9,9 +9,11 @@ from diffusers.models.normalization import AdaGroupNorm from timm.layers import use_fused_attn - from comfy.utils import common_upscale +import folder_paths +import os.path + class Attention(nn.Module): fused_attn = True @@ -203,11 +205,17 @@ class NPNetGoldenNoise: @classmethod def INPUT_TYPES(s): + if "npnet" not in folder_paths.folder_names_and_paths: + folder_paths.folder_names_and_paths["npnet"] = ( + [os.path.join(folder_paths.models_dir, "npnet")], + {".pth", ".safetensors"}, + ) + return { "required": { "noise": ("NOISE",), "prompt": ("CONDITIONING",), - "model_path": ("STRING", {"default": "/path/to/sdxl.pth"}), + "model": (folder_paths.get_filename_list("npnet"),), "device": (["cuda", "cpu"],), } } @@ -235,7 +243,8 @@ class NPNetGoldenNoise: r = common_upscale(r, orig_shape[-1], orig_shape[-2], "nearest-exact", "disabled") return r - def doit(self, noise, prompt, model_path, device): + def doit(self, noise, prompt, model, device): + model_path = folder_paths.get_full_path("npnet", model) if self.npnet is None or self.npnet.pretrained_path != model_path: print("Loading NPNet from", model_path) self.npnet = NPNet(model_path, device=device)