From 903b087896471fca84fbb6f43b2196fc4b375add Mon Sep 17 00:00:00 2001 From: Paul Date: Mon, 24 Jun 2024 22:18:19 +0000 Subject: [PATCH] Add list sampler node --- node.py | 96 +++++++++++++++++++++++++++++++++++++++----------- pyproject.toml | 2 +- 2 files changed, 77 insertions(+), 21 deletions(-) diff --git a/node.py b/node.py index f8bf756..03c7f5a 100644 --- a/node.py +++ b/node.py @@ -4,11 +4,35 @@ import os import random -class PromptFromTemplate: +class BasePrompt: PROMPT_LISTS_DIR = os.path.join( os.path.dirname(os.path.abspath(__file__)), "prompt-lists" ) + def __init__(self): + self.all_lists = {} + for list_name in self.get_all_available_lists(): + self.all_lists[list_name] = self.load_list(list_name) + + @classmethod + def get_all_available_lists(cls): + lists_path = os.path.join(cls.PROMPT_LISTS_DIR, "lists.json") + with open(lists_path, "r", encoding="utf-8") as file: + lists = json.load(file) + return lists + + def load_list(self, list_name): + hyphenated_list = re.sub(r"([a-z])([A-Z])", r"\1-\2", list_name).lower() + category, name = hyphenated_list.split(".") + lists_path = os.path.join(self.PROMPT_LISTS_DIR, f"lists/{category}/{name}.yml") + with open(lists_path, "r", encoding="utf-8") as file: + content = file.read() + list_data = content.split("---")[2].strip().split("\n") + + return list_data + + +class PromptFromTemplate(BasePrompt): @classmethod def INPUT_TYPES(cls): return { @@ -21,11 +45,6 @@ class PromptFromTemplate: } } - def __init__(self): - self.all_lists = {} - for list_name in self.get_all_available_lists(): - self.all_lists[list_name] = self.load_list(list_name) - RETURN_TYPES = ("STRING",) FUNCTION = "generate_prompt_from_template" CATEGORY = "Prompter" @@ -57,29 +76,66 @@ class PromptFromTemplate: print(prompt) return (prompt,) - def get_all_available_lists(self): - lists_path = os.path.join(self.PROMPT_LISTS_DIR, "lists.json") - with open(lists_path, "r", encoding="utf-8") as file: - lists = json.load(file) - return lists - def get_random_list(self): return random.choice(list(self.all_lists.keys())) def get_random_items_from_list(self, list_name, item_count): return random.sample(self.all_lists[list_name], item_count) - def load_list(self, list_name): - hyphenated_list = re.sub(r"([a-z])([A-Z])", r"\1-\2", list_name).lower() - category, name = hyphenated_list.split(".") - lists_path = os.path.join(self.PROMPT_LISTS_DIR, f"lists/{category}/{name}.yml") - with open(lists_path, "r", encoding="utf-8") as file: - content = file.read() - list_data = content.split("---")[2].strip().split("\n") - return list_data +class PromptListSampler(BasePrompt): + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "list_name": (cls.get_all_available_lists(),), + "mode": (["random", "sequential"], {"default": "sequential"}), + "number_of_items": ("INT", {"default": 1, "min": 1}), + "join_with": ("STRING", {"default": ", "}), + "index": ( + "INT", + { + "default": 0, + "min": 0, + "max": 0xFFFFFFFFFFFFFFFF, + "control_after_generate": True, + }, + ), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}), + } + } + + RETURN_TYPES = ("STRING",) + FUNCTION = "get_item_from_list" + CATEGORY = "Prompter" + + def get_item_from_list( + self, + list_name, + index=0, + number_of_items=1, + mode="sequential", + join_with=", ", + seed=0, + ): + random.seed(seed) + + if list_name not in self.all_lists: + return f"[{list_name}]" + + if mode == "random": + items = random.sample( + self.all_lists[list_name], + min(number_of_items, len(self.all_lists[list_name])), + ) + else: + index = index % len(self.all_lists[list_name]) + items = self.all_lists[list_name][index : index + number_of_items] + + return (join_with.join(items),) NODE_CLASS_MAPPINGS = { "Prompt from template 🪴": PromptFromTemplate, + "List sampler 🪴": PromptListSampler, } diff --git a/pyproject.toml b/pyproject.toml index 5222a2c..0985067 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui-prompter-fofrai" description = "A prompt helper for ComfyUI, based on prompter.fofr.ai" -version = "1.0.0" +version = "1.0.1" license = "LICENSE" [project.urls]