52 changed files with 1075 additions and 3799 deletions
-10
View File
@@ -1,10 +0,0 @@
name: Test
on:
workflow_dispatch:
jobs:
build:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v2
- name: Setup upterm session
uses: lhotari/action-upterm@v1
-85
View File
@@ -1,85 +0,0 @@
name: Run Test Workflow
on: [workflow_dispatch]
jobs:
Test:
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
python-version: [ "3.7", "3.8", "3.9", "3.10", "3.11" ]
steps:
- name: Clone Upstream
uses: actions/checkout@v3
with:
repository: comfyanonymous/ComfyUI
ref: master
fetch-depth: 0
- name: Clone Node
uses: actions/checkout@v3
with:
ref: master
fetch-depth: 0
path: custom_nodes/KepPromptLang
- name: Setup Python
uses: actions/setup-python@v4
with:
# Version range or exact version of Python or PyPy to use, using SemVer's version range syntax. Reads from .python-version if unset.
python-version: ${{ matrix.python-version }}
- name: Cache virtualenv
uses: actions/cache@v3
id: cache-venv
with:
path: ./.venv/
key: ${{ runner.os }}-venv-${{ matrix.python-version }}-${{ hashFiles('**/requirements.txt') }}
restore-keys: |
${{ runner.os }}-venv-${{ matrix.python-version }}-
- name: Install Requirements
if: steps.cache-venv.outputs.cache-hit != 'true'
run: |
python -m venv ./.venv
source ./.venv/bin/activate
pip install torch --index-url https://download.pytorch.org/whl/cpu
pip install -r requirements.txt
pip install -r custom_nodes/KepPromptLang/requirements.txt
pip install huggingface_hub websocket-client
# - name: Cache SD Checkpoint
# uses: actions/cache@v3
# with:
# path: |
# models/checkpoints
# key: ${{ runner.os }}-sd-15-checkpoint
- name: Check and Download Model
run: |
source ./.venv/bin/activate
python custom_nodes/KepPromptLang/test_files/check_and_download_model.py
- name: Run in Background
env:
PYTHONUNBUFFERED: 1
run: |
source ./.venv/bin/activate
python main.py --cpu &> server.log &
sleep 10
# - name: Setup upterm session
# uses: lhotari/action-upterm@v1
- name: Run Workflow
run: |
source ./.venv/bin/activate
python custom_nodes/KepPromptLang/test_files/run_workflow.py
- name: Upload Comfy Server Log
if: always()
uses: actions/upload-artifact@v3
with:
name: comfy-server-log-${{ matrix.python-version }}
path: server.log
+1 -100
View File
@@ -1,100 +1 @@
# KepPromptLang
A small DSL for ComfyUI that lets you do math on CLIP token embeddings before they're fed into the text transformer.
```
sum(diff(king|man)|woman)
norm(sum(cat | dog | horse | parrot))
A slerp(cat|dog|0.5) is happy
```
## Install
Clone into `ComfyUI/custom_nodes/`:
```bash
cd ComfyUI/custom_nodes
git clone <repo-url> KepPromptLang
pip install -r KepPromptLang/requirements.txt
```
## Usage
1. Add a **Special CLIP Loader** node and feed it the CLIP output from your **Load Checkpoint**.
2. Pass the wrapped CLIP into a standard **CLIP Text Encode** node.
3. Use the DSL syntax in your prompt.
To debug what your DSL is doing, add a **PromptLang Inspect** node — it shows the per-slot weight, L2 norm, and nearest-vocab words for the resolved embeddings.
See `examples/WIP_Example_workflow.json` for a working workflow.
![Example](assets/first_example.png)
## Syntax
| Element | Syntax | Example |
| --- | --- | --- |
| Plain word | alphanumeric (with `,_.-`) | `cat`, `dog_face` |
| Quoted string | single or double quotes | `"hello world"`, `'it\'s sunny'` |
| Weighted | `(text:weight)` or `emph(text\|weight)` | `(cat:1.3)`, `emph(cat\|1.3)` |
| Embedding (textual inversion) | `embedding:NAME` | `embedding:face_vector` |
| Function | `name(arg \| arg \| ...)` | `sum(king \| woman)` |
Arguments inside a function are separated by `|`. Each arg can itself be plain text, an embedding, a quoted string, or another function call.
### Variables and comments
```
$axis = diff(king|queen); # name an expression
sum(actor|$axis) and reject(doctor|$axis)
```
`$NAME = arg;` binds a name; `$NAME` substitutes it. Single-pass: define before use, no reassignment. `#` comments run to end of line. Substitution is structural — multiple refs share the same parsed action object, but actions are evaluated per occurrence (so `$r = rand(3); $r $r` re-rolls each use).
## Quick examples
- Average two prompts: `avg(The cat is | The dog is | 0.5)`
- Normalize a sum: `norm(sum(cat | dog | horse))`
- King − Man + Woman: `sum(diff(king|man)|woman)` (or `sum(king | neg(man) | woman)`)
- Negate an embedding: `neg(embedding:body_vector)`
## Functions
| Display Name | Action Name | Description | Usage Examples |
| --- | --- | --- | --- |
| Average | avg | Performs a weighted average between two segments or actions. The recommended weight is 0 - 1. | <ul><li>avg(The cat is\|The dog is\|0.5)</li><li>avg(Cat\|Dog\|0.5)</li></ul> |
| Difference | diff | Subtracts the segments in the order they are given. The first segment is subtracted from the second, then the third from the result, and so on. | <ul><li>diff(The cat is\|The dog is)</li><li>diff(Cat\|Dog)</li><li>sum(diff(king\|man)\|woman)</li></ul> |
| Multiply | mult | Multiplies the provided segments or actions by the multiplier. | <ul><li>mult(The cat is\|2.5)</li><li>mult(Cat\|-1)</li></ul> |
| Nearest Vocab | nearest | Snaps a computed vector to the k nearest real vocabulary tokens (by cosine similarity), returning their embeddings concatenated. The input is mean-pooled before lookup. | <ul><li>nearest(sum(diff(king\|man)\|woman))</li><li>nearest(sum(red\|blue)\|3)</li></ul> |
| Negate | neg | Negates the provided segments or actions. | <ul><li>neg(cat)</li><li>sum(king\|neg(man)\|women)</li></ul> |
| Noise | noise | Adds Gaussian noise (mean 0, given std) to the embeddings of the first argument. | <ul><li>A noise(cat\|0.05) on a sunny day</li></ul> |
| Normalize | norm | Normalizes the provided segments or actions. | <ul><li>norm(cat)</li><li>sum(cat\|norm(sum(tiger\|fish)))</li></ul> |
| Positional Embedding Scale | posScale | Scales (multiplies) the positional embeddings of the provided segments or actions by the multiplier. | <ul><li>A posScale(cat\|1.5) on a rainy day</li></ul> |
| Ignore Positional Embeddings | postPos | Prevents positional embeddings from being applied to the provided segments or actions. | <ul><li>A postPos(cat) on a rainy day</li></ul> |
| Project | proj | Projects the first argument onto the direction of the second (mean, unit-normalized). | <ul><li>proj(king\|gender)</li><li>diff(style\|proj(style\|photorealistic))</li></ul> |
| Random Embedding | rand | Returns a random embedding of the specified token length, with the values optionally bounded by the second and third arguments. | <ul><li>A rand(1) cat</li><li>A rand(1\|-1\|1) cat</li></ul> |
| Reject | reject | Removes the component of the first argument along the direction of the second (a - proj(a\|b)). | <ul><li>reject(anime girl\|anime)</li></ul> |
| Renormalize | renorm | Rescales the first argument so each token's L2 norm matches the (mean) L2 norm of the reference. | <ul><li>renorm(sum(king\|neg(man)\|woman)\|queen)</li></ul> |
| Scale Dimensions | scaleDims | Scales the specified dimensions of the input embeddings by the specified amount | <ul><li>The scaleDims(cat\|4,1.5\|76,1.2) is happy</li></ul> |
| Set Dimensions | setDims | Sets the specified dimensions of the input embeddings to the specified value | <ul><li>The setDims(cat\|4, -0.01253\|76, 1.2) is happy</li></ul> |
| Slerp | slerp | Performs a slerp (interpolation) between two segments or actions, with the given weight. The recommended weight is 0 - 1. | <ul><li>The slerp(cat\|dog\|0.5) is happy</li></ul> |
| Sum | sum | Adds the embeddings of the provided segments or actions. | <ul><li>A happy sum(cat\|dog\|shark)</li></ul> |
`lerp(a|b|t)` is also accepted as an alias for `avg(a|b|t)`.
Regenerate the table with `python tools/build_docs.py`.
## Development
Tests are pytest-based and don't require ComfyUI:
```bash
pip install -e ".[dev]"
python -m pytest
```
## Compatibility
- SD1.x (CLIP-L) and SDXL (CLIP-L + CLIP-G).
- SD2 is not supported.
- Two pooler-output actions (`_exp-pooler`, `_exp-pooledAvg`) from earlier versions were experimental and have been removed; they relied on direct HuggingFace transformer access that is no longer how ComfyUI structures its CLIP encoders.
# ClipStuff
+6 -10
View File
@@ -1,15 +1,11 @@
from .nodes import BuildGif, PromptLangInspect, SpecialClipLoader
from .nodes import (
KepAdvTextEncode,
BuildGif,
SpecialClipLoader,
)
NODE_CLASS_MAPPINGS = {
"Kep Adv Text Encode": KepAdvTextEncode,
"Build Gif": BuildGif,
"Special CLIP Loader": SpecialClipLoader,
"PromptLang Inspect": PromptLangInspect,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"Build Gif": "Build GIF (KepPromptLang)",
"Special CLIP Loader": "Special CLIP Loader (KepPromptLang)",
"PromptLang Inspect": "PromptLang Inspect",
}
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.8 MiB

