first commit
This commit is contained in:
+26
@@ -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.
@@ -0,0 +1,4 @@
|
||||
{
|
||||
"name": "CustomScripts",
|
||||
"logging": true
|
||||
}
|
||||
+290
@@ -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.
Binary file not shown.
Binary file not shown.
@@ -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))
|
||||
@@ -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.
@@ -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'],)
|
||||
@@ -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
|
||||
@@ -0,0 +1,2 @@
|
||||
dynamicprompts
|
||||
pandas
|
||||
@@ -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);
|
||||
Reference in New Issue
Block a user