From 2aa1a855e8a6b69e07c9652844fb1f13e690c175 Mon Sep 17 00:00:00 2001 From: rgthree Date: Thu, 22 May 2025 22:49:38 -0400 Subject: [PATCH] Power Puter: Add condition support, passing arguments to special functions, and get trigger words for enabled loras, get connected nodes. --- py/power_lora_loader.py | 34 +++++++++++++- py/power_puter.py | 99 +++++++++++++++++++++++++++++++---------- pyproject.toml | 2 +- 3 files changed, 110 insertions(+), 25 deletions(-) diff --git a/py/power_lora_loader.py b/py/power_lora_loader.py index 9a03f28..6a22d91 100644 --- a/py/power_lora_loader.py +++ b/py/power_lora_loader.py @@ -1,13 +1,19 @@ +import asyncio + from nodes import LoraLoader from .constants import get_category, get_name from .power_prompt_utils import get_lora_by_filename from .utils import FlexibleOptionalInputType, any_type +from .server.utils_info import get_model_info +from .log import log_node_warn + +NODE_NAME = get_name('Power Lora Loader') class RgthreePowerLoraLoader: """ The Power Lora Loader is a powerful, flexible node to add multiple loras to a model/clip.""" - NAME = get_name('Power Lora Loader') + NAME = NODE_NAME CATEGORY = get_category() @classmethod @@ -42,3 +48,29 @@ class RgthreePowerLoraLoader: model, clip = LoraLoader().load_lora(model, clip, lora, strength_model, strength_clip) return (model, clip) + + @classmethod + def get_enabled_loras_from_prompt_node(cls, prompt_node: dict) -> list[dict[str, str | float]]: + return [{ + 'name': lora['lora'], + 'strength': lora['strength'] + } | ({ + 'strength_clip': lora['strengthTwo'] + } if 'strengthTwo' in lora else {}) + for name, lora in prompt_node['inputs'].items() + if name.startswith('lora_') and lora['on']] + + @classmethod + def get_enabled_triggers_from_prompt_node(cls, prompt_node: dict, max_each: int = 1): + loras = [l['name'] for l in cls.get_enabled_loras_from_prompt_node(prompt_node)] + trained_words = [] + for lora in loras: + info = asyncio.run(get_model_info(lora, 'loras')) + if not info or not info.keys(): + log_node_warn(NODE_NAME, f'No info found for lora {lora} when grabbging triggers.') + continue + if 'trainedWords' not in info or not info['trainedWords']: + log_node_warn(NODE_NAME, f'No trained words for lora {lora} when grabbging triggers.') + continue + trained_words += [w for wi in info['trainedWords'][:max_each] if (wi and (w := wi['word']))] + return trained_words diff --git a/py/power_puter.py b/py/power_puter.py index 6d578d8..bd4eeb6 100644 --- a/py/power_puter.py +++ b/py/power_puter.py @@ -19,6 +19,8 @@ from .constants import get_category, get_name from .utils import FlexibleOptionalInputType, any_type, get_dict_value from .log import log_node_warn +from .power_lora_loader import RgthreePowerLoraLoader + @dataclasses.dataclass(frozen=True, kw_only=True) class Function(): @@ -52,8 +54,9 @@ _FUNCTIONS = { Function(name="str", call=str, args=(1, 1)), Function(name="bool", call=bool, args=(1, 1)), # Special - Function(name="node", call='_get_node', args=(1, 1)), + Function(name="node", call='_get_node', args=(0, 1)), Function(name="nodes", call='_get_nodes', args=(0, 1)), + Function(name="input_node", call='_get_input_node', args=(0, 1)), Function(name="dir", call=dir, args=(1, 1)), Function(name="type", call=type, args=(1, 1)), ] @@ -66,17 +69,10 @@ _GLOBAL_FUNCTION_NAMES = ['node', 'nodes'] # Special functions by class type (called from the Attrs.) _SPECIAL_FUNCTIONS = { - "Power Lora Loader (rgthree)": { + RgthreePowerLoraLoader.NAME: { # Get a list of the enabled loras from a power lora loader. - "loras": - lambda node: [{ - 'name': lora['lora'], - 'strength': lora['strength'] - } | ({ - 'strength_clip': lora['strengthTwo'] - } if 'strengthTwo' in lora else {}) - for name, lora in node['inputs'].items() - if name.startswith('lora_') and lora['on']] + "loras": RgthreePowerLoraLoader.get_enabled_loras_from_prompt_node, + "triggers": RgthreePowerLoraLoader.get_enabled_triggers_from_prompt_node, } } @@ -117,8 +113,9 @@ class RgthreePowerPuter: "required": {}, "optional": FlexibleOptionalInputType(any_type), "hidden": { + "unique_id": "UNIQUE_ID", "extra_pnginfo": "EXTRA_PNGINFO", - "prompt": "PROMPT" + "prompt": "PROMPT", }, } @@ -140,6 +137,7 @@ class RgthreePowerPuter: """Does the nodes' work.""" output = kwargs['output'] code = kwargs['code'] + unique_id = kwargs['unique_id'] pnginfo = kwargs['extra_pnginfo'] workflow = pnginfo["workflow"] if "workflow" in pnginfo else {"nodes": []} prompt = kwargs['prompt'] @@ -150,7 +148,7 @@ class RgthreePowerPuter: del ctx['extra_pnginfo'] del ctx['prompt'] - eva = _Puter(code=code, ctx=ctx, workflow=workflow, prompt=prompt) + eva = _Puter(code=code, ctx=ctx, workflow=workflow, prompt=prompt, unique_id=unique_id) value = eva.execute() if value is not None: @@ -174,7 +172,7 @@ class _Puter: See https://www.basicexamples.com/example/python/ast for examples. """ - def __init__(self, *, code: str, ctx: dict[str, Any], workflow, prompt): + def __init__(self, *, code: str, ctx: dict[str, Any], workflow, prompt, unique_id): ctx = ctx or {} self._ctx = {**ctx} self._code = code @@ -183,14 +181,16 @@ class _Puter: self._prompt_nodes = [] if self._prompt: self._prompt_nodes = [{'id': id} | {**node} for id, node in self._prompt.items()] + self._prompt_node = [n for n in self._prompt_nodes if n['id'] == unique_id][0] def execute(self, code=str | None) -> Any: """Evaluates a the code block.""" code = code or self._code node = ast.parse(self._code) + ctx = {**self._ctx} last_value = None for body in node.body: - last_value = self._eval_statement(body) + last_value = self._eval_statement(body, ctx) # If we got a return, then that's it folks. if isinstance(body, ast.Return): break @@ -218,16 +218,29 @@ class _Puter: found = [n for n in nodes if node_id == get_dict_value(n, '_meta.title', '')] return found - def _get_node(self, node_id: int | str | re.Pattern) -> Any | None: + def _get_node(self, node_id: int | str | re.Pattern | None = None) -> Any | None: """Returns a prompt-node from the hidden prompt.""" + if node_id is None: + return self._prompt_node nodes = self._get_nodes(node_id) if nodes and len(nodes) > 1: log_node_warn(_NODE_NAME, f"More than one node found for '{node_id}'. Returning first.") return nodes[0] if nodes else None - def _eval_statement(self, stmt: ast.stmt, ctx: dict | None = None): + def _get_input_node(self, input_name, node=None): + """Gets the (non-muted) node of an input connection from a node (default to the power puter).""" + node = node if node else self._prompt_node + try: + connected_node_id = node['inputs'][input_name][0] + return [n for n in self._prompt_nodes if n['id'] == connected_node_id][0] + except (TypeError, IndexError, KeyError): + log_node_warn(_NODE_NAME, f'No input node found for "{input_name}". ') + + def _eval_statement(self, stmt: ast.stmt, ctx: dict, prev_stmt: ast.stmt | None = None): """Evaluates an ast.stmt.""" - ctx = self._ctx if ctx is None else ctx + + if '__returned__' in ctx: + return ctx['__returned__'] # print('\n\n----: _eval_statement') # print(type(stmt)) @@ -253,7 +266,7 @@ class _Puter: return _OPERATORS[type(stmt.op)](left, right) if isinstance(stmt, ast.UnaryOp): - return _OPERATORS[type(stmt.op)](self._eval_statement(stmt.operand), ctx=ctx) + return _OPERATORS[type(stmt.op)](self._eval_statement(stmt.operand, ctx=ctx)) if isinstance(stmt, (ast.Attribute, ast.Subscript)): # Like: node(14).inputs.sampler_name (Attribute) @@ -271,13 +284,25 @@ class _Puter: # Any special cases in the _SPECIAL_FUNCTIONS class_type = get_dict_value(item, "class_type") if class_type in _SPECIAL_FUNCTIONS and attr in _SPECIAL_FUNCTIONS[class_type]: - val = _SPECIAL_FUNCTIONS[class_type][attr](item) + val = _SPECIAL_FUNCTIONS[class_type][attr] + # If our previous statment was a Call, then send back a tuple of the callable and + # the evaluated item, and it will make the call; perhaps also adding other arguments + # only it knows about. + if isinstance(prev_stmt, ast.Call): + return (val, item) + val = val(item) else: val = None else: raise return val + if isinstance(stmt, ast.List): + value = [] + for elt in stmt.elts: + value.append(self._eval_statement(elt, ctx=ctx)) + return value + # f-strings: https://www.basicexamples.com/example/python/ast-JoinedStr # Note, this will str() all evaluated items in the fstrings, and doesn't handle f-string # directives, like padding, etc. @@ -347,8 +372,13 @@ class _Puter: if isinstance(stmt, ast.Call): call = None + args = [] + kwargs = {} if isinstance(stmt.func, ast.Attribute): - call = self._eval_statement(stmt.func, ctx=ctx) + call = self._eval_statement(stmt.func, prev_stmt=stmt, ctx=ctx) + if isinstance(call, tuple): + args.append(call[1]) + call = call[0] if isinstance(stmt.func, ast.Name): name = stmt.func.id if name in _FUNCTIONS: @@ -364,10 +394,11 @@ class _Puter: # print(type(stmt)) # print(type(stmt.func)) raise ValueError(f'No call for ast.Call {stmt}') - args = [] for arg in stmt.args: args.append(self._eval_statement(arg, ctx=ctx)) - return call(*args) + for kwarg in stmt.keywords: + kwargs[kwarg.arg] = self._eval_statement(kwarg.value, ctx=ctx) + return call(*args, **kwargs) if isinstance(stmt, ast.Compare): l = self._eval_statement(stmt.left, ctx=ctx) @@ -388,6 +419,21 @@ class _Puter: return 1 if l in r else 0 raise NotImplementedError("Operator " + stmt.ops[0].__class__.__name__ + " not supported.") + if isinstance(stmt, (ast.If, ast.IfExp)): + value = self._eval_statement(stmt.test, ctx=ctx) + if value: + # ast.If is a list, ast.IfExp is an object. + bodies = stmt.body if isinstance(stmt.body, list) else [stmt.body] + for body in bodies: + value = self._eval_statement(body, ctx=ctx) + elif stmt.orelse: + # ast.If is a list, ast.IfExp is an object. TBH, I don't know why the If is a list, it's + # only ever one item AFAICT. + orelses = stmt.orelse if isinstance(stmt.orelse, list) else [stmt.orelse] + for orelse in orelses: + value = self._eval_statement(orelse, ctx=ctx) + return value + # Assign a variable and add it to our ctx. if isinstance(stmt, ast.Assign): value = self._eval_statement(stmt.value, ctx=ctx) @@ -401,4 +447,11 @@ class _Puter: ctx[stmt.target.id] = value return value + if isinstance(stmt, ast.Return): + value = self._eval_statement(stmt.value, ctx=ctx) + # Mark that we have a return value, as we may be deeper in evaluation, like going through an + # if condition's body. + ctx['__returned__'] = value + return value + raise TypeError(stmt) diff --git a/pyproject.toml b/pyproject.toml index 09c4b03..1246271 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "rgthree-comfy" description = "Making ComfyUI more comfortable." -version = "1.0.2505212136" +version = "1.0.2505222249" license = { file = "LICENSE" } dependencies = []