File diff suppressed because it is too large Load Diff
+5 -48
View File
@@ -1,49 +1,6 @@
from ..parser.registration import register_action
from .avg import AverageAction
from .diff import DiffAction
from .mult import MultiplyAction
from .nearest import NearestAction
from .neg import NegAction
from .noise import NoiseAction
from .norm import NormAction
from .pos_scale import PosScaleAction
from .post_pos import PostPosAction
from .project import ProjectAction, RejectAction
from .rand import RandAction
from .renorm import RenormAction
from .scale_dims import ScaleDims
from .set_dims import SetDims
from .slerp import SlerpAction
from .sum import SumAction
from custom_nodes.ClipStuff.lib.actions.arith import ArithAction
from custom_nodes.ClipStuff.lib.actions.nudge import NudgeAction
for _action in [
AverageAction,
DiffAction,
MultiplyAction,
NearestAction,
NegAction,
NoiseAction,
NormAction,
PosScaleAction,
PostPosAction,
ProjectAction,
RandAction,
RejectAction,
RenormAction,
ScaleDims,
SetDims,
SlerpAction,
SumAction,
]:
register_action(_action)
class _LerpAlias(AverageAction):
"""`lerp(a|b|t)` is sugar for `avg(a|b|t)`."""
display_name = "Lerp"
action_name = "lerp"
usage_examples = ["lerp(cat|dog|0.5)"]
register_action(_LerpAlias)
ALL_ACTIONS = [NudgeAction, ArithAction]
ALL_START_CHARS = [action.START_CHAR for action in ALL_ACTIONS]
ALL_END_CHARS = [action.END_CHAR for action in ALL_ACTIONS]
-63
View File
@@ -1,63 +0,0 @@
from typing import Callable, List, TypeVar
import torch
from torch import Tensor
from torch.nn import Embedding
from .base import Action
from .types import SegOrAction
from .weighted import WeightedGroup
def embedding_tensor(seg_or_action: SegOrAction, embedding_module: Embedding) -> Tensor:
"""Embeddings for a segment, or get_result() for an action — always a bare tensor."""
if isinstance(seg_or_action, Action):
result = seg_or_action.get_result(embedding_module)
return result[0] if isinstance(result, tuple) else result
if isinstance(seg_or_action, WeightedGroup):
# Weights apply post-transformer, not to embedding math; recurse and drop the weight.
return concat_embeddings(seg_or_action.items, embedding_module)
return seg_or_action.get_embeddings(embedding_module)
def get_total_length(args: List[SegOrAction]) -> int:
return sum(seg_or_action.token_length() for seg_or_action in args)
def concat_embeddings(args: List[SegOrAction], embedding_module: Embedding) -> Tensor:
"""Materialize and concatenate embeddings for a sequence of segments/actions along the seq dim."""
return torch.cat([embedding_tensor(x, embedding_module) for x in args], dim=1)
def add_with_broadcast(result: Tensor, arg_embedding: Tensor, op: str) -> Tensor:
"""Add or subtract arg_embedding into result, averaging arg over the seq dim if shapes mismatch."""
matched = arg_embedding.shape[-2] == 1 or result.shape[-2] == arg_embedding.shape[-2]
if not matched:
print(f"WARNING: shape mismatch when trying to apply {op}, arg will be averaged")
arg_embedding = torch.mean(arg_embedding, dim=1, keepdim=True)
return result.add(arg_embedding) if op == "add" else result.sub(arg_embedding)
T = TypeVar("T")
def parse_numeric_arg(
arg: List[SegOrAction],
*,
action_name: str,
role: str,
cast: Callable[[str], T] = float,
) -> T:
"""Pull a single numeric value out of a one-segment arg, with helpful errors.
Used by every action that takes a scalar weight/multiplier/length.
"""
if len(arg) != 1:
raise ValueError(f"{action_name} {role} should have exactly one segment")
item = arg[0]
if isinstance(item, (Action, WeightedGroup)):
raise ValueError(f"{action_name} {role} must be a plain numeric segment")
try:
return cast(item.text)
except ValueError:
raise ValueError(f"{action_name} {role} should be a {cast.__name__}")
+131
View File
@@ -0,0 +1,131 @@
from typing import Callable, Union
import torch
from torch.nn import Embedding
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment
class ArithAction(Action):
START_CHAR = "<"
END_CHAR = ">"
def __init__(self, base_segment: PromptSegment | Action, ops: dict[str, list[PromptSegment | Action]]):
self.base_segment = base_segment
self.ops = ops
def __repr__(self):
return f"ArithAction(\n\tbase_segment={self.base_segment},\n\tops={self.ops}\n)"
def depth_repr(self, depth=1):
out = "ArithAction(\n"
if isinstance(self.base_segment, Action):
base_segment_repr = self.base_segment.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
elif isinstance(self.base_segment, PromptSegment):
out += "\t" * depth + f'base_segment={self.base_segment.depth_repr(depth)}'
else:
out += "\t" * depth + f'base_segment="{self.base_segment}",'
for op_key, ops in self.ops.items():
for op in ops:
out += "\n" + "\t" * depth + f'"{op_key}":[\n'
if isinstance(op, Action):
op_repr = op.depth_repr(depth + 2)
out += "\t" * (depth + 1) + f"{op_repr}\n"
else:
out += "\t" * (depth + 1) + f'{op.depth_repr()},\n'
out += "\t" * depth + "],"
out += "\n" + "\t" * (depth - 1) + ")"
return out
def token_length(self):
# ArithAction modifies the embeddings of the base segment, so the length is the length of the base segment
if isinstance(self.base_segment, Action):
return self.base_segment.token_length()
return len(self.base_segment.tokens)
def get_all_segments(self):
segments = []
if isinstance(self.base_segment, Action):
segments += self.base_segment.get_all_segments()
else:
segments.append(self.base_segment)
for op_key, ops in self.ops.items():
for op in ops:
if isinstance(op, Action):
segments += op.get_all_segments()
else:
segments.append(op)
return segments
def get_result(self, embedding_module: Embedding):
if isinstance(self.base_segment, Action):
base_segment_result = self.base_segment.get_result(embedding_module)
else:
base_segment_result = self.base_segment.get_embeddings(embedding_module)
for op_key, ops in self.ops.items():
for op in ops:
if isinstance(op, Action):
op_result = op.get_result(embedding_module)
else:
op_result = op.get_embeddings(embedding_module)
if op_result.shape[1] > base_segment_result.shape[1]:
print('[WARN] ArithAction: op_result.shape[1] > base_segment_result.shape[1] - averaging op_result')
op_result = torch.mean(op_result, dim=1, keepdim=True)
if op_key == "+":
base_segment_result.add(op_result)
elif op_key == "-":
base_segment_result.subtract(op_result)
return base_segment_result
@classmethod
def parse_segment(
cls,
tokens: list[str],
start_chars: list[str],
end_chars: list[str],
parent_parser: Callable[[list[str], SD1Tokenizer], Union[PromptSegment, 'Action']],
tokenizer: SD1Tokenizer,
) -> Action:
"""
Parse an arithmetic action from a list of tokens
Supported formats:
<base_segment:+op1-op2-op3>
:param tokens: List of tokens, will be modified
:param start_chars: List of start chars for all actions
:param end_chars: List of end chars for all actions
:param parent_parser: Function to parse segments to allow for nested actions
:return:
"""
token = tokens.pop(0)
assert token == cls.START_CHAR, "ArithAction must start with " + cls.START_CHAR + " but got " + token
# Parse base segment
base_segment = parent_parser(tokens, tokenizer)
token = tokens.pop(0)
assert token == ":", "ArithAction must have a ':' after the base segment" + " but got " + token
# Parse ops string
ops = {'+': [], '-': []}
while tokens[0] != cls.END_CHAR:
op_char = tokens.pop(0)
assert op_char in ["+", "-"], "ArithAction must have a '+' or '-' as an op char but got " + op_char
ops[op_char].append(parent_parser(tokens, tokenizer))
token = tokens.pop(0)
assert token == cls.END_CHAR, "ArithAction must end with " + cls.END_CHAR + " but got " + token
return cls(base_segment, ops)
-49
View File
@@ -1,49 +0,0 @@
from typing import List
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length, parse_numeric_arg
from .base import MultiArgAction
from .types import SegOrAction
class AverageAction(MultiArgAction):
grammar = 'avg(" arg "|" arg "|" arg ")"'
display_name = "Average"
action_name = "avg"
description = "Performs a weighted average between two segments or actions. The recommended weight is 0 - 1."
usage_examples = [
"avg(The cat is|The dog is|0.5)",
"avg(Cat|Dog|0.5)",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 3:
raise ValueError("Average action expects exactly three arguments (2 vectors and a weight)")
self.first_arg = args[0]
self.second_arg = args[1]
self.parsed_weight = parse_numeric_arg(
args[2], action_name="Average", role="weight", cast=float
)
first_len = get_total_length(self.first_arg)
second_len = get_total_length(self.second_arg)
if first_len != second_len:
raise ValueError(
f"Average start and end arguments should have the same length. Got {first_len} and {second_len}"
)
if self.parsed_weight < 0 or self.parsed_weight > 1:
print(f"WARNING: Average weight should be between 0 and 1. Got {self.parsed_weight}")
def token_length(self) -> int:
return get_total_length(self.first_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
start = concat_embeddings(self.first_arg, embedding_module)
end = concat_embeddings(self.second_arg, embedding_module)
return start * (1 - self.parsed_weight) + end * self.parsed_weight
+73 -51
View File
@@ -1,73 +1,95 @@
from abc import ABC, abstractmethod
from dataclasses import dataclass
from enum import Enum
from typing import List, Optional, Tuple, Union
from typing import Callable, Union
import torch
from torch import Tensor
from torch.nn import Embedding
class ActionArity(Enum):
NONE = 0
SINGLE = 1
MULTI = 2
@dataclass
class PostModifiers:
"""Optional position-embedding tweaks an action can request for its token range.
`start_idx` / `end_idx` are filled in by the encoder once the action's position in
the final token stream is known.
"""
position_embed_scale: Optional[float] = None
bypass_pos_embed: bool = False
start_idx: int = 0
end_idx: int = 0
ActionResult = Union[Tensor, Tuple[Tensor, PostModifiers]]
# Tokenizer placeholder for the 2nd..Nth slots of a multi-token Action, so each
# row stays exactly max_length entries (required for comfy's per-position weight
# indexing). process_tokens drops these; the Action's tensor fills the slots.
ACTION_CONTINUATION = object()
from comfy.sd1_clip import SD1Tokenizer
class Action(ABC):
arity: ActionArity = ActionArity.NONE
display_name: str = ""
action_name: str = ""
description: str = ""
grammar: str = ""
usage_examples: List[str] = []
@property
@abstractmethod
def START_CHAR(self):
pass
@property
@abstractmethod
def END_CHAR(self):
pass
@abstractmethod
def __init__(self, *args, **kwargs) -> None: ...
def token_length(self):
pass
@abstractmethod
def token_length(self) -> int: ...
def get_all_segments(self):
pass
@abstractmethod
def get_result(self, embedding_module: Embedding) -> ActionResult: ...
def get_result(self, embedding_module: Embedding):
pass
@classmethod
@abstractmethod
def parse_segment(
cls,
tokens: list[str],
start_chars: list[str],
end_chars: list[str],
parent_parser: Callable[[list[str], SD1Tokenizer], Union[str, 'Action']],
tokenizer: SD1Tokenizer,
) -> 'Action':
pass
class SingleArgAction(Action, ABC):
arity = ActionArity.SINGLE
def depth_repr(self, depth=1):
raise NotImplementedError()
def __init__(self, arg: List):
self.arg = arg
class PromptSegment:
def __init__(self, text: str, tokens: list[Union[int, Tensor]]):
self.text = text
self.tokens = tokens
def __repr__(self) -> str:
return f"{self.action_name}({self.arg})"
def token_length(self):
return len(self.tokens)
def get_embeddings(self, embedding_module: Embedding):
tensors = torch.LongTensor(self.tokens).to(torch.device('cpu'))
unsqueezed_tensors = tensors.unsqueeze(0)
return embedding_module(unsqueezed_tensors)
class MultiArgAction(Action, ABC):
arity = ActionArity.MULTI
def depth_repr(self, depth=1):
out = f'"{self.text}"('
def __init__(self, args: List[List]):
self.all_args = args
cleaned_tokens = list(map(lambda x: str(x) if isinstance(x, int) else "EMBD", self.tokens))
out += ", ".join(cleaned_tokens)
def __repr__(self) -> str:
joined = " | ".join(str(a) for a in self.all_args)
return f"{self.action_name}({joined})"
out += ")"
return out
def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment:
split_text = text.split(" ")
tokens = []
for word in split_text:
if word.startswith(tokenizer.embedding_identifier) and tokenizer.embedding_directory is not None:
embedding_name = word[len(tokenizer.embedding_identifier):].strip('\n')
get_embed_ret = tokenizer._try_get_embedding(embedding_name)
embedding = get_embed_ret[0]
leftover = get_embed_ret[1]
if embedding is None:
print(f"warning, embedding:{embedding_name} does not exist, ignoring")
else:
if len(embedding.shape) == 1:
tokens.append(embedding)
else:
tokens.extend(embedding)
if leftover != "":
word = leftover
else:
continue
tokens.extend(tokenizer.tokenizer(word)["input_ids"][1:-1])
return PromptSegment(text, tokens)
-39
View File
@@ -1,39 +0,0 @@
from typing import List
from torch import Tensor
from torch.nn import Embedding
from .action_utils import add_with_broadcast, concat_embeddings
from .base import MultiArgAction
from .types import SegOrAction
class DiffAction(MultiArgAction):
grammar = 'diff(" arg ("|" arg)* ")"'
display_name = "Difference"
action_name = "diff"
description = (
"Subtracts the segments in the order they are given. "
"The first segment is subtracted from the second, then the third from the result, and so on."
)
usage_examples = [
"diff(The cat is|The dog is)",
"diff(Cat|Dog)",
"sum(diff(king|man)|woman)",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
self.base_arg = args[0]
self.additional_args = args[1:]
def token_length(self) -> int:
return sum(s.token_length() for s in self.base_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
result = concat_embeddings(self.base_arg, embedding_module)
for arg in self.additional_args:
arg_embedding = concat_embeddings(arg, embedding_module)
result = add_with_broadcast(result, arg_embedding, op="sub")
return result
+17
View File
@@ -0,0 +1,17 @@
from custom_nodes.ClipStuff.lib.actions import ALL_START_CHARS, ALL_END_CHARS
from custom_nodes.ClipStuff.lib.actions.base import Action
def is_action_segment(action_class: Action.__class__, segment: str):
if not issubclass(action_class, Action):
raise Exception(
f"action_class must be a subclass of Action, got {action_class}"
)
return (
segment[0] == action_class.START_CHAR and segment[-1] == action_class.END_CHAR
)
def is_any_action_segment(segment: str):
return segment[0] in ALL_START_CHARS and segment[-1] in ALL_END_CHARS
-36
View File
@@ -1,36 +0,0 @@
from typing import List
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length, parse_numeric_arg
from .base import MultiArgAction
from .types import SegOrAction
class MultiplyAction(MultiArgAction):
grammar = 'mult(" arg+ ")"'
display_name = "Multiply"
action_name = "mult"
description = "Multiplies the provided segments or actions by the multiplier."
usage_examples = [
"mult(The cat is|2.5)",
"mult(Cat|-1)",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 2:
raise ValueError("Multiply action expects exactly two arguments")
self.target_arg = args[0]
self.parsed_multiplier = parse_numeric_arg(
args[1], action_name="Multiply", role="multiplier", cast=float
)
def token_length(self) -> int:
return get_total_length(self.target_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
return concat_embeddings(self.target_arg, embedding_module) * self.parsed_multiplier
-45
View File
@@ -1,45 +0,0 @@
from typing import List
import torch
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, parse_numeric_arg
from .base import MultiArgAction
from .types import SegOrAction
class NearestAction(MultiArgAction):
grammar = 'nearest(" arg ("|" arg)? ")"'
display_name = "Nearest Vocab"
action_name = "nearest"
description = (
"Snaps a computed vector to the k nearest real vocabulary tokens (by cosine similarity), "
"returning their embeddings concatenated. The input is mean-pooled before lookup."
)
usage_examples = [
"nearest(sum(diff(king|man)|woman))",
"nearest(sum(red|blue)|3)",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) not in (1, 2):
raise ValueError("nearest expects one or two arguments: nearest(expr) or nearest(expr|k)")
self.expr_arg = args[0]
self.k = parse_numeric_arg(args[1], action_name="nearest", role="k", cast=int) if len(args) == 2 else 1
def token_length(self) -> int:
return self.k
def get_result(self, embedding_module: Embedding) -> Tensor:
weight = embedding_module.weight.to(torch.float32)
weight_norm = torch.nn.functional.normalize(weight, dim=-1)
expr = concat_embeddings(self.expr_arg, embedding_module).to(torch.float32)
query = torch.nn.functional.normalize(expr.mean(dim=1), dim=-1)
sims = query @ weight_norm.T
top_ids = sims.topk(self.k, dim=-1).indices.squeeze(0)
return weight[top_ids].unsqueeze(0)
-23
View File
@@ -1,23 +0,0 @@
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length
from .base import SingleArgAction
class NegAction(SingleArgAction):
grammar = 'neg(" arg+ ")"'
display_name = "Negate"
action_name = "neg"
description = "Negates the provided segments or actions."
usage_examples = [
"neg(cat)",
"sum(king|neg(man)|women)",
]
def token_length(self) -> int:
return get_total_length(self.arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
return concat_embeddings(self.arg, embedding_module) * -1
-34
View File
@@ -1,34 +0,0 @@
from typing import List
import torch
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length, parse_numeric_arg
from .base import MultiArgAction
from .types import SegOrAction
class NoiseAction(MultiArgAction):
grammar = 'noise(" arg "|" arg ")"'
display_name = "Noise"
action_name = "noise"
description = "Adds Gaussian noise (mean 0, given std) to the embeddings of the first argument."
usage_examples = [
"A noise(cat|0.05) on a sunny day",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 2:
raise ValueError("noise expects exactly two arguments: noise(a|std)")
self.a_arg = args[0]
self.std = parse_numeric_arg(args[1], action_name="noise", role="std", cast=float)
def token_length(self) -> int:
return get_total_length(self.a_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
a = concat_embeddings(self.a_arg, embedding_module)
return a + torch.randn_like(a) * self.std
-25
View File
@@ -1,25 +0,0 @@
import torch
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length
from .base import SingleArgAction
class NormAction(SingleArgAction):
grammar = 'norm(" arg+ ")"'
display_name = "Normalize"
action_name = "norm"
description = "Normalizes the provided segments or actions."
usage_examples = [
"norm(cat)",
"sum(cat|norm(sum(tiger|fish)))",
]
def token_length(self) -> int:
return get_total_length(self.arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
embeddings = concat_embeddings(self.arg, embedding_module)
return torch.div(embeddings, torch.norm(embeddings, dim=-1, keepdim=True))
+128
View File
@@ -0,0 +1,128 @@
from typing import Optional, Union, Callable
import torch
from torch.nn import Embedding
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment
class NudgeAction(Action):
def token_length(self):
# Nudge nudges the embeddings of the base segment, so the length is the length of the base segment
if isinstance(self.base_segment, Action):
return self.base_segment.token_length()
return len(self.base_segment.tokens)
def get_all_segments(self):
segments = []
if isinstance(self.base_segment, Action):
segments += self.base_segment.get_all_segments()
else:
segments.append(self.base_segment)
if isinstance(self.target, Action):
segments += self.target.get_all_segments()
else:
segments.append(self.target)
return segments
def get_result(self, embedding_module: Embedding):
if isinstance(self.base_segment, Action):
base_segment_result = self.base_segment.get_result(embedding_module)
else:
base_segment_result = self.base_segment.get_embeddings(embedding_module)
if isinstance(self.target, Action):
target_segment_result = self.target.get_result(embedding_module)
else:
target_segment_result = self.target.get_embeddings(embedding_module)
base_mean = torch.mean(base_segment_result, dim=1, keepdim=True)
if target_segment_result.shape[1] == 1:
translation_vector = target_segment_result - base_mean
else:
translation_vector = torch.mean(target_segment_result, dim=1, keepdim=True) - base_mean
return base_segment_result.add(translation_vector, alpha=self.weight)
START_CHAR = "["
END_CHAR = "]"
def __init__(
self,
base_segment: PromptSegment | Action,
target: Union[PromptSegment, Action],
weight: Optional[float] = None,
):
self.base_segment = base_segment
self.weight = weight
self.target = target
def __repr__(self):
return f"NudgeAction(\n\tbase_segment={self.base_segment},\n\ttarget={self.target},\n\tweight={self.weight}\n)"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_segment, Action):
base_segment_repr = self.base_segment.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f'base_segment={self.base_segment.depth_repr()},\n'
if isinstance(self.target, Action):
target_repr = self.target.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.target.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
@classmethod
def parse_segment(
cls,
tokens: list[str],
start_chars: list[str],
end_chars: list[str],
parent_parser: Callable[[list[str], SD1Tokenizer], PromptSegment | Action],
tokenizer: SD1Tokenizer,
) -> Action:
"""
Parse a nudge action from a list of tokens
Supported formats:
[base_segment:target_segment]
[base_segment:target_segment:weight]
Weight is optional, if not provided it will be None
:param tokens: List of tokens, will be modified
:param start_chars: List of start chars for all actions
:param end_chars: List of end chars for all actions
:param parent_parser: Function to parse segments to allow for nested actions
:return:
"""
token = tokens.pop(0)
assert token == cls.START_CHAR, "NudgeAction must start with " + cls.START_CHAR + " got " + token
# Parse base segment
base_segment = parent_parser(tokens, tokenizer)
token = tokens.pop(0)
assert token == ":", "NudgeAction must have a ':' after the base segment" + " but got " + token
# Parse target segment
target_segment = parent_parser(tokens, tokenizer)
# Parse weight if it exists
weight = None
if tokens[0] == ":":
# Parse weight
tokens.pop(0)
weight = float(tokens.pop(0))
token = tokens.pop(0)
assert token == cls.END_CHAR, "NudgeAction must end with " + cls.END_CHAR + " got " + token
return cls(base_segment, target_segment, weight)
-37
View File
@@ -1,37 +0,0 @@
from typing import List, Tuple
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length, parse_numeric_arg
from .base import MultiArgAction, PostModifiers
from .types import SegOrAction
class PosScaleAction(MultiArgAction):
grammar = 'posScale(" arg+ ")"'
display_name = "Positional Embedding Scale"
action_name = "posScale"
description = (
"Scales (multiplies) the positional embeddings of the provided segments or actions by the multiplier."
)
usage_examples = [
"A posScale(cat|1.5) on a rainy day",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 2:
raise ValueError("PosScale action expects exactly two arguments")
self.target_arg = args[0]
self.parsed_multiplier = parse_numeric_arg(
args[1], action_name="PosScale", role="multiplier", cast=float
)
def token_length(self) -> int:
return get_total_length(self.target_arg)
def get_result(self, embedding_module: Embedding) -> Tuple[Tensor, PostModifiers]:
target_embeddings = concat_embeddings(self.target_arg, embedding_module)
return target_embeddings, PostModifiers(position_embed_scale=self.parsed_multiplier)
-24
View File
@@ -1,24 +0,0 @@
from typing import Tuple
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length
from .base import PostModifiers, SingleArgAction
class PostPosAction(SingleArgAction):
grammar = 'postPos(" arg+ ")"'
display_name = "Ignore Positional Embeddings"
action_name = "postPos"
description = "Prevents positional embeddings from being applied to the provided segments or actions."
usage_examples = [
"A postPos(cat) on a rainy day",
]
def token_length(self) -> int:
return get_total_length(self.arg)
def get_result(self, embedding_module: Embedding) -> Tuple[Tensor, PostModifiers]:
return concat_embeddings(self.arg, embedding_module), PostModifiers(bypass_pos_embed=True)
-72
View File
@@ -1,72 +0,0 @@
from typing import List
import torch
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length
from .base import MultiArgAction
from .types import SegOrAction
def _direction(args: List[SegOrAction], embedding_module: Embedding) -> Tensor:
"""Mean unit direction of an arg's embeddings: [1, 1, hidden]."""
emb = concat_embeddings(args, embedding_module)
mean = emb.mean(dim=1, keepdim=True)
return torch.nn.functional.normalize(mean, dim=-1)
def _project(a: Tensor, b_hat: Tensor) -> Tensor:
coeff = (a * b_hat).sum(dim=-1, keepdim=True)
return coeff * b_hat
class ProjectAction(MultiArgAction):
grammar = 'proj(" arg "|" arg ")"'
display_name = "Project"
action_name = "proj"
description = "Projects the first argument onto the direction of the second (mean, unit-normalized)."
usage_examples = [
"proj(king|gender)",
"diff(style|proj(style|photorealistic))",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 2:
raise ValueError("proj expects exactly two arguments: proj(a|b)")
self.a_arg = args[0]
self.b_arg = args[1]
def token_length(self) -> int:
return get_total_length(self.a_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
a = concat_embeddings(self.a_arg, embedding_module)
return _project(a, _direction(self.b_arg, embedding_module))
class RejectAction(MultiArgAction):
grammar = 'reject(" arg "|" arg ")"'
display_name = "Reject"
action_name = "reject"
description = "Removes the component of the first argument along the direction of the second (a - proj(a|b))."
usage_examples = [
"reject(anime girl|anime)",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 2:
raise ValueError("reject expects exactly two arguments: reject(a|b)")
self.a_arg = args[0]
self.b_arg = args[1]
def token_length(self) -> int:
return get_total_length(self.a_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
a = concat_embeddings(self.a_arg, embedding_module)
return a - _project(a, _direction(self.b_arg, embedding_module))
-53
View File
@@ -1,53 +0,0 @@
from typing import List
import torch
from torch import Tensor
from torch.nn import Embedding
from .action_utils import parse_numeric_arg
from .base import MultiArgAction
from .types import SegOrAction
class RandAction(MultiArgAction):
grammar = 'rand(" arg ")"'
display_name = "Random Embedding"
action_name = "rand"
description = (
"Returns a random embedding of the specified token length, "
"with the values optionally bounded by the second and third arguments."
)
usage_examples = [
"A rand(1) cat",
"A rand(1|-1|1) cat",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) not in (1, 3):
raise ValueError("Random action expects exactly one or three arguments")
self.parsed_token_length = parse_numeric_arg(
args[0], action_name="Random", role="first argument (token length)", cast=int
)
if len(args) == 3:
self.range_min = parse_numeric_arg(
args[1], action_name="Random", role="second argument (min)", cast=int
)
self.range_max = parse_numeric_arg(
args[2], action_name="Random", role="third argument (max)", cast=int
)
if self.range_min > self.range_max:
raise ValueError("Random action min must be <= max")
else:
self.range_min = 0
self.range_max = 1
def token_length(self) -> int:
return self.parsed_token_length
def get_result(self, embedding_module: Embedding) -> Tensor:
return torch.empty(
1, self.parsed_token_length, embedding_module.embedding_dim
).uniform_(self.range_min, self.range_max)
-37
View File
@@ -1,37 +0,0 @@
from typing import List
import torch
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length
from .base import MultiArgAction
from .types import SegOrAction
class RenormAction(MultiArgAction):
grammar = 'renorm(" arg "|" arg ")"'
display_name = "Renormalize"
action_name = "renorm"
description = "Rescales the first argument so each token's L2 norm matches the (mean) L2 norm of the reference."
usage_examples = [
"renorm(sum(king|neg(man)|woman)|queen)",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 2:
raise ValueError("renorm expects exactly two arguments: renorm(a|ref)")
self.a_arg = args[0]
self.ref_arg = args[1]
def token_length(self) -> int:
return get_total_length(self.a_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
a = concat_embeddings(self.a_arg, embedding_module)
ref = concat_embeddings(self.ref_arg, embedding_module)
a_norm = torch.norm(a, dim=-1, keepdim=True).clamp(min=1e-8)
ref_norm = torch.norm(ref, dim=-1, keepdim=True).mean()
return a * (ref_norm / a_norm)
-68
View File
@@ -1,68 +0,0 @@
from typing import List, Tuple
from torch import Tensor
from torch.nn import Embedding
from ..parser.prompt_segment import PromptSegment
from .action_utils import concat_embeddings, get_total_length
from .base import Action, MultiArgAction
from .types import SegOrAction
class ScaleDims(MultiArgAction):
grammar = 'scaleDims(" arg ("|" arg)* ")"'
display_name = "Scale Dimensions"
action_name = "scaleDims"
description = "Scales the specified dimensions of the input embeddings by the specified amount"
usage_examples = [
"The scaleDims(cat|4,1.5|76,1.2) is happy",
]
def __init__(self, args: List[List[SegOrAction]]):
super().__init__(args)
self.base_arg = args[0]
self.scale_args: List[Tuple[int, float]] = _parse_dim_value_pairs(args[1:], action_name="ScaleDims")
def token_length(self) -> int:
return get_total_length(self.base_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
embeddings = concat_embeddings(self.base_arg, embedding_module)
for dim, scale in self.scale_args:
embeddings[0, :, dim] *= scale
return embeddings
def _parse_dim_value_pairs(
args: List[List[SegOrAction]],
*,
action_name: str,
) -> List[Tuple[int, float]]:
"""Parse args of the form `<dim>,<value>` into `(int, float)` pairs.
Used by both scaleDims and setDims.
"""
pairs: List[Tuple[int, float]] = []
for arg in args:
if isinstance(arg, Action):
raise ValueError(f"{action_name} args must be in the format <dim>,<value> but got an action")
if len(arg) != 1:
raise ValueError(f"{action_name} args must be a single segment of <dim>,<value>")
seg = arg[0]
assert isinstance(seg, PromptSegment)
if "," not in seg.text:
raise ValueError(f"{action_name} args must be <dim>,<value> but got: {seg.text!r}")
dim_str, value_str = seg.text.split(",", 1)
try:
dim = int(dim_str)
except ValueError:
raise ValueError(f"{action_name} dim must be an integer; got {dim_str!r}")
try:
value = float(value_str)
except ValueError:
raise ValueError(f"{action_name} value must be a float; got {value_str!r}")
pairs.append((dim, value))
return pairs
-34
View File
@@ -1,34 +0,0 @@
from typing import List, Tuple
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length
from .base import MultiArgAction
from .scale_dims import _parse_dim_value_pairs
from .types import SegOrAction
class SetDims(MultiArgAction):
grammar = 'setDims(" arg ("|" arg)* ")"'
display_name = "Set Dimensions"
action_name = "setDims"
description = "Sets the specified dimensions of the input embeddings to the specified value"
usage_examples = [
"The setDims(cat|4, -0.01253|76, 1.2) is happy",
]
def __init__(self, args: List[List[SegOrAction]]):
super().__init__(args)
self.base_arg = args[0]
self.value_args: List[Tuple[int, float]] = _parse_dim_value_pairs(args[1:], action_name="SetDims")
def token_length(self) -> int:
return get_total_length(self.base_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
embeddings = concat_embeddings(self.base_arg, embedding_module)
for dim, value in self.value_args:
embeddings[0, :, dim] = value
return embeddings
-52
View File
@@ -1,52 +0,0 @@
from typing import List
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length, parse_numeric_arg
from .base import MultiArgAction
from .types import SegOrAction
from .utils import slerp
class SlerpAction(MultiArgAction):
grammar = 'slerp(" arg "|" arg "|" arg ")"'
display_name = "Slerp"
action_name = "slerp"
description = (
"Performs a slerp (interpolation) between two segments or actions, with the given weight. "
"The recommended weight is 0 - 1."
)
usage_examples = [
"The slerp(cat|dog|0.5) is happy",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 3:
raise ValueError("Slerp action expects exactly three arguments (2 vectors and a weight)")
self.start_argument = args[0]
self.end_argument = args[1]
self.parsed_weight = parse_numeric_arg(
args[2], action_name="Slerp", role="weight", cast=float
)
start_len = get_total_length(self.start_argument)
end_len = get_total_length(self.end_argument)
if start_len != end_len:
raise ValueError(
f"Slerp start and end arguments should have the same length. Got {start_len} and {end_len}"
)
if self.parsed_weight < 0 or self.parsed_weight > 1:
print(f"WARNING: Slerp weight should be between 0 and 1. Got {self.parsed_weight}")
def token_length(self) -> int:
return get_total_length(self.start_argument)
def get_result(self, embedding_module: Embedding) -> Tensor:
start = concat_embeddings(self.start_argument, embedding_module)
end = concat_embeddings(self.end_argument, embedding_module)
return slerp(self.parsed_weight, start, end)
-34
View File
@@ -1,34 +0,0 @@
from typing import List
from torch import Tensor
from torch.nn import Embedding
from .action_utils import add_with_broadcast, concat_embeddings
from .base import MultiArgAction
from .types import SegOrAction
class SumAction(MultiArgAction):
grammar = 'sum(" arg ("|" arg)+ ")"'
display_name = "Sum"
action_name = "sum"
description = "Adds the embeddings of the provided segments or actions."
usage_examples = [
"A happy sum(cat|dog|shark)",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
self.base_arg = args[0]
self.additional_args = args[1:]
def token_length(self) -> int:
return sum(s.token_length() for s in self.base_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
result = concat_embeddings(self.base_arg, embedding_module)
for arg in self.additional_args:
arg_embedding = concat_embeddings(arg, embedding_module)
result = add_with_broadcast(result, arg_embedding, op="add")
return result
-7
View File
@@ -1,7 +0,0 @@
from typing import Union
from ..parser.prompt_segment import PromptSegment
from .base import Action
from .weighted import WeightedGroup
SegOrAction = Union[PromptSegment, Action, WeightedGroup]
-23
View File
@@ -1,23 +0,0 @@
import torch
def slerp(val: float, low: torch.Tensor, high: torch.Tensor, epsilon: float = 1e-5) -> torch.Tensor:
"""Spherical linear interpolation between two tensors along the last dim."""
val_t = torch.tensor(val, dtype=torch.float32, device=low.device).clamp(0, 1)
low_norm = low / torch.norm(low, dim=-1, keepdim=True)
high_norm = high / torch.norm(high, dim=-1, keepdim=True)
dot = (low_norm * high_norm).sum(-1, keepdim=True).clamp(-1, 1)
omega = torch.acos(dot)
sin_omega = torch.sin(omega)
scale_low = torch.sin((1.0 - val_t) * omega) / (sin_omega + epsilon)
scale_high = torch.sin(val_t * omega) / (sin_omega + epsilon)
# Fall back to linear interp where the angle is too small for stable slerp.
close = sin_omega < epsilon
scale_low = torch.where(close, 1.0 - val_t, scale_low)
scale_high = torch.where(close, val_t, scale_high)
return scale_low * low + scale_high * high
-15
View File
@@ -1,15 +0,0 @@
from typing import List
class WeightedGroup:
"""A group of segments/actions sharing an attention weight (the `(text:1.2)` syntax)."""
def __init__(self, items: List, weight: float):
self.items = items
self.weight = weight
def token_length(self) -> int:
return sum(item.token_length() for item in self.items)
def __repr__(self) -> str:
return f"({self.items}:{self.weight})"
+25
View File
@@ -0,0 +1,25 @@
{
"_name_or_path": "openai/clip-vit-large-patch14",
"architectures": [
"CLIPTextModel"
],
"attention_dropout": 0.0,
"bos_token_id": 0,
"dropout": 0.0,
"eos_token_id": 2,
"hidden_act": "quick_gelu",
"hidden_size": 768,
"initializer_factor": 1.0,
"initializer_range": 0.02,
"intermediate_size": 3072,
"layer_norm_eps": 1e-05,
"max_position_embeddings": 77,
"model_type": "clip_text_model",
"num_attention_heads": 12,
"num_hidden_layers": 12,
"pad_token_id": 1,
"projection_dim": 768,
"torch_dtype": "float32",
"transformers_version": "4.24.0",
"vocab_size": 49408
}
+175 -104
View File
@@ -1,129 +1,200 @@
"""DSL-aware CLIP text encoders.
The tokenizer emits ComfyUI's native `(token, weight)` format with one twist:
a `token` can also be a lazily-evaluated `Action`. We resolve those to tensors
here in `process_tokens` (where the embedding module is available) and delegate
everything else — embedding lookup, mask building, splice — to the stock
`SDClipModel.process_tokens`.
`posScale` / `postPos` actions return a `PostModifiers` alongside their tensor.
ComfyUI's `CLIPTextModel_.forward` adds the position embedding inline whenever
`embeds` is supplied, so we pre-bake `(modified - default)` into `embeds` such
that the transformer's add nets to `+ modified`.
"""
import dataclasses
from typing import List
import contextlib
import os
from typing import Union
import torch
from torch import Tensor
from torch.nn import Embedding
from transformers import CLIPTextConfig, modeling_utils
from comfy import sd1_clip, sdxl_clip
from comfy import model_management
import comfy.ops
from comfy.sd import CLIP
from custom_nodes.ClipStuff.lib.actions.base import PromptSegment, Action
from custom_nodes.ClipStuff.lib.fun_clip_stuff import MyCLIPTextModel
from custom_nodes.ClipStuff.lib.tokenizer import TokenDict
from .actions.base import ACTION_CONTINUATION, Action, PostModifiers
class SD1FunClipModel(torch.nn.Module):
"""Uses the CLIP transformer encoder for text (from huggingface)"""
LAYERS = [
"last",
"pooled",
"hidden"
]
def __init__(self, version="openai/clip-vit-large-patch14", device="cpu", max_length=77,
freeze=True, layer="last", layer_idx=None, textmodel_json_config=None,
textmodel_path=None): # clip-vit-base-patch32
super().__init__()
assert layer in self.LAYERS
self.num_layers = 12
if textmodel_path is not None:
self.transformer = MyCLIPTextModel.from_pretrained(textmodel_path)
else:
if textmodel_json_config is None:
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "clip_config.json")
config = CLIPTextConfig.from_json_file(textmodel_json_config)
self.num_layers = config.num_hidden_layers
with comfy.ops.use_comfy_ops():
with modeling_utils.no_init_weights():
self.transformer = MyCLIPTextModel(config)
class PromptLangSDClipModel(sd1_clip.SDClipModel):
def process_tokens(self, tokens, device): # type: ignore[override]
embedding_module = self.transformer.get_input_embeddings()
self.max_length = max_length
if freeze:
self.freeze()
self.layer = layer
self.layer_idx = None
self.empty_tokens = [[49406] + [49407] * 76]
self.text_projection = None
self.layer_norm_hidden_state = True
if layer == "hidden":
assert layer_idx is not None
assert abs(layer_idx) <= self.num_layers
self.clip_layer(layer_idx)
self.layer_default = (self.layer, self.layer_idx)
def freeze(self):
self.transformer = self.transformer.eval()
# self.train = disabled_train
for param in self.parameters():
param.requires_grad = False
def clip_layer(self, layer_idx):
if abs(layer_idx) >= self.num_layers:
self.layer = "last"
else:
self.layer = "hidden"
self.layer_idx = layer_idx
def reset_clip_layer(self):
self.layer = self.layer_default[0]
self.layer_idx = self.layer_default[1]
def set_up_textual_embeddings(self, tokens: list[list[PromptSegment | Action]], current_embeds):
next_new_token = token_dict_size = current_embeds.weight.shape[0] - 1
embedding_weights = []
# For each batch
for batch in tokens:
for seg_or_action in batch:
if isinstance(seg_or_action, Action):
segments = seg_or_action.get_all_segments()
else:
segments = [seg_or_action]
for segment in segments:
tokens_temp = []
segment_length = segment.token_length()
for tid_or_tensor in segment.tokens:
if isinstance(tid_or_tensor, int):
if tid_or_tensor == token_dict_size: # Is EOS token
tid_or_tensor = -1 # Set to -1 so that it can be replaced with the EOS token later
tokens_temp += [tid_or_tensor]
else:
if tid_or_tensor.shape[0] == current_embeds.weight.shape[1]:
embedding_weights += [tid_or_tensor]
tokens_temp += [next_new_token]
next_new_token += 1
else:
print("WARNING: shape mismatch when trying to apply embedding, embedding will be ignored",
tid_or_tensor.shape[0], current_embeds.weight.shape[1])
if len(tokens_temp) < segment_length:
# Pretty sure this is only needed if the embedding is not the same size as the CLIP embedding
print("WARNING: segment length mismatch, padding with EOS token")
tokens_temp.extend([self.empty_tokens[0][-1] * (segment_length - len(tokens_temp))])
segment.tokens = tokens_temp
n = token_dict_size
if len(embedding_weights) > 0:
# Create new embedding, with size of current embedding + number of new embeddings
new_embedding = torch.nn.Embedding(next_new_token + 1, current_embeds.weight.shape[1],
device=current_embeds.weight.device, dtype=current_embeds.weight.dtype)
# Copy current embedding weights to new embedding
new_embedding.weight[:token_dict_size] = current_embeds.weight[:-1]
# Add new embeddings
for embed in embedding_weights:
new_embedding.weight[n] = embed
n += 1
# Set re-add the EOS token
new_embedding.weight[n] = current_embeds.weight[-1] # EOS embedding
self.transformer.set_input_embeddings(new_embedding)
resolved: List[list] = []
pos_modifiers_per_batch: List[List[PostModifiers]] = []
for batch in tokens:
row: list = []
modifiers: List[PostModifiers] = []
position = 0
for entry in batch:
if entry is ACTION_CONTINUATION:
# Slot already accounted for by the preceding Action's `position += length`.
continue
if isinstance(entry, Action):
length = entry.token_length()
result = entry.get_result(embedding_module)
if isinstance(result, tuple):
tensor, mods = result
modifiers.append(
dataclasses.replace(mods, start_idx=position, end_idx=position + length)
)
else:
tensor = result
row.append(tensor)
position += length
for seg_or_action in batch:
if isinstance(seg_or_action, Action):
segments = seg_or_action.get_all_segments()
else:
row.append(entry)
position += 1
resolved.append(row)
pos_modifiers_per_batch.append(modifiers)
segments = [seg_or_action]
embeds, attention_mask, num_tokens, embeds_info = super().process_tokens(resolved, device)
for segment in segments:
for tokenIdx in range(len(segment.tokens)):
if segment.tokens[tokenIdx] == -1:
segment.tokens[tokenIdx] = n
if any(pos_modifiers_per_batch):
embeds = _apply_pos_modifiers(
embeds, pos_modifiers_per_batch, self._get_position_embedding()
)
def forward(self, tokens, **kwargs):
backup_embeds = self.transformer.get_input_embeddings()
device = backup_embeds.weight.device
self.set_up_textual_embeddings(tokens, backup_embeds)
# tokens = torch.LongTensor(tokens).to(device)
return embeds, attention_mask, num_tokens, embeds_info
def _get_position_embedding(self) -> Embedding:
"""Isolated so a ComfyUI internal layout change only needs one fix."""
return self.transformer.text_model.embeddings.position_embedding
if backup_embeds.weight.dtype != torch.float32:
precision_scope = torch.autocast
else:
precision_scope = contextlib.nullcontext
class PromptLangSDXLClipG(sdxl_clip.SDXLClipG, PromptLangSDClipModel):
"""SDXL's larger CLIP-G text encoder, with our DSL-aware process_tokens."""
if (kwargs.get("position_ids", None) is not None):
position_ids = torch.LongTensor(kwargs["position_ids"]).to(device)
else:
position_ids = None
def _apply_pos_modifiers(
embeds: Tensor,
pos_modifiers_per_batch: List[List[PostModifiers]],
position_embedding: Embedding,
) -> Tensor:
seq_len = embeds.shape[1]
pos_weights = position_embedding.weight[:seq_len].to(device=embeds.device, dtype=embeds.dtype)
with precision_scope(model_management.get_autocast_device(device)):
outputs = self.transformer(input_ids=tokens, output_hidden_states=self.layer == "hidden",
position_ids=position_ids)
self.transformer.set_input_embeddings(backup_embeds)
out = embeds.clone()
for batch_idx, modifiers in enumerate(pos_modifiers_per_batch):
for mod in modifiers:
default_slice = pos_weights[mod.start_idx:mod.end_idx]
if mod.bypass_pos_embed:
modified_slice = torch.zeros_like(default_slice)
elif mod.position_embed_scale is not None:
modified_slice = default_slice * float(mod.position_embed_scale)
if self.layer == "last":
z = outputs.last_hidden_state
elif self.layer == "pooled":
z = outputs.pooler_output[:, None, :]
else:
continue
z = outputs.hidden_states[self.layer_idx]
if self.layer_norm_hidden_state:
z = self.transformer.text_model.final_layer_norm(z)
# The transformer will add `default_slice` back; net effect is `+ modified_slice`.
out[batch_idx, mod.start_idx:mod.end_idx] += modified_slice - default_slice
pooled_output = outputs.pooler_output
if self.text_projection is not None:
pooled_output = pooled_output.to(self.text_projection.device) @ self.text_projection
return z.float(), pooled_output.float()
return out
def encode(self, tokens, **kwargs):
return self(tokens, **kwargs)
def load_sd(self, sd):
return self.transformer.load_state_dict(sd, strict=False)
class PromptLangSD1ClipModel(sd1_clip.SD1ClipModel):
def __init__(self, device="cpu", dtype=None, model_options=None, **kwargs):
super().__init__(
device=device,
dtype=dtype,
model_options=model_options or {},
clip_name="l",
clip_model=PromptLangSDClipModel,
**kwargs,
)
def encode_token_weights(self, prompt_segments: list[list[Union[PromptSegment | Action]]], **kwargs):
to_encode = [[PromptSegment(text="_Empty Batch_", tokens=self.empty_tokens[0])]]
for batch in prompt_segments:
to_encode.append(batch)
out, pooled = self.encode(to_encode, **kwargs)
z_empty = out[0:1]
if pooled.shape[0] > 1:
first_pooled = pooled[1:2]
else:
first_pooled = pooled[0:1]
class PromptLangSDXLClipModel(sdxl_clip.SDXLClipModel):
def __init__(self, device="cpu", dtype=None, model_options=None) -> None:
torch.nn.Module.__init__(self)
opts = model_options or {}
self.clip_l = PromptLangSDClipModel(
layer="hidden",
layer_idx=-2,
device=device,
dtype=dtype,
layer_norm_hidden_state=False,
model_options=opts,
)
self.clip_g = PromptLangSDXLClipG(device=device, dtype=dtype, model_options=opts)
self.dtypes = {dtype} if dtype is not None else set()
output = []
for k in range(1, out.shape[0]):
z = out[k:k + 1]
# for i in range(len(z)):
# for j in range(len(z[i])):
# weight = token_dicts[k - 1][j][0].weight
# z[i][j] = (z[i][j] - z_empty[0][j]) * weight + z_empty[0][j]
output.append(z)
if (len(output) == 0):
return z_empty.cpu(), first_pooled.cpu()
return torch.cat(output, dim=-2).cpu(), first_pooled.cpu()
+217
View File
@@ -0,0 +1,217 @@
from typing import Optional, Tuple, Union
import torch
from torch import device
from transformers import CLIPTextConfig
from transformers.modeling_outputs import BaseModelOutputWithPooling
from transformers.models.clip.modeling_clip import _expand_mask, CLIPTextEmbeddings, CLIPTextTransformer, \
CLIPTextModel
from custom_nodes.ClipStuff.lib.actions.base import PromptSegment, Action
from custom_nodes.ClipStuff.lib.tokenizer import TokenDict
def slerp(val, low, high):
low = low.unsqueeze(0)
high = high.unsqueeze(0)
low_norm = low/torch.norm(low, dim=1, keepdim=True)
high_norm = high/torch.norm(high, dim=1, keepdim=True)
omega = torch.acos((low_norm*high_norm).sum(1))
so = torch.sin(omega)
res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high
return res
class MyCLIPTextEmbeddings(CLIPTextEmbeddings):
def __init__(self, config: CLIPTextConfig):
super().__init__(config)
def forward(
self,
input_dicts: Optional[list[list[tuple[TokenDict]]]] = None,
input_ids: Optional[torch.LongTensor] = None,
position_ids: Optional[torch.LongTensor] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
) -> torch.Tensor:
batches = []
for batch_idx, batch in enumerate(input_dicts):
results = []
for seg_or_action in batch:
if isinstance(seg_or_action, Action):
results.append(seg_or_action.get_result(self.token_embedding))
else:
results.append(seg_or_action.get_embeddings(self.token_embedding))
batches.append(results)
seq_length = batches[0][0].shape[-2]
if position_ids is None:
position_ids = self.position_ids[:, :seq_length]
# if inputs_embeds is None:
# inputs_embeds = self.token_embedding(input_ids)
# for batch_idx, batch in enumerate(input_dicts):
# for token_idx, token in enumerate(batch):
# if token[0].nudge_id is not None:
# nudged_embed = inputs_embeds[batch_idx, token_idx][:] + self.token_embedding(torch.LongTensor([token[0].nudge_id]).to(torch.device('cpu')))[0]
# if token[0].nudge_index_start is not None and token[0].nudge_index_stop is not None:
# nudge_start = token[0].nudge_index_start
# nudge_end = token[0].nudge_index_stop
# else:
# nudge_start = 0
# nudge_end = 768
# inputs_embeds[batch_idx, token_idx][nudge_start:nudge_end] = (slerp(token[0].nudge_weight, inputs_embeds[batch_idx, token_idx][:], nudged_embed)[0][nudge_start:nudge_end])
# elif token[0].arith_ops is not None:
# for op, id_list in token[0].arith_ops.items():
# if op == '+':
# for this_id in id_list:
# inputs_embeds[batch_idx, token_idx] += self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0]
# elif op == '-':
# for this_id in id_list:
# inputs_embeds[batch_idx, token_idx] -= self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0]
embeds = []
for batch in batches:
if len(batch) == 1:
embeds.append(batch[0])
else:
embeds.append(torch.cat(batch, dim=-2))
position_embeddings = self.position_embedding(position_ids)
embeddings = torch.cat(embeds, dim=0) + position_embeddings
return embeddings
class MyCLIPTextTransformer(CLIPTextTransformer):
def __init__(self, config: CLIPTextConfig):
super().__init__(config)
self.embeddings = MyCLIPTextEmbeddings(config)
def forward(
self,
input_ids: Optional[list[list[PromptSegment | Action]]] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
) -> Union[Tuple, BaseModelOutputWithPooling]:
r"""
Returns:
"""
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
)
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
if input_ids is None:
raise ValueError("You have to specify input_ids")
# input_shape = input_ids.size()
# input_ids = input_ids.view(-1, input_shape[-1])
hidden_states = self.embeddings(input_dicts=input_ids)
bsz = len(input_ids)
# TODO: Properly gather this
seq_len = 77
# bsz, seq_len = input_shape
# CLIP's text model uses causal mask, prepare it here.
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(
hidden_states.device
)
# expand attention_mask
if attention_mask is not None:
# [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
attention_mask = _expand_mask(attention_mask, hidden_states.dtype)
encoder_outputs = self.encoder(
inputs_embeds=hidden_states,
attention_mask=attention_mask,
causal_attention_mask=causal_attention_mask,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
return_dict=return_dict,
)
last_hidden_state = encoder_outputs[0]
last_hidden_state = self.final_layer_norm(last_hidden_state)
# Hacky way to get idx of first EOT token
eot_idx = [1]
for batch in input_ids[1:]:
idx = 0
for seg_or_action in batch:
if isinstance(seg_or_action, Action):
idx += seg_or_action.token_length()
else:
if seg_or_action.text == '__PAD__':
break
eot_idx.append(idx)
# text_embeds.shape = [batch_size, sequence_length, transformer.width]
# take features from the eot embedding (eot_token is the highest number in each sequence)
# casting to torch.int for onnx compatibility: argmax doesn't support int64 inputs with opset 14
# TODO: Get the index of the first EOT token
pooled_output = last_hidden_state[
torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device),
eot_idx
]
if not return_dict:
return (last_hidden_state, pooled_output) + encoder_outputs[1:]
return BaseModelOutputWithPooling(
last_hidden_state=last_hidden_state,
pooler_output=pooled_output,
hidden_states=encoder_outputs.hidden_states,
attentions=encoder_outputs.attentions,
)
class MyCLIPTextModel(CLIPTextModel):
def __init__(self, config: CLIPTextConfig):
super().__init__(config)
self.text_model = MyCLIPTextTransformer(config)
def forward(
self,
input_ids: Optional[list[list[tuple[TokenDict]]]] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
) -> Union[Tuple, BaseModelOutputWithPooling]:
r"""
Returns:
Examples:
```python
>>> from transformers import AutoTokenizer, CLIPTextModel
>>> model = CLIPTextModel.from_pretrained("openai/clip-vit-base-patch32")
>>> tokenizer = AutoTokenizer.from_pretrained("openai/clip-vit-base-patch32")
>>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="pt")
>>> outputs = model(**inputs)
>>> last_hidden_state = outputs.last_hidden_state
>>> pooled_output = outputs.pooler_output # pooled (EOS token) states
```"""
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
return self.text_model(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
return_dict=return_dict,
)
-68
View File
@@ -1,68 +0,0 @@
"""Debug helper: report what a DSL prompt resolves to at the embedding layer."""
from typing import List, Tuple
import torch
from .actions.base import ACTION_CONTINUATION, Action
def inspect_prompt(clip, text: str, top_k: int = 3) -> str:
"""Tokenize + resolve actions and report per-slot L2 norm and nearest vocab tokens.
Runs only the embedding lookup (no transformer forward), so it's cheap.
"""
inner_clip, inner_tok = _unwrap(clip)
embedding_module = inner_clip.transformer.get_input_embeddings()
weight = embedding_module.weight.to(torch.float32)
weight_norm = torch.nn.functional.normalize(weight, dim=-1)
batches = inner_tok.tokenize_with_weights(text)
lines = [f"Prompt: {text!r}", ""]
for batch_idx, batch in enumerate(batches):
lines.append(f"-- batch {batch_idx} ({len(batch)} entries) --")
lines.append(f"{'idx':>3} {'w':>5} {'src':<24} {'L2':>6} nearest")
position = 0
for token, w in batch:
if token is ACTION_CONTINUATION:
position += 1
continue
embeds, source = _resolve(token, embedding_module)
for row in embeds:
norm = torch.norm(row).item()
nearest = _nearest_vocab(row, weight_norm, inner_tok, top_k)
lines.append(f"{position:>3} {w:>5.2f} {source:<24.24} {norm:>6.3f} {nearest}")
position += 1
lines.append("")
return "\n".join(lines)
def _unwrap(clip):
"""Dig past SD1ClipModel/SDXL wrappers to the underlying SDClipModel + SDTokenizer."""
cond = clip.cond_stage_model
tok = clip.tokenizer
inner_clip = getattr(cond, getattr(cond, "clip", "clip_l"), cond)
inner_tok = getattr(tok, getattr(tok, "clip", "clip_l"), tok)
return inner_clip, inner_tok
def _resolve(token, embedding_module) -> Tuple[torch.Tensor, str]:
"""Map a token entry to its `[N, hidden]` embedding rows and a short source label."""
if isinstance(token, Action):
result = token.get_result(embedding_module)
tensor = result[0] if isinstance(result, tuple) else result
return tensor.reshape(-1, tensor.shape[-1]).to(torch.float32), repr(token)
if isinstance(token, int):
return embedding_module.weight[token : token + 1].to(torch.float32), f"tok#{token}"
# Inline TI tensor.
return token.reshape(-1, token.shape[-1]).to(torch.float32), "embedding:"
def _nearest_vocab(row: torch.Tensor, weight_norm: torch.Tensor, tokenizer, top_k: int) -> str:
row_norm = torch.nn.functional.normalize(row.unsqueeze(0), dim=-1)
sims = (row_norm @ weight_norm.T).squeeze(0)
top_ids: List[int] = sims.topk(top_k).indices.tolist()
inv_vocab = getattr(tokenizer, "inv_vocab", {})
return ", ".join(inv_vocab.get(tid, f"#{tid}") for tid in top_ids)
-5
View File
@@ -1,5 +0,0 @@
from lark import Lark
from .grammar import grammar
PromptParser = Lark(grammar, start="start", parser="earley")
-38
View File
@@ -1,38 +0,0 @@
grammar = r"""
?start: stmt+
?stmt: assign
| item
assign: "$" NAME "=" arg ";"
item: embedding
| WORD
| generic_function
| QUOTED_STRING
| weighted
| ref
generic_function: FUNC_NAME "(" arg ("|" arg)* ")"
weighted: "(" arg ":" SIGNED_NUMBER ")"
ref: "$" NAME
arg: item+
embedding: "embedding:" WORD
// NAME and WORD overlap on bare identifiers; the earley parser's dynamic lexer
// disambiguates by grammar context (the leading "$" forces NAME). This breaks
// under a basic/contextual lexer, so keep parser="earley" in __init__.py.
FUNC_NAME: /[A-Za-z_-]+/
NAME: /[A-Za-z_][A-Za-z0-9_]*/
WORD: /[A-Za-z0-9,_\.-]+/
QUOTED_STRING: /"([^"\\]*(\\.[^"\\]*)*)"|'([^'\\]*(\\.[^'\\]*)*)'/
SIGNED_NUMBER: /-?\d+(\.\d+)?/
COMMENT: /#[^\n]*/
%import common.WS
%ignore WS
%ignore COMMENT
"""
-30
View File
@@ -1,30 +0,0 @@
from typing import List, Union
import torch
from torch import Tensor
from torch.nn import Embedding
class PromptSegment:
"""A run of contiguous tokens from the user's prompt, possibly with inline TI tensor entries."""
def __init__(self, text: str, tokens: List[Union[int, Tensor]]):
self.text = text
self.tokens = tokens
def __repr__(self) -> str:
cleaned = ", ".join(str(t) if isinstance(t, int) else "EMBD" for t in self.tokens)
return f'"{self.text}"({cleaned})'
def token_length(self) -> int:
return len(self.tokens)
def get_embeddings(self, embedding_module: Embedding) -> Tensor:
"""Look up embeddings for plain int tokens.
Inline TI tensors aren't handled here; the encoder splices them in at a higher level.
"""
ids = torch.LongTensor([t for t in self.tokens if isinstance(t, int)]).to(
embedding_module.weight.device
)
return embedding_module(ids.unsqueeze(0))
-18
View File
@@ -1,18 +0,0 @@
from typing import Dict, Type
from ..actions.base import Action
action_registry: Dict[str, Type[Action]] = {}
def register_action(action: Type[Action]) -> None:
name = str(action.action_name)
if name in action_registry:
raise ValueError(f"Action {name} already registered")
action_registry[name] = action
def get_action_by_name(name: str) -> Type[Action]:
if name not in action_registry:
raise ValueError(f"Action {name} not found in registry")
return action_registry[name]
-81
View File
@@ -1,81 +0,0 @@
from typing import List
from lark import Token, Transformer
from comfy.sd1_clip import SDTokenizer
from ..actions.action_utils import parse_numeric_arg
from ..actions.base import Action, ActionArity
from ..actions.weighted import WeightedGroup
from .prompt_segment import PromptSegment
from .registration import get_action_by_name
from .utils import build_prompt_segment
class PromptTransformer(Transformer):
"""Maps the Lark parse tree into a flat list of PromptSegments and Actions."""
def __init__(self, tokenizer: SDTokenizer):
super().__init__()
self.tokenizer = tokenizer
self.vars: dict = {}
def assign(self, items):
name = str(items[0])
if name in self.vars:
raise ValueError(f"Variable ${name} is already defined")
self.vars[name] = items[1]
return None
def ref(self, items):
name = str(items[0])
if name not in self.vars:
raise ValueError(f"Variable ${name} referenced before assignment")
# Weight 1.0 makes the group transparent: _flatten and embedding_tensor
# already recurse through WeightedGroup, so no new container type needed.
return WeightedGroup(self.vars[name], weight=1.0)
def item(self, items: List[Token]):
for item in items:
if isinstance(item, (Action, PromptSegment, WeightedGroup)):
return item
if item.type == "WORD":
return build_prompt_segment(str(item), self.tokenizer)
if item.type == "QUOTED_STRING":
# Strip surrounding quotes, unescape \" and \'.
unquoted = item[1:-1]
unescaped = unquoted.replace('\\"', '"').replace("\\'", "'")
return build_prompt_segment(unescaped, self.tokenizer)
raise ValueError(f"Unknown item type: {item.type}")
def arg(self, items):
return items
def weighted(self, items):
arg_items, weight_token = items
return WeightedGroup(arg_items, float(weight_token))
def embedding(self, items):
return build_prompt_segment(
f"{self.tokenizer.embedding_identifier}{items[0]}",
self.tokenizer,
)
def generic_function(self, items):
# `emph(text|w)` is sugar for `(text:w)`; handled here so it doesn't need
# to fit the Action ABC (it changes weights, not embeddings).
if str(items[0]) == "emph":
if len(items) != 3:
raise ValueError("emph expects exactly two arguments: emph(text|weight)")
weight = parse_numeric_arg(items[2], action_name="emph", role="weight", cast=float)
return WeightedGroup(items[1], weight)
action = get_action_by_name(items[0])
if action.arity == ActionArity.SINGLE:
if len(items) != 2:
raise ValueError(f"Action {action.action_name} expects exactly one argument")
return action(items[1])
if action.arity == ActionArity.MULTI:
return action(items[1:])
raise ValueError(f"Unknown action arity: {action.arity}")
-41
View File
@@ -1,41 +0,0 @@
from lark import Token
from comfy.sd1_clip import SDTokenizer
from .prompt_segment import PromptSegment
def flatten_tree(tree):
if isinstance(tree, Token):
return [str(tree)]
return [str(tree.data)] + sum([flatten_tree(child) for child in tree.children], [])
def build_prompt_segment(text: str, tokenizer: SDTokenizer) -> PromptSegment:
"""Tokenize a chunk of plain text into a PromptSegment, expanding `embedding:NAME` refs to tensors."""
tokens = []
for word in text.split(" "):
if word.startswith(tokenizer.embedding_identifier) and tokenizer.embedding_directory is not None:
embedding_name = word[len(tokenizer.embedding_identifier):].strip("\n")
embedding, leftover = tokenizer._try_get_embedding(embedding_name)
if embedding is None:
print(f"warning, embedding:{embedding_name} does not exist, ignoring")
elif embedding.shape[1] != tokenizer.embedding_size:
print(
f"warning, embedding:{embedding_name} has size {embedding.shape[1]}, "
f"expected {tokenizer.embedding_size}, ignoring"
)
else:
if len(embedding.shape) == 1:
tokens.append(embedding)
else:
tokens.extend(embedding)
if leftover != "":
word = leftover
else:
continue
# Strip the SOT/EOT bracketing tokens added by the underlying CLIP tokenizer.
tokens.extend(tokenizer.tokenizer(word)["input_ids"][1:-1])
return PromptSegment(text, tokens)
+161 -105
View File
@@ -1,122 +1,178 @@
"""DSL-aware tokenizers.
import re
from typing import Union
Override `tokenize_with_weights` to parse our DSL and emit ComfyUI's native
`List[List[(token, weight)]]` format, where `token` is an int id, an inline
TI tensor, a lazily-evaluated `Action`, or `ACTION_CONTINUATION`.
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.ClipStuff.lib.actions import (
NudgeAction,
ArithAction,
ALL_START_CHARS,
ALL_END_CHARS,
ALL_ACTIONS,
)
from custom_nodes.ClipStuff.lib.actions.base import (
Action,
PromptSegment,
build_prompt_segment,
)
from custom_nodes.ClipStuff.lib.actions.lib import (
is_any_action_segment,
is_action_segment,
)
from custom_nodes.ClipStuff.lib.actions.utils import batch_size_info
Row alignment matters: comfy's stock `encode_token_weights` indexes weights by
post-transformer position, so each row must be exactly `max_length` entries.
A multi-slot Action is therefore emitted as one `(action, w)` entry followed by
`(ACTION_CONTINUATION, w)` placeholders; `process_tokens` drops the placeholders
and the action's tensor expands to fill those slots.
"""
arith_action = r'(<[a-zA-Z0-9\-_]+:[a-zA-Z0-9\-_]+>)'
from typing import Dict, Iterable, List, Tuple, Union
from lark import Tree
from comfy.sd1_clip import SD1Tokenizer, SDTokenizer
from .actions.base import ACTION_CONTINUATION, Action
from .actions.weighted import WeightedGroup
from .parser import PromptParser
from .parser.prompt_segment import PromptSegment
from .parser.transformer import PromptTransformer
# Side-effect import: registers all built-in actions with the parser.
from . import actions # noqa: F401
TokenEntry = Tuple[Union[int, "Action", object], float]
# TODO: Get embedding identifier from tokenizer
tokenizer_regex = re.compile(
fr"""
\d+\.\d+ # Capture decimals
|
(?:(?!embedding:)[\w\s]|embedding:[a-zA-Z0-9_]+)+ # Capture sequences of characters, including "embedding:"
|
\d+ # Capture whole numbers
|
[:+-{re.escape("".join(ALL_START_CHARS))}{re.escape("".join(ALL_END_CHARS))}] # Capture special characters including start and end characters
""",
re.VERBOSE
)
def tokenize(text: str) -> list[str]:
# Captures:
# 1. Words
# 2. Numbers(1.0, 1)
# 3. Special characters(ALL_START_CHARS, ALL_END_CHARS, :, +, -)
tokens = re.findall(tokenizer_regex, text)
print(tokens)
return [token.strip() for token in tokens]
def _flatten(item, weight: float) -> Iterable[TokenEntry]:
"""Walk the parsed item tree, yielding one (token, weight) entry per output slot."""
if isinstance(item, WeightedGroup):
for sub in item.items:
yield from _flatten(sub, weight * item.weight)
elif isinstance(item, Action):
yield (item, weight)
for _ in range(item.token_length() - 1):
yield (ACTION_CONTINUATION, weight)
elif isinstance(item, PromptSegment):
for tok in item.tokens:
yield (tok, weight)
else:
raise TypeError(f"Unexpected parse item {item!r} ({type(item).__name__})")
def parse_segment(tokens: list[str], tokenizer: SD1Tokenizer) -> PromptSegment | Action:
print("Parse segment: Checking token: " + tokens[0])
for action in ALL_ACTIONS:
if tokens[0] == action.START_CHAR:
return action.parse_segment(tokens, ALL_START_CHARS, ALL_END_CHARS, parse_segment, tokenizer)
# If we get here, it's a text segment
return build_prompt_segment(tokens.pop(0), tokenizer)
def parse(tokens: list[str], tokenizer: SD1Tokenizer) -> list[PromptSegment | Action]:
parsed = []
while tokens:
if tokens[0] == '':
tokens.pop(0)
continue
print("Parse: Checking token: " + tokens[0])
if tokens[0] in ALL_START_CHARS:
parsed.append(parse_segment(tokens, tokenizer))
else:
parsed.append(build_prompt_segment(tokens.pop(0), tokenizer))
return parsed
class PromptLangSDTokenizer(SDTokenizer):
def tokenize_with_weights( # type: ignore[override]
self, text: str, return_word_ids: bool = False, **kwargs
) -> List[List[TokenEntry]]:
# SDXL passes a pre-parsed tree to avoid re-running Lark per sub-tokenizer.
tree = kwargs.pop("_parsed_tree", None) or PromptParser.parse(text)
return self._batch_from_tree(tree)
def parse_special_tokens(string) -> list[str]:
out = []
current = ""
def _batch_from_tree(self, tree) -> List[List[TokenEntry]]:
pad_token = self.end_token if self.pad_with_end else 0
parsed = PromptTransformer(self).transform(tree)
items = parsed.children if isinstance(parsed, Tree) else [parsed]
# assign stmts return None (they only populate the transformer's var table).
items = [i for i in items if i is not None]
batches: List[List[TokenEntry]] = []
current: List[TokenEntry] = [(self.start_token, 1.0)]
def close(row: List[TokenEntry]) -> None:
row.append((self.end_token, 1.0))
row.extend([(pad_token, 1.0)] * (self.max_length - len(row)))
batches.append(row)
for item in items:
entries = list(_flatten(item, 1.0))
if len(current) + len(entries) > self.max_length - 1:
close(current)
current = [(self.start_token, 1.0)]
current.extend(entries)
close(current)
return batches
for char in string:
if char in ALL_START_CHARS:
out += [current]
current = char
elif char in ALL_END_CHARS:
out += [current + char]
current = ""
else:
current += char
out += [current]
return out
class PromptLangSD1Tokenizer(SD1Tokenizer):
def __init__(self, embedding_directory=None, tokenizer_data=None, clip_name="l", tokenizer=PromptLangSDTokenizer):
super().__init__(
embedding_directory=embedding_directory,
tokenizer_data=tokenizer_data or {},
clip_name=clip_name,
tokenizer=tokenizer,
)
def parse_segment_actions(string, tokenizer: SD1Tokenizer) -> list[PromptSegment | NudgeAction | ArithAction]:
tokens = tokenize(string)
parsed = parse(tokens, tokenizer)
return parsed
class TokenDict:
def __init__(self,
token_id: int,
weight: float = None,
nudge_id=None, nudge_weight=None, nudge_start: int = None, nudge_end: int = None,
arith_ops: dict[str, list[str]] = None):
if weight is None:
self.weight = 1.0
else:
self.weight = weight
self.token_id = token_id
self.nudge_id = nudge_id
self.nudge_weight = nudge_weight
self.nudge_index_start = nudge_start
self.nudge_index_stop = nudge_end
self.arith_ops = arith_ops
class PromptLangSDXLClipGTokenizer(PromptLangSDTokenizer):
def __init__(self, tokenizer_path=None, embedding_directory=None, tokenizer_data=None):
super().__init__(
tokenizer_path=tokenizer_path,
pad_with_end=False,
embedding_directory=embedding_directory,
embedding_size=1280,
embedding_key="clip_g",
tokenizer_data=tokenizer_data or {},
)
class MyTokenizer(SD1Tokenizer):
def __init__(self, tokenizer_path=None, max_length=77, pad_with_end=True, embedding_directory=None, embedding_size=768, embedding_key='clip_l', special_tokens=None):
super().__init__(tokenizer_path, max_length, pad_with_end, embedding_directory, embedding_size, embedding_key)
"""
Doesn't actually tokenize...
Returns batches of segments and actions
:return: List of list(batches) of segments and actions
"""
def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs) -> list[list[PromptSegment | Action]]:
if self.pad_with_end:
pad_token = self.end_token
else:
pad_token = 0
class PromptLangSDXLTokenizer:
def __init__(self, embedding_directory=None, tokenizer_data=None) -> None:
td = tokenizer_data or {}
self.clip_l = PromptLangSDTokenizer(embedding_directory=embedding_directory, tokenizer_data=td)
self.clip_g = PromptLangSDXLClipGTokenizer(embedding_directory=embedding_directory, tokenizer_data=td)
parsed_actions = parse_segment_actions(text, self)
def tokenize_with_weights(self, text: str, return_word_ids: bool = False, **kwargs) -> Dict[str, List[List[TokenEntry]]]:
tree = PromptParser.parse(text)
return {
"g": self.clip_g.tokenize_with_weights(text, return_word_ids, _parsed_tree=tree, **kwargs),
"l": self.clip_l.tokenize_with_weights(text, return_word_ids, _parsed_tree=tree, **kwargs),
}
# nudge_start = kwargs.get("nudge_start")
# nudge_end = kwargs.get("nudge_end")
#
# if nudge_start is not None and nudge_end is not None:
# nudge_start = int(nudge_start)
# nudge_end = int(nudge_end)
#
# # tokenize words
for segment in parsed_actions:
if isinstance(segment, Action):
print(segment.depth_repr())
else:
print(segment.depth_repr())
def untokenize(self, token_weight_pair):
return self.clip_g.untokenize(token_weight_pair)
# reshape token array to CLIP input size
batched_segments = []
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token])]
# batched_segments.append(batch)
batch_size = 1
for segment in parsed_actions:
num_tokens = segment.token_length()
# determine if we're going to try and keep the tokens in a single batch
is_large = num_tokens >= self.max_word_length
def state_dict(self):
return {}
# If the segment is too large to fit in a single batch, pad the current batch and start a new one
if num_tokens + batch_size > self.max_length - 1:
remaining_length = self.max_length - batch_size - 1 # -1 for end token
# Pad batch
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length - 1))
batched_segments.append(batch)
# start new batch
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token]), segment]
batch_size = num_tokens + 1 # +1 for start token
continue
# If the segment is small enough to fit in the current batch, add it
batch.append(segment)
batch_size += num_tokens
# Pad the last batch
remaining_length = self.max_length - batch_size - 1 # -1 for end token
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length))
batched_segments.append(batch)
for batch in batched_segments:
batch_size_info(batch)
return batched_segments
+136 -162
View File
@@ -1,24 +1,23 @@
import os
from dataclasses import dataclass
from typing import Any, List, Tuple
import random
import numpy as np
from PIL import Image
import comfy.sd
import folder_paths
from comfy.supported_models_base import ClipTarget
import comfy.sd
import comfy.ops
from custom_nodes.ClipStuff.lib.clip_model import SD1FunClipModel
from .lib.clip_model import PromptLangSD1ClipModel, PromptLangSDXLClipModel
from .lib.inspect import inspect_prompt
from .lib.tokenizer import PromptLangSD1Tokenizer, PromptLangSDXLTokenizer
from custom_nodes.ClipStuff.lib.tokenizer import MyTokenizer
class EmptyClass:
pass
class SpecialClipLoader:
"""Wraps a loaded CLIP with our DSL-aware tokenizer + text encoder."""
@classmethod
def INPUT_TYPES(cls): # type: ignore[no-untyped-def]
def INPUT_TYPES(s):
return {
"required": {
"source_clip": ("CLIP",),
@@ -27,200 +26,175 @@ class SpecialClipLoader:
RETURN_TYPES = ("CLIP",)
FUNCTION = "load_clip"
OUTPUT_IS_LIST = (False,)
CATEGORY = "conditioning"
@staticmethod
def load_clip(source_clip: comfy.sd.CLIP) -> Tuple[comfy.sd.CLIP]:
is_sdxl = hasattr(source_clip.cond_stage_model, "clip_g") and hasattr(source_clip.cond_stage_model, "clip_l")
embedding_directory = source_clip.tokenizer.clip_l.embedding_directory
def load_clip(source_clip):
clip_target = EmptyClass()
clip_target.params = {}
clip_target.clip = SD1FunClipModel
clip_target.tokenizer = MyTokenizer
if is_sdxl:
target = ClipTarget(PromptLangSDXLTokenizer, PromptLangSDXLClipModel)
else:
target = ClipTarget(PromptLangSD1Tokenizer, PromptLangSD1ClipModel)
new_clip = comfy.sd.CLIP(target=target, embedding_directory=embedding_directory)
new_clip.cond_stage_model.load_state_dict(source_clip.cond_stage_model.state_dict())
new_clip.layer_idx = source_clip.layer_idx
return (new_clip,)
# TODO: Extract embedding directory from source_clip
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.embedding_directory)
comfy.sd.load_clip_weights(
clip.cond_stage_model, source_clip.cond_stage_model.state_dict()
)
return (clip,)
class PromptLangInspect:
"""Shows what a DSL prompt resolves to at the embedding layer: per-slot weight, L2 norm, nearest vocab."""
class KepAdvTextEncode:
@classmethod
def INPUT_TYPES(cls): # type: ignore[no-untyped-def]
def INPUT_TYPES(s):
return {
"required": {
"clip": ("CLIP",),
"text": ("STRING", {"multiline": True}),
"top_k": ("INT", {"default": 3, "min": 1, "max": 10}),
"clip": ("CLIP",),
"nudge_start": ("INT", {}),
"nudge_end": ("INT", {}),
"split_newlines": ("BOOL", {"default": True}),
}
}
RETURN_TYPES = ("STRING",)
FUNCTION = "inspect"
RETURN_TYPES = ("CONDITIONING",)
FUNCTION = "encode"
OUTPUT_IS_LIST = (True,)
CATEGORY = "conditioning"
OUTPUT_NODE = True
def inspect(self, clip, text: str, top_k: int):
report = inspect_prompt(clip, text, top_k=top_k)
return {"ui": {"text": [report]}, "result": (report,)}
@staticmethod
def encode(clip, text, nudge_start, nudge_end, split_newlines):
ret = []
if split_newlines:
prompts = text.split("\n")
else:
prompts = [text]
for prompt in prompts:
if prompt.strip() == "":
continue
tokens = clip.tokenizer.tokenize_with_weights(
text,
return_word_ids=False,
nudge_start=nudge_start,
nudge_end=nudge_end,
)
cond, pooled = clip.encode_from_tokens(
tokens, return_pooled=True, position_ids=[0] * 77
)
cond = [[cond, {"pooled_output": pooled}]]
ret.append(cond)
return (ret,)
def tensor2img(tensor_img) -> Image.Image:
arr = (255.0 * tensor_img.cpu().numpy()).clip(0, 255).astype(np.uint8)
return Image.fromarray(arr)
def tensor2img(tensor_img):
i = 255.0 * tensor_img.cpu().numpy()
i_np_arr = np.clip(i, 0, 255, out=i).astype(np.uint8, copy=False)
return Image.fromarray(i_np_arr)
class BuildGif:
"""Builds an animated webp from a list of image batches.
Two output modes:
- "Big Grid": tiles batches across the X axis and chunks across the Y axis,
producing a single animated webp where each frame is the next image in a chunk.
- "One Per Split": one animation per (split, batch_index) combination.
"""
def __init__(self) -> None:
self.output_dir = folder_paths.get_output_directory()
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls): # type: ignore[no-untyped-def]
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
"split_every": ("INT", {"default": -1}),
"frame_duration": ("INT", {"default": 125}),
"output_mode": (["One Per Split", "Big Grid"], {"default": "Big Grid"}),
"output_mode": (
["One Per Split", "Big Grid"],
{"default": "Big Grid"},
),
}
}
RETURN_TYPES = ()
RELOAD_INST = True
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("Gifs",)
INPUT_IS_LIST = True
FUNCTION = "build_gif"
OUTPUT_NODE = True
OUTPUT_IS_LIST = (True,)
# OUTPUT_NODE = False
CATEGORY = "List Stuff"
def build_gif(
self,
images: List[Any],
split_every: List[int],
frame_duration: List[int],
output_mode: List[str],
):
@staticmethod
def build_gif(images: list, split_every: list[int], output_mode: str):
print("Build GIF called!")
print(f"{type(images)}")
if len(split_every) > 1:
raise ValueError("List input for split_every is not supported.")
if len(output_mode) > 1:
raise ValueError("List input for output_mode is not supported.")
if len(frame_duration) > 1:
raise ValueError("List input for frame_duration is not supported.")
mode = output_mode[0]
duration = frame_duration[0]
split_requested = split_every[0]
full_output_folder, filename, counter, subfolder, _ = folder_paths.get_save_image_path(
filename_prefix="Gif", output_dir=self.output_dir, image_width=0, image_height=0
)
raise Exception("List input for split every is not supported.")
split_every = split_every[0]
batch_size = images[0].size()[0]
# split_every=-1 means "don't split": one chunk containing everything.
if split_requested == -1:
if split_every == -1:
split_chunks = 1
chunk_len = len(images)
split_every = len(images)
else:
chunk_len = split_requested
split_chunks = len(images) // chunk_len
split_chunks = int(len(images) / split_every)
out = []
num_wide = batch_size
num_tall = split_chunks
chunked_batches = [
images[chunk_len * i : chunk_len * (i + 1)]
for i in range(split_chunks)
images[split_every * chunk_idx : split_every * (chunk_idx + 1)]
for chunk_idx in range(split_chunks)
]
results = []
ctx = _SaveContext(
images=images,
chunked_batches=chunked_batches,
chunk_len=chunk_len,
batch_size=batch_size,
split_chunks=split_chunks,
full_output_folder=full_output_folder,
filename=filename,
counter=counter,
subfolder=subfolder,
duration=duration,
)
if mode == "Big Grid":
results.append(self._save_big_grid(ctx))
elif mode == "One Per Split":
results.extend(self._save_one_per_split(ctx))
return {"ui": {"images": results}}
def _save_big_grid(self, ctx):
img_shape = ctx.images[0][0].shape
frames = []
for idx_in_chunk in range(ctx.chunk_len):
img_frame = Image.new(
"RGB", size=(ctx.batch_size * img_shape[0], ctx.split_chunks * img_shape[1])
)
for split_idx in range(ctx.split_chunks):
for batch_idx, img_tensor in enumerate(ctx.chunked_batches[split_idx][idx_in_chunk]):
img_frame.paste(
tensor2img(img_tensor),
(batch_idx * img_shape[0], split_idx * img_shape[1]),
)
frames.append(img_frame)
file = f"{ctx.filename}_{ctx.counter:05}_"
save_path = os.path.join(ctx.full_output_folder, file)
frames[0].save(
f"{save_path}.webp",
lossless=True,
save_all=True,
append_images=frames[1:],
optimize=False,
duration=ctx.duration,
loop=0,
)
return {"filename": f"{file}.webp", "subfolder": ctx.subfolder, "type": "output"}
def _save_one_per_split(self, ctx):
results = []
counter = ctx.counter
for split_idx in range(ctx.split_chunks):
split_start = ctx.chunk_len * split_idx
split_end = ctx.chunk_len * (split_idx + 1)
for batch_idx in range(ctx.batch_size):
file = f"{ctx.filename}_{counter:05}_"
save_path = os.path.join(ctx.full_output_folder, file)
counter += 1
tensor2img(ctx.images[split_start][batch_idx]).save(
f"{save_path}.webp",
save_all=True,
append_images=[
tensor2img(nested[batch_idx])
for nested in ctx.images[split_start + 1 : split_end]
],
optimize=False,
duration=ctx.duration,
loop=0,
if output_mode == "Big Grid":
# For every image in gif
for idx_in_chunk in range(split_every):
img_shape = images[0][0].shape
img_frame = Image.new(
"RGB", size=(num_wide * img_shape[0], num_tall * img_shape[1])
)
results.append({"filename": f"{file}.webp", "subfolder": ctx.subfolder, "type": "output"})
return results
# For every chunk of images
for split_idx in range(split_chunks):
img_chunk = chunked_batches[split_idx]
for batch_idx, img_tensor in enumerate(img_chunk[idx_in_chunk]):
img = tensor2img(img_tensor)
img_frame.paste(
img, (batch_idx * img_shape[0], split_idx * img_shape[1])
)
frames.append(img_frame)
@dataclass
class _SaveContext:
images: Any
chunked_batches: Any
chunk_len: int
batch_size: int
split_chunks: int
full_output_folder: str
filename: str
counter: int
subfolder: str
duration: int
save_path = (
f"{folder_paths.get_output_directory()}/{random.randint(1, 100)}"
)
frames[0].save(
f"{save_path}.webp",
# quality=100,
# method=6,
lossless=True,
save_all=True,
append_images=frames[1:],
optimize=False,
duration=125,
loop=0,
)
elif output_mode == "One Per Split":
for split_idx in range(int(split_chunks)):
split_start = split_every * split_idx
split_end = split_every * (split_idx + 1)
for batch_idx in range(batch_size):
save_path = f"{folder_paths.get_output_directory()}/-{batch_idx}-{random.randint(1, 100)}"
print(save_path)
tensor2img(images[split_start][batch_idx]).save(
f"{save_path}.webp",
save_all=True,
append_images=[
tensor2img(nested_batch[batch_idx])
for nested_batch in images[split_start + 1 : split_end]
],
optimize=False,
duration=125,
loop=0,
)
return (out,)
-26
View File
@@ -1,26 +0,0 @@
[project]
name = "keppromptlang"
version = "0.2.0"
description = "A small DSL for ComfyUI that lets you do math on CLIP token embeddings before they're fed into the text transformer."
readme = "README.md"
license = { text = "MIT" }
requires-python = ">=3.10"
dependencies = [
"lark",
]
[project.optional-dependencies]
dev = [
"pytest",
"torch",
]
[project.urls]
Repository = "https://github.com/M1kep/KepPromptLang"
[tool.comfy]
PublisherId = "m1kep"
DisplayName = "KepPromptLang"
[tool.pytest.ini_options]
testpaths = ["tests"]
-1
View File
@@ -1 +0,0 @@
lark
-105
View File
@@ -1,105 +0,0 @@
"""Test setup that runs before any tests are collected.
Two things make this tricky:
1. The project's runtime imports use ComfyUI (`comfy.sd1_clip`), which we don't want to require for unit tests.
2. The package is normally installed under `custom_nodes/KepPromptLang/`, so we register `KepPromptLang` as a package alias.
"""
import os
import sys
import types
import pytest
REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, os.path.dirname(REPO_ROOT))
def _install_runtime_stubs():
"""Stub out runtime deps (numpy/PIL/comfy/folder_paths) so test imports of the package work.
Tests don't exercise the ComfyUI nodes; they only need the parser and action math layers.
"""
# Only stub modules that aren't actually installed; real numpy/torch must take precedence.
if "PIL" not in sys.modules:
try:
import PIL # noqa: F401
except ImportError:
pil = types.ModuleType("PIL")
pil.Image = types.ModuleType("PIL.Image")
sys.modules["PIL"] = pil
sys.modules["PIL.Image"] = pil.Image
if "numpy" not in sys.modules:
try:
import numpy # noqa: F401
except ImportError:
sys.modules["numpy"] = types.ModuleType("numpy")
if "folder_paths" not in sys.modules:
sys.modules["folder_paths"] = types.ModuleType("folder_paths")
if "comfy" in sys.modules:
return
comfy = types.ModuleType("comfy")
comfy_sd = types.ModuleType("comfy.sd")
comfy_sd.CLIP = type("CLIP", (), {})
comfy_supported = types.ModuleType("comfy.supported_models_base")
comfy_supported.ClipTarget = type("ClipTarget", (), {})
sdxl_clip = types.ModuleType("comfy.sdxl_clip")
sdxl_clip.SDXLClipModel = type("SDXLClipModel", (), {})
sdxl_clip.SDXLClipG = type("SDXLClipG", (), {})
sd1_clip = types.ModuleType("comfy.sd1_clip")
class SDTokenizer: # minimal stand-in
embedding_identifier = "embedding:"
def __init__(self, *args, **kwargs):
self.embedding_directory = None
self.embedding_size = 768
self.start_token = 49406
self.end_token = 49407
self.pad_with_end = True
self.max_length = 77
self.tokenizer = _FakeTokenizer()
def _try_get_embedding(self, name):
return None, ""
class SD1Tokenizer:
def __init__(self, *args, **kwargs):
pass
sd1_clip.SDTokenizer = SDTokenizer
sd1_clip.SD1Tokenizer = SD1Tokenizer
sd1_clip.SDClipModel = type("SDClipModel", (), {})
sd1_clip.SD1ClipModel = type("SD1ClipModel", (), {})
comfy.sd1_clip = sd1_clip
comfy.sd = comfy_sd
comfy.sdxl_clip = sdxl_clip
comfy.supported_models_base = comfy_supported
sys.modules["comfy"] = comfy
sys.modules["comfy.sd1_clip"] = sd1_clip
sys.modules["comfy.sdxl_clip"] = sdxl_clip
sys.modules["comfy.sd"] = comfy_sd
sys.modules["comfy.supported_models_base"] = comfy_supported
class _FakeTokenizer:
"""Tokenize each whitespace-separated word into a single deterministic int id."""
def __call__(self, word):
# Deterministic: sum of character codepoints, modulo a small range. SOT/EOT bracketing.
token = (sum(ord(c) for c in word) % 49000) + 100
return {"input_ids": [49406, token, 49407]}
_install_runtime_stubs()
@pytest.fixture
def tokenizer():
from comfy.sd1_clip import SDTokenizer
return SDTokenizer()
-160
View File
@@ -1,160 +0,0 @@
"""Action-level tests using a tiny in-memory torch.nn.Embedding.
These verify the math/shape contracts of each action without needing ComfyUI or a real CLIP.
"""
import pytest
torch = pytest.importorskip("torch")
from KepPromptLang.lib.actions.avg import AverageAction
from KepPromptLang.lib.actions.diff import DiffAction
from KepPromptLang.lib.actions.mult import MultiplyAction
from KepPromptLang.lib.actions.neg import NegAction
from KepPromptLang.lib.actions.norm import NormAction
from KepPromptLang.lib.actions.pos_scale import PosScaleAction
from KepPromptLang.lib.actions.post_pos import PostPosAction
from KepPromptLang.lib.actions.rand import RandAction
from KepPromptLang.lib.actions.scale_dims import ScaleDims
from KepPromptLang.lib.actions.set_dims import SetDims
from KepPromptLang.lib.actions.slerp import SlerpAction
from KepPromptLang.lib.actions.sum import SumAction
from KepPromptLang.lib.actions.utils import slerp
from KepPromptLang.lib.parser.prompt_segment import PromptSegment
EMBED_DIM = 4
VOCAB = 100
@pytest.fixture
def embedding():
torch.manual_seed(0)
emb = torch.nn.Embedding(VOCAB, EMBED_DIM)
return emb
def seg(*token_ids):
return PromptSegment(text="x", tokens=list(token_ids))
def test_sum_adds_embeddings(embedding):
a = seg(1, 2)
b = seg(3, 4)
action = SumAction([[a], [b]])
expected = embedding(torch.LongTensor([[1, 2]])) + embedding(torch.LongTensor([[3, 4]]))
assert torch.allclose(action.get_result(embedding), expected)
def test_diff_subtracts_embeddings(embedding):
a = seg(1, 2)
b = seg(3, 4)
action = DiffAction([[a], [b]])
expected = embedding(torch.LongTensor([[1, 2]])) - embedding(torch.LongTensor([[3, 4]]))
assert torch.allclose(action.get_result(embedding), expected)
def test_neg_negates(embedding):
action = NegAction([seg(1, 2)])
expected = -embedding(torch.LongTensor([[1, 2]]))
assert torch.allclose(action.get_result(embedding), expected)
def test_mult_scales(embedding):
# The multiplier is read from PromptSegment.text (mimicking parser output).
target_seg = PromptSegment(text="3.5", tokens=[1])
action = MultiplyAction([[seg(1, 2)], [target_seg]])
expected = embedding(torch.LongTensor([[1, 2]])) * 3.5
assert torch.allclose(action.get_result(embedding), expected)
def test_norm_unit_length(embedding):
action = NormAction([seg(1, 2)])
result = action.get_result(embedding)
norms = torch.norm(result, dim=-1)
assert torch.allclose(norms, torch.ones_like(norms), atol=1e-5)
def test_avg_weighted_mix(embedding):
weight_seg = PromptSegment(text="0.25", tokens=[1])
action = AverageAction([[seg(1, 2)], [seg(3, 4)], [weight_seg]])
expected = embedding(torch.LongTensor([[1, 2]])) * 0.75 + embedding(torch.LongTensor([[3, 4]])) * 0.25
assert torch.allclose(action.get_result(embedding), expected)
def test_avg_mismatched_lengths_errors():
weight_seg = PromptSegment(text="0.5", tokens=[1])
with pytest.raises(ValueError, match="same length"):
AverageAction([[seg(1, 2)], [seg(3)], [weight_seg]])
def test_slerp_endpoints(embedding):
weight0 = PromptSegment(text="0.0", tokens=[1])
weight1 = PromptSegment(text="1.0", tokens=[1])
a, b = seg(1, 2), seg(3, 4)
a_emb = embedding(torch.LongTensor([[1, 2]]))
b_emb = embedding(torch.LongTensor([[3, 4]]))
assert torch.allclose(SlerpAction([[a], [b], [weight0]]).get_result(embedding), a_emb, atol=1e-5)
assert torch.allclose(SlerpAction([[a], [b], [weight1]]).get_result(embedding), b_emb, atol=1e-5)
def test_slerp_helper_endpoints_and_midpoint():
low = torch.tensor([1.0, 0.0])
high = torch.tensor([0.0, 1.0]) # 90 degrees apart, both unit length
assert torch.allclose(slerp(0.0, low, high), low, atol=1e-5)
assert torch.allclose(slerp(1.0, low, high), high, atol=1e-5)
midpoint = slerp(0.5, low, high)
# Midpoint of orthogonal unit vectors on the unit sphere is (sqrt(2)/2, sqrt(2)/2).
expected = torch.tensor([2 ** 0.5 / 2, 2 ** 0.5 / 2])
assert torch.allclose(midpoint, expected, atol=1e-5)
def test_rand_token_length_and_bounds():
length_seg = PromptSegment(text="3", tokens=[1])
min_seg = PromptSegment(text="-2", tokens=[1])
max_seg = PromptSegment(text="2", tokens=[1])
action = RandAction([[length_seg], [min_seg], [max_seg]])
emb = torch.nn.Embedding(VOCAB, EMBED_DIM)
result = action.get_result(emb)
assert result.shape == (1, 3, EMBED_DIM)
assert (result >= -2).all() and (result <= 2).all()
def test_scale_dims_modifies_only_target_dim(embedding):
pair_seg = PromptSegment(text="0,3.0", tokens=[1])
action = ScaleDims([[seg(1, 2)], [pair_seg]])
base = embedding(torch.LongTensor([[1, 2]])).clone()
result = action.get_result(embedding)
assert torch.allclose(result[0, :, 0], base[0, :, 0] * 3.0)
assert torch.allclose(result[0, :, 1:], base[0, :, 1:])
def test_set_dims_overwrites_value(embedding):
pair_seg = PromptSegment(text="2,-9.5", tokens=[1])
action = SetDims([[seg(1, 2)], [pair_seg]])
result = action.get_result(embedding)
assert torch.allclose(result[0, :, 2], torch.tensor([-9.5, -9.5]))
def test_pos_scale_returns_modifier(embedding):
multiplier = PromptSegment(text="1.5", tokens=[1])
action = PosScaleAction([[seg(1, 2)], [multiplier]])
tensor, modifiers = action.get_result(embedding)
assert tensor.shape == (1, 2, EMBED_DIM)
assert modifiers.position_embed_scale == 1.5
def test_post_pos_returns_bypass(embedding):
action = PostPosAction([seg(1, 2)])
tensor, modifiers = action.get_result(embedding)
assert tensor.shape == (1, 2, EMBED_DIM)
assert modifiers.bypass_pos_embed is True
def test_action_token_lengths():
a, b = seg(1, 2, 3), seg(4, 5, 6)
assert SumAction([[a], [b]]).token_length() == 3
assert DiffAction([[a], [b]]).token_length() == 3
assert NegAction([a]).token_length() == 3
assert NormAction([a]).token_length() == 3
-96
View File
@@ -1,96 +0,0 @@
import pytest
torch = pytest.importorskip("torch")
from KepPromptLang.lib.actions.nearest import NearestAction
from KepPromptLang.lib.actions.noise import NoiseAction
from KepPromptLang.lib.actions.project import ProjectAction, RejectAction
from KepPromptLang.lib.actions.renorm import RenormAction
from KepPromptLang.lib.parser.prompt_segment import PromptSegment
from KepPromptLang.lib.parser.registration import get_action_by_name
EMBED_DIM = 4
VOCAB = 50
@pytest.fixture
def embedding():
torch.manual_seed(0)
return torch.nn.Embedding(VOCAB, EMBED_DIM)
def seg(*token_ids):
return PromptSegment(text="x", tokens=list(token_ids))
def test_proj_plus_reject_reconstructs_input(embedding):
a, b = seg(1, 2), seg(3)
proj = ProjectAction([[a], [b]]).get_result(embedding)
rej = RejectAction([[a], [b]]).get_result(embedding)
a_emb = embedding(torch.LongTensor([[1, 2]]))
assert torch.allclose(proj + rej, a_emb, atol=1e-5)
def test_reject_is_orthogonal_to_b(embedding):
a, b = seg(1, 2), seg(3)
rej = RejectAction([[a], [b]]).get_result(embedding)
b_dir = torch.nn.functional.normalize(
embedding(torch.LongTensor([[3]])).mean(dim=1, keepdim=True), dim=-1
)
dots = (rej * b_dir).sum(dim=-1)
assert torch.allclose(dots, torch.zeros_like(dots), atol=1e-5)
def test_renorm_matches_ref_norm(embedding):
a, ref = seg(1, 2), seg(3)
out = RenormAction([[a], [ref]]).get_result(embedding)
ref_norm = torch.norm(embedding(torch.LongTensor([[3]])), dim=-1).mean()
out_norms = torch.norm(out, dim=-1)
assert torch.allclose(out_norms, ref_norm.expand_as(out_norms), atol=1e-5)
def test_noise_shape_and_mean(embedding):
std = PromptSegment(text="0.01", tokens=[1])
out = NoiseAction([[seg(1, 2)], [std]]).get_result(embedding)
base = embedding(torch.LongTensor([[1, 2]]))
assert out.shape == base.shape
# Perturbation magnitude bounded (5σ with margin); std=0.01, EMBED_DIM=4.
assert (out - base).abs().max() < 0.2
def test_noise_zero_std_is_identity(embedding):
std = PromptSegment(text="0.0", tokens=[1])
out = NoiseAction([[seg(1, 2)], [std]]).get_result(embedding)
base = embedding(torch.LongTensor([[1, 2]]))
assert torch.allclose(out, base)
def test_nearest_returns_exact_token_for_that_token(embedding):
out = NearestAction([[seg(7)]]).get_result(embedding)
assert out.shape == (1, 1, EMBED_DIM)
assert torch.allclose(out[0, 0], embedding.weight[7])
def test_nearest_k_tokens(embedding):
k = PromptSegment(text="3", tokens=[1])
action = NearestAction([[seg(7)], [k]])
assert action.token_length() == 3
out = action.get_result(embedding)
assert out.shape == (1, 3, EMBED_DIM)
# First match should be the token itself.
assert torch.allclose(out[0, 0], embedding.weight[7])
def test_lerp_is_registered_as_avg_alias():
lerp_cls = get_action_by_name("lerp")
avg_cls = get_action_by_name("avg")
assert issubclass(lerp_cls, avg_cls)
def test_token_lengths():
a, b = seg(1, 2, 3), seg(4)
assert ProjectAction([[a], [b]]).token_length() == 3
assert RejectAction([[a], [b]]).token_length() == 3
assert RenormAction([[a], [b]]).token_length() == 3
std = PromptSegment(text="0.1", tokens=[1])
assert NoiseAction([[a], [std]]).token_length() == 3
-58
View File
@@ -1,58 +0,0 @@
from KepPromptLang.lib.actions.diff import DiffAction
from KepPromptLang.lib.actions.norm import NormAction
from KepPromptLang.lib.actions.sum import SumAction
from KepPromptLang.lib.parser import PromptParser
from KepPromptLang.lib.parser.prompt_segment import PromptSegment
from KepPromptLang.lib.parser.transformer import PromptTransformer
def parse(text, tokenizer):
tree = PromptParser.parse(text)
return PromptTransformer(tokenizer).transform(tree)
def test_plain_words_become_segments(tokenizer):
result = parse("hello world", tokenizer)
items = result.children
assert len(items) == 2
assert all(isinstance(i, PromptSegment) for i in items)
assert items[0].text == "hello"
assert items[1].text == "world"
def test_sum_action_parses(tokenizer):
action = parse("sum(king|man|woman)", tokenizer)
items = action.children if hasattr(action, "children") else [action]
assert len(items) == 1
assert isinstance(items[0], SumAction)
assert len(items[0].all_args) == 3
def test_nested_actions(tokenizer):
action = parse("sum(diff(king|man)|woman)", tokenizer)
items = action.children if hasattr(action, "children") else [action]
outer = items[0]
assert isinstance(outer, SumAction)
inner = outer.all_args[0][0]
assert isinstance(inner, DiffAction)
def test_norm_single_arg(tokenizer):
action = parse("norm(cat)", tokenizer)
items = action.children if hasattr(action, "children") else [action]
assert isinstance(items[0], NormAction)
def test_quoted_string(tokenizer):
result = parse('"hello world"', tokenizer)
items = result.children if hasattr(result, "children") else [result]
assert isinstance(items[0], PromptSegment)
assert items[0].text == "hello world"
def test_unknown_action_errors(tokenizer):
import pytest
from lark.exceptions import VisitError
with pytest.raises((ValueError, VisitError), match="not found in registry"):
parse("nonexistentAction(cat)", tokenizer)
-93
View File
@@ -1,93 +0,0 @@
"""Verify the tokenizer emits ComfyUI's native (token, weight) format with lazy Actions
and per-position weights.
Uses the comfy stub from conftest, so no real ComfyUI needed.
"""
import pytest
torch = pytest.importorskip("torch")
from KepPromptLang.lib.actions.base import ACTION_CONTINUATION, Action
from KepPromptLang.lib.actions.sum import SumAction
from KepPromptLang.lib.actions.weighted import WeightedGroup
from KepPromptLang.lib.tokenizer import PromptLangSDTokenizer
@pytest.fixture
def tok():
return PromptLangSDTokenizer()
def test_plain_text_is_int_tuples_at_max_length(tok):
[row] = tok.tokenize_with_weights("hello world")
assert all(isinstance(t, int) and w == 1.0 for t, w in row)
assert row[0] == (tok.start_token, 1.0)
assert len(row) == tok.max_length
def test_action_emits_one_entry_plus_continuations(tok):
[row] = tok.tokenize_with_weights("a sum(king|man|woman) here")
assert len(row) == tok.max_length
actions = [t for t, _ in row if isinstance(t, Action)]
continuations = [t for t, _ in row if t is ACTION_CONTINUATION]
assert len(actions) == 1
assert isinstance(actions[0], SumAction)
assert len(continuations) == actions[0].token_length() - 1
def test_paren_weight_syntax(tok):
[row] = tok.tokenize_with_weights("a (cat:1.3) here")
weighted = [(t, w) for t, w in row if w != 1.0]
# "cat" is one token under the fake tokenizer.
assert len(weighted) == 1
assert weighted[0][1] == pytest.approx(1.3)
assert isinstance(weighted[0][0], int)
def test_paren_weight_on_action_propagates_to_continuations(tok):
[row] = tok.tokenize_with_weights("(sum(king|man|woman):0.7)")
action_entry = next((t, w) for t, w in row if isinstance(t, Action))
cont_weights = [w for t, w in row if t is ACTION_CONTINUATION]
assert action_entry[1] == pytest.approx(0.7)
assert all(w == pytest.approx(0.7) for w in cont_weights)
def test_nested_paren_weights_multiply(tok):
[row] = tok.tokenize_with_weights("((cat:1.2):0.5)")
weighted = [(t, w) for t, w in row if w != 1.0]
assert len(weighted) == 1
assert weighted[0][1] == pytest.approx(0.6)
def test_emph_is_alias_for_paren_weight(tok):
[row] = tok.tokenize_with_weights("emph(cat|1.3)")
weighted = [(t, w) for t, w in row if w != 1.0]
assert len(weighted) == 1
assert weighted[0][1] == pytest.approx(1.3)
def test_nested_actions_stay_nested(tok):
[row] = tok.tokenize_with_weights("sum(diff(king|man)|woman)")
actions = [t for t, _ in row if isinstance(t, Action)]
assert len(actions) == 1
assert isinstance(actions[0], SumAction)
from KepPromptLang.lib.actions.diff import DiffAction
assert isinstance(actions[0].all_args[0][0], DiffAction)
def test_overflow_splits_into_multiple_batches(tok):
text = " ".join(f"w{i}" for i in range(80))
batches = tok.tokenize_with_weights(text)
assert len(batches) >= 2
for row in batches:
assert len(row) == tok.max_length
assert row[0] == (tok.start_token, 1.0)
def test_weighted_group_token_length():
from KepPromptLang.lib.parser.prompt_segment import PromptSegment
grp = WeightedGroup([PromptSegment("a", [1, 2]), PromptSegment("b", [3])], 1.5)
assert grp.token_length() == 3
-77
View File
@@ -1,77 +0,0 @@
import pytest
torch = pytest.importorskip("torch")
from KepPromptLang.lib.actions.base import Action
from KepPromptLang.lib.actions.sum import SumAction
from KepPromptLang.lib.tokenizer import PromptLangSDTokenizer
@pytest.fixture
def tok():
return PromptLangSDTokenizer()
def content_tokens(row, tok):
"""Non-SOT/EOT/pad int tokens from a row, in order."""
return [
t for t, _ in row
if isinstance(t, int) and t not in (tok.start_token, tok.end_token, 0)
]
def test_var_substitutes_at_top_level(tok):
[direct] = tok.tokenize_with_weights("a cat dog")
[via_var] = tok.tokenize_with_weights("$x = cat dog; a $x")
assert content_tokens(via_var, tok) == content_tokens(direct, tok)
def test_var_holding_action(tok):
[row] = tok.tokenize_with_weights("$axis = sum(king|man); $axis")
actions = [t for t, _ in row if isinstance(t, Action)]
assert len(actions) == 1
assert isinstance(actions[0], SumAction)
def test_var_inside_function_arg(tok):
[row] = tok.tokenize_with_weights("$a = king; sum($a|woman)")
actions = [t for t, _ in row if isinstance(t, Action)]
assert len(actions) == 1
# token_length should be 1 (single-token base arg via the fake tokenizer)
assert actions[0].token_length() == 1
def test_var_under_weight(tok):
[row] = tok.tokenize_with_weights("$x = cat; ($x:1.5)")
weighted = [w for t, w in row if isinstance(t, int) and w != 1.0]
assert weighted == [pytest.approx(1.5)]
def test_var_ref_before_assign_errors(tok):
from lark.exceptions import VisitError
with pytest.raises((ValueError, VisitError), match="referenced before assignment"):
tok.tokenize_with_weights("$x and then $x = cat;")
def test_var_reassignment_errors(tok):
from lark.exceptions import VisitError
with pytest.raises((ValueError, VisitError), match="already defined"):
tok.tokenize_with_weights("$x = cat; $x = dog; $x")
def test_var_chains(tok):
[direct] = tok.tokenize_with_weights("cat")
[chained] = tok.tokenize_with_weights("$a = cat; $b = $a; $b")
assert content_tokens(chained, tok) == content_tokens(direct, tok)
def test_comments_ignored(tok):
[a] = tok.tokenize_with_weights("cat dog")
[b] = tok.tokenize_with_weights("cat # this is ignored\ndog")
assert content_tokens(a, tok) == content_tokens(b, tok)
def test_assign_only_produces_empty_prompt(tok):
[row] = tok.tokenize_with_weights("$x = cat;")
# SOT + EOT + padding only
assert content_tokens(row, tok) == []
-72
View File
@@ -1,72 +0,0 @@
"""Regenerate the action table in README.md.
Loads each action file by path so the docs can be regenerated without ComfyUI installed.
"""
import importlib
import inspect
import os
import sys
import types
from typing import List, Type
REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
ACTIONS_DIR = os.path.join(REPO_ROOT, "lib", "actions")
EXCLUDED = {"__init__.py", "base.py", "types.py", "action_utils.py", "utils.py"}
def _stub_runtime_deps():
"""Stub modules whose only purpose is to satisfy the package's top-level imports."""
sys.path.insert(0, os.path.dirname(REPO_ROOT))
# Avoid pulling in ComfyUI's nodes.py during action discovery.
pkg_init = sys.modules.get("KepPromptLang")
if pkg_init is None:
pkg = types.ModuleType("KepPromptLang")
pkg.__path__ = [REPO_ROOT]
sys.modules["KepPromptLang"] = pkg
def find_action_classes() -> List[Type]:
_stub_runtime_deps()
base_class = importlib.import_module("KepPromptLang.lib.actions.base").Action
found: List[Type] = []
for filename in sorted(os.listdir(ACTIONS_DIR)):
if not filename.endswith(".py") or filename in EXCLUDED:
continue
mod = importlib.import_module(f"KepPromptLang.lib.actions.{filename[:-3]}")
for _, cls in inspect.getmembers(mod, inspect.isclass):
if (
issubclass(cls, base_class)
and cls is not base_class
and cls.__module__ == mod.__name__
):
found.append(cls)
return found
def render_table(classes: List[Type]) -> str:
rows = []
for cls in sorted(classes, key=lambda c: c.action_name):
examples = "<ul>" + "".join(
f"<li>{ex.replace('|', chr(92) + '|')}</li>"
for ex in (cls.usage_examples or [])
) + "</ul>"
cells = [
(cls.display_name or "").replace("|", "\\|"),
(cls.action_name or "").replace("|", "\\|"),
(cls.description or "").replace("|", "\\|"),
examples,
]
rows.append("| " + " | ".join(cells) + " |")
return (
"| Display Name | Action Name | Description | Usage Examples |\n"
"| --- | --- | --- | --- |\n"
+ "\n".join(rows)
+ "\n"
)
if __name__ == "__main__":
print(render_table(find_action_classes()))