first commit

This commit is contained in:
Nakarin Satayasai
2023-11-01 22:20:03 +07:00
commit 927cb8ffb9
17 changed files with 857 additions and 0 deletions
+26
View File
@@ -0,0 +1,26 @@
from .py.io.multi_lora_loader import *
from .py.io.csv_prompts_loader import *
from .py.conditioning.random_prompt import *
from .py.conditioning.combine_prompts import *
# from .malefish import init
NODE_CLASS_MAPPINGS = {
"MultiLoraLoader": MultiLoraLoader,
"RandomPrompt": RandomPrompt,
"CombinePrompt": CombinePrompts,
"CSVPromptsLoader": CSVPromptsLoader
}
# A dictionary that contains the friendly/humanly readable titles for the nodes
NODE_DISPLAY_NAME_MAPPINGS = {
"MultiLoraLoader": "Multi Lora Loader",
"RandomPrompt": "Random (Prompt)",
"CombinePrompt": "Combine (Prompt)",
"CSVPromptsLoader": "CSV Prompts Loader"
}
# init()
WEB_DIRECTORY = "./web"
__all__ = ["NODE_CLASS_MAPPINGS",
"NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"]
Binary file not shown.
Binary file not shown.
Binary file not shown.
+4
View File
@@ -0,0 +1,4 @@
{
"name": "CustomScripts",
"logging": true
}
+290
View File
@@ -0,0 +1,290 @@
import asyncio
import os
import json
import shutil
import inspect
import aiohttp
from server import PromptServer
from tqdm import tqdm
config = None
def is_logging_enabled():
config = get_extension_config()
if "logging" not in config:
return False
return config["logging"]
def log(message, type=None, always=False, name=None):
if not always and not is_logging_enabled():
return
if type is not None:
message = f"[{type}] {message}"
if name is None:
name = get_extension_config()["name"]
print(f"(malefish:{name}) {message}")
def get_ext_dir(subpath=None, mkdir=False):
dir = os.path.dirname(__file__)
if subpath is not None:
dir = os.path.join(dir, subpath)
dir = os.path.abspath(dir)
if mkdir and not os.path.exists(dir):
os.makedirs(dir)
return dir
def get_comfy_dir(subpath=None, mkdir=False):
dir = os.path.dirname(inspect.getfile(PromptServer))
if subpath is not None:
dir = os.path.join(dir, subpath)
dir = os.path.abspath(dir)
if mkdir and not os.path.exists(dir):
os.makedirs(dir)
return dir
def get_web_ext_dir():
config = get_extension_config()
name = config["name"]
dir = get_comfy_dir("web/extensions/malefish")
if not os.path.exists(dir):
os.makedirs(dir)
dir = os.path.join(dir, name)
return dir
def get_extension_config(reload=False):
global config
if reload == False and config is not None:
return config
config_path = get_ext_dir("malefish.json")
if not os.path.exists(config_path):
log("Missing malefish.json, this extension may not work correctly. Please reinstall the extension.",
type="ERROR", always=True, name="???")
print(f"Extension path: {get_ext_dir()}")
return {"name": "Unknown", "version": -1}
with open(config_path, "r") as f:
config = json.loads(f.read())
return config
def link_js(src, dst):
src = os.path.abspath(src)
dst = os.path.abspath(dst)
if os.name == "nt":
try:
import _winapi
_winapi.CreateJunction(src, dst)
return True
except:
pass
try:
os.symlink(src, dst)
return True
except:
import logging
logging.exception('')
return False
def is_junction(path):
if os.name != "nt":
return False
try:
return bool(os.readlink(path))
except OSError:
return False
def install_js():
src_dir = get_ext_dir("web/js")
if not os.path.exists(src_dir):
log("No JS")
return
should_install = should_install_js()
if should_install:
log("it looks like you're running an old version of ComfyUI that requires manual setup of web files, it is recommended you update your installation.", "warning", True)
dst_dir = get_web_ext_dir()
linked = os.path.islink(dst_dir) or is_junction(dst_dir)
if linked or os.path.exists(dst_dir):
if linked:
if should_install:
log("JS already linked")
else:
os.unlink(dst_dir)
log("JS unlinked, PromptServer will serve extension")
elif not should_install:
shutil.rmtree(dst_dir)
log("JS deleted, PromptServer will serve extension")
return
if not should_install:
log("JS skipped, PromptServer will serve extension")
return
if link_js(src_dir, dst_dir):
log("JS linked")
return
log("Copying JS files")
shutil.copytree(src_dir, dst_dir, dirs_exist_ok=True)
def should_install_js():
return not hasattr(PromptServer.instance, "supports") or "custom_nodes_from_web" not in PromptServer.instance.supports
def init(check_imports=None):
log("Init")
if check_imports is not None:
import importlib.util
for imp in check_imports:
spec = importlib.util.find_spec(imp)
if spec is None:
log(f"{imp} is required, please check requirements are installed.",
type="ERROR", always=True)
return False
install_js()
return True
def get_async_loop():
loop = None
try:
loop = asyncio.get_event_loop()
except:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
return loop
def get_http_session():
loop = get_async_loop()
return aiohttp.ClientSession(loop=loop)
async def download(url, stream, update_callback=None, session=None):
close_session = False
if session is None:
close_session = True
session = get_http_session()
try:
async with session.get(url) as response:
size = int(response.headers.get('content-length', 0)) or None
with tqdm(
unit='B', unit_scale=True, miniters=1, desc=url.split('/')[-1], total=size,
) as progressbar:
perc = 0
async for chunk in response.content.iter_chunked(2048):
stream.write(chunk)
progressbar.update(len(chunk))
if update_callback is not None and progressbar.total is not None and progressbar.total != 0:
last = perc
perc = round(progressbar.n / progressbar.total, 2)
if perc != last:
last = perc
await update_callback(perc)
finally:
if close_session and session is not None:
await session.close()
async def download_to_file(url, destination, update_callback=None, is_ext_subpath=True, session=None):
if is_ext_subpath:
destination = get_ext_dir(destination)
with open(destination, mode='wb') as f:
download(url, f, update_callback, session)
def wait_for_async(async_fn, loop=None):
res = []
async def run_async():
r = await async_fn()
res.append(r)
if loop is None:
try:
loop = asyncio.get_event_loop()
except:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
loop.run_until_complete(run_async())
return res[0]
def update_node_status(client_id, node, text, progress=None):
if client_id is None:
client_id = PromptServer.instance.client_id
if client_id is None:
return
PromptServer.instance.send_sync("malefish/update_status", {
"node": node,
"progress": progress,
"text": text
}, client_id)
async def update_node_status_async(client_id, node, text, progress=None):
if client_id is None:
client_id = PromptServer.instance.client_id
if client_id is None:
return
await PromptServer.instance.send("malefish/update_status", {
"node": node,
"progress": progress,
"text": text
}, client_id)
def get_config_value(key, default=None, throw=False):
split = key.split(".")
obj = get_extension_config()
for s in split:
if s in obj:
obj = obj[s]
else:
if throw:
raise KeyError("Configuration key missing: " + key)
else:
return default
return obj
def is_inside_dir(root_dir, check_path):
root_dir = os.path.abspath(root_dir)
if not os.path.isabs(check_path):
check_path = os.path.abspath(os.path.join(root_dir, check_path))
return os.path.commonpath([check_path, root_dir]) == root_dir
def get_child_dir(root_dir, child_path, throw_if_outside=True):
child_path = os.path.abspath(os.path.join(root_dir, child_path))
if is_inside_dir(root_dir, child_path):
return child_path
if throw_if_outside:
raise NotADirectoryError(
"Saving outside the target folder is not allowed.")
return None
Binary file not shown.
+149
View File
@@ -0,0 +1,149 @@
from pathlib import Path
import folder_paths
import logging
from collections.abc import Iterable
from abc import ABC
from dynamicprompts.sampling_context import SamplingContext
from dynamicprompts.enums import SamplingMethod
from dynamicprompts.generators import RandomPromptGenerator
from dynamicprompts.wildcards.wildcard_manager import WildcardManager
import re
logger = logging.getLogger(__name__)
wildcardManager = WildcardManager(
Path(folder_paths.get_folder_paths("wildcards")[0]))
class CombinePrompts(ABC):
def __init__(self):
self._current_full_prompt = None
self._current_face_prompt = None
self.lora_spec_re = re.compile("(<(?:lora|lyco):[^>]+>)")
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"styles": ("STRING", {
"multiline": True,
"default": "masterpiece"
}),
"face": ("STRING", {
"multiline": True,
"default": "1girl"
}),
"cloth": ("STRING", {
"multiline": True,
"default": "dress"
}),
"action": ("STRING", {
"multiline": True,
"default": "standing"
}),
"face_action": ("STRING", {
"multiline": True,
"default": "smile"
}),
"environment": ("STRING", {
"multiline": True,
"default": "forest"
}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff})
},
"optional": {
"added_styles": ("STRING", {"forceInput": True}),
"added_face": ("STRING", {"forceInput": True}),
"added_cloth": ("STRING", {"forceInput": True}),
"added_action": ("STRING", {"forceInput": True}),
"added_face_action": ("STRING", {"forceInput": True}),
"added_environment": ("STRING", {"forceInput": True}),
}
}
@classmethod
def IS_CHANGED(cls, text):
# Force re-evaluation of the node
return float("NaN")
RETURN_TYPES = ("STRING", "STRING", "STRING", "STRING",)
RETURN_NAMES = ("full", "face_only", "full_loras", "face_loras")
FUNCTION = "process"
# OUTPUT_NODE = False
CATEGORY = "Malefish/conditioning"
@property
def generator(self) -> RandomPromptGenerator:
return RandomPromptGenerator(
wildcard_manager=wildcardManager,
)
def process(self,
styles, face, cloth, action, face_action, environment, seed,
added_styles="", added_face="", added_cloth="", added_action="", added_face_action="", added_environment=""
):
full = ""
faceOnly = ""
if styles != "":
full += styles
faceOnly += styles
if added_styles != "":
full += ", \n" + added_styles
faceOnly += ", \n" + added_styles
if face != "":
full += ", \n" + face
faceOnly += ", \n" + face
if added_face != "":
full += ", \n" + added_face
faceOnly += ", \n" + added_face
if cloth != "":
full += ", \n" + cloth
if added_cloth != "":
full += ", \n" + added_cloth
if action != "":
full += ", \n" + action
if added_action != "":
full += ", \n" + added_action
if face_action != "":
full += ", \n" + face_action
faceOnly += ", \n" + face_action
if added_face_action != "":
full += ", \n" + added_face_action
faceOnly += ", \n" + added_face_action
if environment != "":
full += ", \n" + environment
if added_environment != "":
full += ", \n" + added_environment
if full.strip() == "" and faceOnly.strip() == "":
return ("",)\
full = self.generator.generate(full, num_images=1, seeds=seed)[0]
faceOnly = self.generator.generate(
faceOnly, num_images=1, seeds=seed)[0]
fullExtractedLoras = self.lora_spec_re.findall(full)
fullFilteredText = self.lora_spec_re.sub("", full)
faceExtractedLoras = self.lora_spec_re.findall(faceOnly)
faceFilteredText = self.lora_spec_re.sub("", faceOnly)
return (fullFilteredText, faceFilteredText, "\n".join(fullExtractedLoras), "\n".join(faceExtractedLoras))
+98
View File
@@ -0,0 +1,98 @@
from pathlib import Path
import folder_paths
import logging
from collections.abc import Iterable
from abc import ABC
from dynamicprompts.sampling_context import SamplingContext
from dynamicprompts.enums import SamplingMethod
from dynamicprompts.wildcards.wildcard_manager import WildcardManager
import re
logger = logging.getLogger(__name__)
wildcardManager = WildcardManager(
Path(folder_paths.get_folder_paths("wildcards")[0]))
class RandomPrompt(ABC):
def __init__(self):
self._current_prompt = None
self.lora_spec_re = re.compile("(<(?:lora|lyco):[^>]+>)")
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"text": ("STRING", {
"multiline": True,
"default": ""
}),
},
}
@classmethod
def IS_CHANGED(cls, text):
# Force re-evaluation of the node
return float("NaN")
RETURN_TYPES = ("STRING", "STRING",)
RETURN_NAMES = ("prompt", "loras",)
FUNCTION = "process"
# OUTPUT_NODE = False
CATEGORY = "Malefish/conditioning"
@property
def context(self) -> SamplingContext:
return SamplingContext(
wildcard_manager=wildcardManager,
default_sampling_method=SamplingMethod.RANDOM,
)
def _get_next_prompt(self, prompts: Iterable[str], current_prompt: str) -> str:
"""
Get the next prompt from the prompts generator.
"""
try:
return next(prompts)
except (StopIteration, RuntimeError):
self._prompts = self.context.sample_prompts(current_prompt)
try:
return next(prompts)
except StopIteration:
logger.exception("No more prompts to generate!")
return ""
def has_prompt_changed(self, text: str) -> bool:
"""
Check if the prompt has changed.
"""
return self._current_prompt != text
def process(self, text):
if text.strip() == "":
return ("",)
if self.has_prompt_changed(text):
self._current_prompt = text
self._prompts = self.context.sample_prompts(self._current_prompt)
if self._prompts is None:
logger.exception("Something went wrong. Prompts is None!")
return ("",)
if self._current_prompt is None:
logger.exception("Something went wrong. Current prompt is None!")
return ("",)
new_prompt = self._get_next_prompt(self._prompts, self._current_prompt)
extracted_loras = self.lora_spec_re.findall(new_prompt)
filtered_text = self.lora_spec_re.sub("", new_prompt)
print(f"New prompt: {new_prompt}")
return (filtered_text, "\n".join(extracted_loras),)
Binary file not shown.
Binary file not shown.
+71
View File
@@ -0,0 +1,71 @@
import logging
import os
import pathlib
import pandas as pd
logger = logging.getLogger(__name__)
class CSVPromptsLoader():
def __init__(self):
pass
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"path": ("STRING", {
"multiline": False,
"default": ""
}),
"key": ("STRING", {
"multiline": False,
"default": ""
}),
"remove_extension": ([
"On",
"Off"],)
},
}
RETURN_TYPES = ("STRING", "STRING", "STRING", "STRING", "STRING", )
RETURN_NAMES = ("styles", "face", "cloth", "action", "face_action",)
FUNCTION = "process"
# OUTPUT_NODE = False
CATEGORY = "Malefish/io"
@classmethod
def IS_CHANGED(cls, text):
# Force re-evaluation of the node
return float("NaN")
def process(self, path, key, remove_extension):
if path.strip() == "":
return ("", "", "", "", "",)
if not os.path.exists(path):
logger.exception("Path is not exist")
return ("", "", "", "", "",)
file_extension = pathlib.Path(path).suffix
if file_extension != ".csv":
logger.exception("file is not csv")
return ("", "", "", "", "",)
df = pd.read_csv(path, index_col=0, dtype={
"name": 'string', "styles": 'string', "face": 'string', "cloth": 'string', "action": 'string', "face_action": 'string'})
if remove_extension == "On":
key = os.path.splitext(key)[0]
try:
values = df.loc[key]
except KeyError:
logger.exception("Key is not exist")
return ("", "", "", "", "",)
return (values['styles'], values['face'], values['cloth'], values['action'], values['face_action'],)
+162
View File
@@ -0,0 +1,162 @@
import torch
import folder_paths
import comfy.utils
import comfy.sd
import os
import re
class MultiLoraLoader:
def __init__(self):
self.lora_items = []
@classmethod
def INPUT_TYPES(s):
return {"required": {"model": ("MODEL",),
"clip": ("CLIP", ),
"text": ("STRING", {
"multiline": True,
"default": ""}),
}}
RETURN_TYPES = ("MODEL", "CLIP")
FUNCTION = "load_loras"
CATEGORY = "Malefish/io"
def load_loras(self, model, clip, text):
result = (model, clip)
available_loras = self.available_loras()
self.update_current_lora_items_with_new_items(
self.items_from_lora_text_with_available_loras(text, available_loras))
if len(self.lora_items) > 0:
for item in self.lora_items:
if item.lora_name in available_loras:
result = item.apply_lora(result[0], result[1])
else:
raise ValueError(
f"Unable to find lora with name '{item.lora_name}'")
return result
def available_loras(self):
return folder_paths.get_filename_list("loras")
def items_from_lora_text_with_available_loras(self, lora_text, available_loras):
return LoraItemsParser.parse_lora_items_from_text(lora_text, self.dictionary_with_short_names_for_loras(available_loras))
def dictionary_with_short_names_for_loras(self, available_loras):
result = {}
for path in available_loras:
result[os.path.splitext(os.path.basename(path))[0]] = path
return result
def update_current_lora_items_with_new_items(self, lora_items):
if self.lora_items != lora_items:
existing_by_name = dict(
[(existing_item.lora_name, existing_item) for existing_item in self.lora_items])
for new_item in lora_items:
new_item.move_resources_from(existing_by_name)
self.lora_items = lora_items
class LoraItemsParser:
@classmethod
def parse_lora_items_from_text(cls, lora_text, loras_by_short_names={}, default_weight=1, weight_separator=":"):
return cls(lora_text, loras_by_short_names, default_weight, weight_separator).execute()
def __init__(self, lora_text, loras_by_short_names, default_weight, weight_separator):
self.lora_text = lora_text
self.loras_by_short_names = loras_by_short_names
self.default_weight = default_weight
self.weight_separator = weight_separator
self.prefix_trim_re = re.compile("\A<(lora|lyco):")
self.comment_trim_re = re.compile("\s*#.*\Z")
def execute(self):
return [LoraItem(elements[0], elements[1], elements[2])
for line in self.lora_text.splitlines()
for elements in [self.parse_lora_description(self.description_from_line(line))] if elements[0] is not None]
def parse_lora_description(self, description):
if description is None:
return (None,)
lora_name = None
strength_model = self.default_weight
strength_clip = None
remaining, sep, strength = description.rpartition(
self.weight_separator)
if sep == self.weight_separator:
lora_name = remaining
strength_model = float(strength)
remaining, sep, strength = remaining.rpartition(
self.weight_separator)
if sep == self.weight_separator:
strength_clip = strength_model
strength_model = float(strength)
lora_name = remaining
else:
lora_name = description
if strength_clip is None:
strength_clip = strength_model
return (self.loras_by_short_names.get(lora_name, lora_name), strength_model, strength_clip)
def description_from_line(self, line):
result = self.comment_trim_re.sub("", line.strip())
result = self.prefix_trim_re.sub("", result.removesuffix(">"))
return result if len(result) > 0 else None
class LoraItem:
def __init__(self, lora_name, strength_model, strength_clip):
self.lora_name = lora_name
self.strength_model = strength_model
self.strength_clip = strength_clip
self._loaded_lora = None
def __eq__(self, other):
return self.lora_name == other.lora_name and self.strength_model == other.strength_model and self.strength_clip == other.strength_clip
def get_lora_path(self):
return folder_paths.get_full_path("loras", self.lora_name)
def move_resources_from(self, lora_items_by_name):
existing = lora_items_by_name.get(self.lora_name)
if existing is not None:
self._loaded_lora = existing._loaded_lora
existing._loaded_lora = None
def apply_lora(self, model, clip):
if self.is_noop:
return (model, clip)
model_lora, clip_lora = comfy.sd.load_lora_for_models(
model, clip, self.lora_object, self.strength_model, self.strength_clip)
return (model_lora, clip_lora)
@property
def lora_object(self):
if self._loaded_lora is None:
lora_path = self.get_lora_path()
if lora_path is None:
raise ValueError(
f"Unable to get file path for lora with name '{self.lora_name}'")
self._loaded_lora = comfy.utils.load_torch_file(
lora_path, safe_load=True)
return self._loaded_lora
@property
def is_noop(self):
return self.strength_model == 0 and self.strength_clip == 0
+2
View File
@@ -0,0 +1,2 @@
dynamicprompts
pandas
+55
View File
@@ -0,0 +1,55 @@
import { app } from "../scripts/app.js";
const ext = {
// Unique name for the extension
name: "Example.LoggingExtension",
async init(app) {
// Any initial setup to run as soon as the page loads
console.log("[logging]", "extension init");
},
async setup(app) {
// Any setup to run after the app is created
console.log("[logging]", "extension setup");
},
async addCustomNodeDefs(defs, app) {
// Add custom node definitions
// These definitions will be configured and registered automatically
// defs is a lookup core nodes, add yours into this
console.log("[logging]", "add custom node definitions", "current nodes:", Object.keys(defs));
},
async getCustomWidgets(app) {
// Return custom widget types
// See ComfyWidgets for widget examples
console.log("[logging]", "provide custom widgets");
},
async beforeRegisterNodeDef(nodeType, nodeData, app) {
// Run custom logic before a node definition is registered with the graph
console.log("[logging]", "before register node: ", nodeType, nodeData);
// This fires for every node definition so only log once
delete ext.beforeRegisterNodeDef;
},
async registerCustomNodes(app) {
// Register any custom node implementations here allowing for more flexability than a custom node def
console.log("[logging]", "register custom nodes");
},
loadedGraphNode(node, app) {
// Fires for each node when loading/dragging/etc a workflow json or png
// If you break something in the backend and want to patch workflows in the frontend
// This is the place to do this
console.log("[logging]", "loaded graph node: ", node);
// This fires for every node on each load so only log once
delete ext.loadedGraphNode;
},
nodeCreated(node, app) {
// Fires every time a node is constructed
// You can modify widgets/add handlers/etc here
console.log("[logging]", "node created: ", node);
// This fires for every node so only log once
delete ext.nodeCreated;
}
};
app.registerExtension(ext);