Author SHA1 Message Date
Michael Poutre 61f84e06ed feat(func/slerp): Add function 2023-08-31 22:31:07 -07:00
49 changed files with 1443 additions and 2112 deletions
+68 -87
View File
@@ -1,100 +1,81 @@
# KepPromptLang
# ClipStuff
## Basic Instructions.
Clone repo into custom_nodes folder.
Install the requirements.txt file via pip.
A small DSL for ComfyUI that lets you do math on CLIP token embeddings before they're fed into the text transformer.
Pass CLIP output from Load Checkpoint into SpecialClipLoader node, then use the outputted clip with standard Clip Text Encode.
```
sum(diff(king|man)|woman)
norm(sum(cat | dog | horse | parrot))
A slerp(cat|dog|0.5) is happy
```
See example workflow in examples folder.
## 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)`
### Example Photo
![Example Photo](assets/first_example.png)
## 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> |
## Syntax Elements
`lerp(a|b|t)` is also accepted as an alias for `avg(a|b|t)`.
1. **Embedding**:
- Syntax: `embedding:WORD`
- Example: `embedding:face_vector`
- Represents a named vector embedding(Textual Inversion).
Regenerate the table with `python tools/build_docs.py`.
2. **Word**:
- Syntax: Any alphanumeric word including characters such as `,`, `_`, and `-`.
- Example: `cat, dog_face, id_123`
- Represents simple words or identifiers.
## Development
3. **Quoted String**:
- Syntax: A string enclosed within double or single quotes. You can escape quotes inside the string using a backslash (`\`).
- Example: `"Hello World"`, `'It\'s a sunny day'`
- Represents string literals.
Tests are pytest-based and don't require ComfyUI:
## Functions
```bash
pip install -e ".[dev]"
python -m pytest
```
Here are the available functions and their usage:
## Compatibility
1. **Sum Function**:
- Syntax: `sum(arg1 | arg2 | ... | argN)`
- Adds together multiple embeddings.
- Example: `sum(embedding:face1 | dog)`
- 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.
2. **Negation Function**:
- Syntax: `neg(arg)`
- Negates the output.
- Example: `neg(A embedding:happycats outside)`
3. **Normalization Function**:
- Syntax: `norm(arg)`
- Normalizes the given vector embedding.
- Example: `norm(sum(embedding:face1 | embedding:face2))`
4. **Difference Function**:
- Syntax: `diff(arg1 | arg2 | ... | argN)`
- Computes the difference between multiple vector embeddings.
- Example: `diff(embedding:face1 | embedding:face2)`
### Notes on Arguments:
- Each function takes one or more arguments.
- An argument (`arg`) can be an embedding, a word, another function, or a quoted string.
- For functions that accept multiple arguments, they are separated by the `|` symbol.
## Examples
1. Add two embeddings and normalize the result:
```
norm(sum(cat | dog | horse | parrot))
```
2. Negate an embedding:
```
neg(embedding:body_vector)
```
3. King - Man + Woman = Queen:
```
sum(diff(king|man)|woman)
```
or
```
sum(king|neg(man)|woman)
```
```
+4 -10
View File
@@ -1,15 +1,9 @@
from .nodes import BuildGif, PromptLangInspect, SpecialClipLoader
from .nodes import (
BuildGif,
SpecialClipLoader,
)
NODE_CLASS_MAPPINGS = {
"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"]
+16
View File
@@ -0,0 +1,16 @@
from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction
from custom_nodes.KepPromptLang.lib.actions.mult import MultiplyAction
from custom_nodes.KepPromptLang.lib.actions.neg import NegAction
from custom_nodes.KepPromptLang.lib.actions.norm import NormAction
from custom_nodes.KepPromptLang.lib.actions.rand import RandAction
from custom_nodes.KepPromptLang.lib.actions.slerp import SlerpAction
from custom_nodes.KepPromptLang.lib.actions.sum import SumAction
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
register_action(DiffAction)
register_action(MultiplyAction)
register_action(NegAction)
register_action(NormAction)
register_action(RandAction)
register_action(SumAction)
register_action(SlerpAction)
View File
+117
View File
@@ -0,0 +1,117 @@
from abc import ABC, abstractmethod
from enum import Enum
from typing import Union, List
from torch import Tensor
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class ActionArity(Enum):
NONE = 0
SINGLE = 1
MULTI = 2
class Action(ABC):
@property
@abstractmethod
def chars(self) -> Union[List[str], None]:
pass
@property
@abstractmethod
def arity(self) -> ActionArity:
"""
Determines the arity of the action. This is used to determine how many arguments the action supports.
:return:
"""
pass
@property
@abstractmethod
def name(self) -> str:
pass
@property
@abstractmethod
def grammar(self) -> str:
"""
The grammar for this action. This is used to parse the action from the prompt.
:return:
"""
pass
@abstractmethod
def token_length(self) -> int:
"""
The length of the tokens that this action will add to the prompt.
:return:
"""
pass
@abstractmethod
def get_all_segments(self) -> List[PromptSegment]:
"""
Get all segments, including nested segments.
:return:
"""
pass
@abstractmethod
def __init__(self, *args, **kwargs) -> None:
"""
Initialize the action. This is called when the action is parsed from the prompt.
:param args: The arguments for the action.
"""
pass
@abstractmethod
def get_result(self, embedding_module: Embedding) -> Tensor:
"""
Get the result of this action. This is called when the embeddings are being calculated.
:param embedding_module: The embedding module to use to get the base embeddings for tokens.
:return:
"""
pass
def depth_repr(self, depth: int = 1) -> str:
raise NotImplementedError()
class SingleArgAction(Action, ABC):
arity = ActionArity.SINGLE
def get_all_segments(self) -> List[PromptSegment]:
segments = []
for seg_or_action in self.arg:
if isinstance(seg_or_action, Action):
segments.extend(seg_or_action.get_all_segments())
else:
segments.append(seg_or_action)
return segments
def __init__(self, arg: List[Union[PromptSegment, Action]]):
# TODO: Target is a list now... what does this mean for us..
self.arg = arg
def __repr__(self) -> str:
return f"{self.name}({self.arg})"
class MultiArgAction(Action, ABC):
arity = ActionArity.MULTI
def get_all_segments(self) -> List[PromptSegment]:
segments = []
for arg in self.all_args:
for seg_or_action in arg:
if isinstance(seg_or_action, Action):
segments.extend(seg_or_action.get_all_segments())
else:
segments.append(seg_or_action)
return segments
def __init__(
self,
args: List[List[Union[PromptSegment, Action]]],
):
self.all_args = args
-49
View File
@@ -1,49 +0,0 @@
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
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)
+4 -56
View File
@@ -1,63 +1,11 @@
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
from custom_nodes.KepPromptLang.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
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."""
def get_embedding(seg_or_action: SegOrAction, embedding_module: Embedding) -> 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_result(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__}")
-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
View File
@@ -1,73 +0,0 @@
from abc import ABC, abstractmethod
from dataclasses import dataclass
from enum import Enum
from typing import List, Optional, Tuple, Union
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()
class Action(ABC):
arity: ActionArity = ActionArity.NONE
display_name: str = ""
action_name: str = ""
description: str = ""
grammar: str = ""
usage_examples: List[str] = []
@abstractmethod
def __init__(self, *args, **kwargs) -> None: ...
@abstractmethod
def token_length(self) -> int: ...
@abstractmethod
def get_result(self, embedding_module: Embedding) -> ActionResult: ...
class SingleArgAction(Action, ABC):
arity = ActionArity.SINGLE
def __init__(self, arg: List):
self.arg = arg
def __repr__(self) -> str:
return f"{self.action_name}({self.arg})"
class MultiArgAction(Action, ABC):
arity = ActionArity.MULTI
def __init__(self, args: List[List]):
self.all_args = args
def __repr__(self) -> str:
joined = " | ".join(str(a) for a in self.all_args)
return f"{self.action_name}({joined})"
+62 -24
View File
@@ -1,39 +1,77 @@
from typing import List
from typing import Union, List
from torch import Tensor
import torch
from torch.nn import Embedding
from .action_utils import add_with_broadcast, concat_embeddings
from .base import MultiArgAction
from .types import SegOrAction
from custom_nodes.KepPromptLang.lib.action.base import Action, MultiArgAction
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class DiffAction(MultiArgAction):
grammar = 'diff(" arg ("|" arg)* ")"'
name = "diff"
chars = ["-", "-"]
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:
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
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)
def token_length(self) -> int:
# Sum adds to the embeddings of the base segment, so the length is the length of the base segment
return sum(seg_or_action.token_length() for seg_or_action in self.base_arg)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the base segment
all_base_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.base_arg
]
result = torch.cat(all_base_embeddings, dim=1)
for arg in self.additional_args:
arg_embedding = concat_embeddings(arg, embedding_module)
result = add_with_broadcast(result, arg_embedding, op="sub")
all_arg_embeddings = [
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
]
arg_embedding = torch.cat(all_arg_embeddings, dim=1)
if (
arg_embedding.shape[-2] == 1
or result.shape[-2] == arg_embedding.shape[-2]
):
result = result.sub(arg_embedding)
else:
print(
"WARNING: shape mismatch when trying to apply sum, arg will be averaged"
)
result = result.sub(torch.mean(arg_embedding, dim=1, keepdim=True))
return result
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\tadditional_args={self.additional_args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.additional_args))})"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_arg, Action):
base_segment_repr = self.base_arg.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.additional_args, Action):
target_repr = self.additional_args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
+45 -19
View File
@@ -1,36 +1,62 @@
from typing import List
from torch import Tensor
import torch
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 custom_nodes.KepPromptLang.lib.action.base import (
Action,
MultiArgAction,
)
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
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)",
]
name = "mult"
chars = ["[", "]"]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 2:
raise ValueError("Multiply action expects exactly two arguments")
raise ValueError("Multiply action should have exactly two arguments")
self.target_arg = args[0]
self.parsed_multiplier = parse_numeric_arg(
args[1], action_name="Multiply", role="multiplier", cast=float
)
self._parse_multiplier(args[1])
def _parse_multiplier(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Multiply action first argument should have exactly one segment")
multiplier_seg_or_action = arg[0]
if isinstance(multiplier_seg_or_action, Action):
raise ValueError("Multiply action should not have an action as an argument")
try:
self.parsed_multiplier = float(multiplier_seg_or_action.text)
except ValueError:
raise ValueError("Multiply action should have an integer/float as the first argument")
def token_length(self) -> int:
return get_total_length(self.target_arg)
"""
Mult multiplies the embeddings of the base segment, so the length is the length of the base segment
:return:
"""
total_length = 0
for seg_or_action in self.target_arg:
total_length += seg_or_action.token_length()
return total_length
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
all_embeddings = []
for seg_or_action in self.target_arg:
if isinstance(seg_or_action, Action):
all_embeddings.append(seg_or_action.get_result(embedding_module))
else:
all_embeddings.append(seg_or_action.get_embeddings(embedding_module))
target_embeddings = torch.cat(all_embeddings, dim=1)
return target_embeddings * self.parsed_multiplier
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)
+25 -14
View File
@@ -1,23 +1,34 @@
from torch import Tensor
import torch
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length
from .base import SingleArgAction
from custom_nodes.KepPromptLang.lib.action.base import Action, SingleArgAction
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
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)",
]
name = "neg"
chars = ["[", "]"]
def token_length(self) -> int:
return get_total_length(self.arg)
"""
Neg negates the embeddings of the base segment, so the length is the length of the base segment
:return:
"""
total_length = 0
for seg_or_action in self.arg:
total_length += seg_or_action.token_length()
return total_length
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
all_embeddings = []
for seg_or_action in self.arg:
if isinstance(seg_or_action, Action):
all_embeddings.append(seg_or_action.get_result(embedding_module))
else:
all_embeddings.append(seg_or_action.get_embeddings(embedding_module))
target_embeddings = torch.cat(all_embeddings, dim=1)
return target_embeddings * -1
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
+28 -15
View File
@@ -1,25 +1,38 @@
import torch
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length
from .base import SingleArgAction
from custom_nodes.KepPromptLang.lib.action.base import (
Action,
SingleArgAction,
)
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
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)))",
]
name = "norm"
chars = None
def token_length(self) -> int:
return get_total_length(self.arg)
"""
Norm normalizes the embeddings of the base segment, so the length is the length of the base segment
:return:
"""
total_length = 0
for seg_or_action in self.arg:
total_length += seg_or_action.token_length()
return total_length
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
all_embeddings = []
for seg_or_action in self.arg:
if isinstance(seg_or_action, Action):
target_embeddings = seg_or_action.get_result(embedding_module)
else:
target_embeddings = seg_or_action.get_embeddings(embedding_module)
all_embeddings.append(target_embeddings)
target_embeddings = torch.cat(all_embeddings, dim=1)
return torch.div(target_embeddings, torch.norm(target_embeddings, dim=-1, keepdim=True))
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))
-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))
+68 -34
View File
@@ -1,53 +1,87 @@
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
from custom_nodes.KepPromptLang.lib.action.base import (
Action,
MultiArgAction,
)
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class RandAction(MultiArgAction):
grammar = 'rand(" arg ")"'
name = "rand"
chars = None
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",
]
parsed_token_length = 0
range_min = 0
range_max = 1
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")
if len(args) != 1 and len(args) != 3:
raise ValueError("Random action should have exactly one argument or three arguments")
self._parse_token_length(args[0])
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
self._parse_range(args[1], args[2])
def _parse_token_length(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Random action first argument should have exactly one segment")
token_length_seg_or_action = arg[0]
if isinstance(token_length_seg_or_action, Action):
raise ValueError("Random action should not have an action as an argument")
try:
self.parsed_token_length = int(token_length_seg_or_action.text)
except ValueError:
raise ValueError("Random action should have an integer as the first argument")
def _parse_range(self, min_arg: List[SegOrAction], max_arg: List[SegOrAction]) -> None:
if len(min_arg) != 1:
raise ValueError("Random action second argument should have exactly one segment")
if len(max_arg) != 1:
raise ValueError("Random action third argument should have exactly one segment")
min_seg_or_action = min_arg[0]
max_seg_or_action = max_arg[0]
if isinstance(min_seg_or_action, Action):
raise ValueError("Random action should not have an action as an argument")
if isinstance(max_seg_or_action, Action):
raise ValueError("Random action should not have an action as an argument")
try:
self.range_min = int(min_seg_or_action.text)
except ValueError:
raise ValueError("Random action should have an integer as the second argument")
try:
self.range_max = int(max_seg_or_action.text)
except ValueError:
raise ValueError("Random action should have an integer as the third argument")
if self.range_min > self.range_max:
raise ValueError("Random action should have the second argument be less than the third argument")
def token_length(self) -> int:
"""
Random returns a random embedding whose length is the number in the argument
:return:
"""
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)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Create random tensor of size
result = torch.empty(1, self.parsed_token_length, embedding_module.embedding_dim).uniform_(self.range_min, self.range_max)
return result
-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
+80 -32
View File
@@ -1,52 +1,100 @@
from typing import List
from typing import List, Union
from torch import Tensor
import torch
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
from custom_nodes.KepPromptLang.lib.action.base import MultiArgAction, Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.actions.utils import slerp
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class SlerpAction(MultiArgAction):
grammar = 'slerp(" arg "|" arg "|" arg ")"'
name = "slerp"
chars = ["+", "+"]
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:
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
super().__init__(args)
if len(args) != 3:
raise ValueError("Slerp action expects exactly three arguments (2 vectors and a weight)")
raise ValueError("Slerp action should have 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
)
self._parse_weight(args[2])
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}"
)
self._validate_args()
def _parse_weight(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Slerp weight should have exactly one segment")
weight_seg_or_action = arg[0]
if isinstance(weight_seg_or_action, Action):
raise ValueError("Slerp weight should not have an action as an argument")
try:
self.parsed_weight = float(weight_seg_or_action.text)
except ValueError:
raise ValueError("Slerp should have an integer/float as the weight")
def _validate_args(self) -> None:
start_arg_token_length = sum(seg_or_action.token_length() for seg_or_action in self.start_argument)
end_arg_token_length = sum(seg_or_action.token_length() for seg_or_action in self.end_argument)
if start_arg_token_length != end_arg_token_length:
raise ValueError(f"Slerp start and end arguments should have the same length. Got {start_arg_token_length} and {end_arg_token_length}")
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)
# Slerp interpolates between the embeddings of the start and end segments, so the length is the length of the start segment
return sum(seg_or_action.token_length() for seg_or_action in 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)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the start segment
all_start_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.start_argument
]
start_embedding = torch.cat(all_start_embeddings, dim=1)
# Calculate the embeddings for the end segment
all_end_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.end_argument
]
end_embedding = torch.cat(all_end_embeddings, dim=1)
# Perform the slerp
result = slerp(self.parsed_weight, start_embedding, end_embedding)
return result
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.additional_args))})"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_arg, Action):
base_segment_repr = self.base_arg.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.additional_args, Action):
target_repr = self.additional_args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
+63 -18
View File
@@ -1,34 +1,79 @@
from typing import List
from typing import Union, List
from torch import Tensor
import torch
from torch.nn import Embedding
from .action_utils import add_with_broadcast, concat_embeddings
from .base import MultiArgAction
from .types import SegOrAction
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.action.base import Action, MultiArgAction
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class SumAction(MultiArgAction):
grammar = 'sum(" arg ("|" arg)+ ")"'
name = "sum"
chars = ["+", "+"]
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:
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> 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)
# Sum adds to the embeddings of the base segment, so the length is the length of the base segment
return sum(seg_or_action.token_length() for seg_or_action in self.base_arg)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the base segment
all_base_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.base_arg
]
result = torch.cat(all_base_embeddings, dim=1)
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")
all_arg_embeddings = [
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
]
arg_embedding = torch.cat(all_arg_embeddings, dim=1)
if (
arg_embedding.shape[-2] == 1
or result.shape[-2] == arg_embedding.shape[-2]
):
result = result.add(arg_embedding)
else:
print(
"WARNING: shape mismatch when trying to apply sum, arg will be averaged"
)
result = result.add(
torch.mean(arg_embedding, dim=1, keepdim=True)
)
return result
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.additional_args))})"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_arg, Action):
base_segment_repr = self.base_arg.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.additional_args, Action):
target_repr = self.additional_args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
+3 -4
View File
@@ -1,7 +1,6 @@
from typing import Union
from ..parser.prompt_segment import PromptSegment
from .base import Action
from .weighted import WeightedGroup
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.action.base import Action
SegOrAction = Union[PromptSegment, Action, WeightedGroup]
SegOrAction = Union[PromptSegment, Action]
+27 -12
View File
@@ -1,23 +1,38 @@
from typing import List
import torch
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
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)
def batch_size_info(batch: List[SegOrAction]):
for segment in batch:
print("Token Len: " + str(segment.token_length()))
print(segment.depth_repr())
def slerp(val: float, low: torch.Tensor, high: torch.Tensor, epsilon=1e-5):
# Convert val to tensor and clamp between 0 and 1
val = torch.tensor(val, dtype=torch.float32).clamp(0, 1)
# Normalize the vectors
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)
# Calculate the cosine of the angle between the vectors
dot = (low_norm * high_norm).sum(-1, keepdim=True)
# Clamp to prevent numerical errors
dot = torch.clamp(dot, -1, 1)
omega = torch.acos(dot)
# Slerp formula
sin_omega = torch.sin(omega)
scale_0 = torch.sin((1.0 - val) * omega) / (sin_omega + epsilon)
scale_1 = torch.sin(val * omega) / (sin_omega + epsilon)
scale_low = torch.sin((1.0 - val_t) * omega) / (sin_omega + epsilon)
scale_high = torch.sin(val_t * omega) / (sin_omega + epsilon)
# Handle the case where omega is small (the vectors are close)
close_condition = sin_omega < epsilon
scale_0 = torch.where(close_condition, 1.0 - val, scale_0)
scale_1 = torch.where(close_condition, val, scale_1)
# 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
return scale_0 * low + scale_1 * 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
}
+189 -107
View File
@@ -1,129 +1,211 @@
"""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
import contextlib
import os
from typing import List
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 .actions.base import ACTION_CONTINUATION, Action, PostModifiers
from comfy import model_management
import comfy.ops
from custom_nodes.KepPromptLang.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.fun_clip_stuff import PromptLangTextModel
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class PromptLangSDClipModel(sd1_clip.SDClipModel):
def process_tokens(self, tokens, device): # type: ignore[override]
embedding_module = self.transformer.get_input_embeddings()
# Methods with no comment can be assumed to be the same as comfy.sd1_clip.SD1ClipModel
class PromptLangClipModel(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, dtype=None): # clip-vit-base-patch32
super().__init__()
assert layer in self.LAYERS
self.num_layers = 12
if textmodel_path is not None:
# Our transformer
self.transformer = PromptLangTextModel.from_pretrained(textmodel_path)
else:
if textmodel_json_config is None:
# TODO: Maybe re-use clip config?
# Config could come from cond_stage_model.transformer.config
# Copied clip_config
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(device, dtype):
with modeling_utils.no_init_weights():
# Our transformer
self.transformer = PromptLangTextModel(config)
if dtype is not None:
self.transformer.to(dtype)
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 = torch.nn.Parameter(torch.eye(self.transformer.get_input_embeddings().weight.shape[1]))
self.logit_scale = torch.nn.Parameter(torch.tensor(4.6055))
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]
# Completely changed to support Segments and actions
def set_up_textual_embeddings(self, tokens: List[List[SegOrAction]], 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()
)
# Support our set_up_textual_embeddings which modifies the input embeddings
def forward(self, tokens):
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
if backup_embeds.weight.dtype != torch.float32:
precision_scope = torch.autocast
else:
precision_scope = contextlib.nullcontext
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
with precision_scope(model_management.get_autocast_device(device)):
outputs = self.transformer(input_ids=tokens, output_hidden_states=self.layer == "hidden")
self.transformer.set_input_embeddings(backup_embeds)
class PromptLangSDXLClipG(sdxl_clip.SDXLClipG, PromptLangSDClipModel):
"""SDXL's larger CLIP-G text encoder, with our DSL-aware process_tokens."""
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)
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.float().to(self.text_projection.device) @ self.text_projection.float()
return z.float(), pooled_output.float()
return out
def encode(self, tokens):
return self(tokens)
def load_sd(self, sd):
if "text_projection" in sd:
self.text_projection[:] = sd.pop("text_projection")
if "text_projection.weight" in sd:
self.text_projection[:] = sd.pop("text_projection.weight").transpose(0, 1)
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,
)
# Changed from comfy.sd1_clip.ClipTokenWeightEncoder
# Changed to use PromptSegments
def encode_token_weights(self, prompt_segments: List[List[SegOrAction]]):
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)
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()
+191
View File
@@ -0,0 +1,191 @@
from typing import Optional, Tuple, Union, List
import torch
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.KepPromptLang.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
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 PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
def __init__(self, config: CLIPTextConfig):
super().__init__(config)
def forward(
self,
input_dicts: Optional[List[List[SegOrAction]]] = None,
input_ids: Optional[torch.LongTensor] = None,
position_ids: Optional[torch.LongTensor] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
) -> torch.Tensor:
if input_dicts is None:
raise ValueError("You have to specify input_dicts")
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]
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 PrompLangCLIPTextTransformer(CLIPTextTransformer):
def __init__(self, config: CLIPTextConfig):
super().__init__(config)
self.embeddings = PromptLangCLIPTextEmbeddings(config)
def forward(
self,
input_ids: Optional[List[List[SegOrAction]]] = 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
input_shape = torch.Size([bsz, seq_len])
# CLIP's text model uses causal mask, prepare it here.
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
## VERSION DIFF ##
# transformers < 4.30.0
if hasattr(self, "_build_causal_attention_mask"):
print("Using transformers < 4.30.0")
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(hidden_states.device)
else:
# transformers >= 4.30.0
print("Using transformers >= 4.30.0")
from transformers.models.clip.modeling_clip import _make_causal_mask
causal_attention_mask = _make_causal_mask(input_shape, hidden_states.dtype, device=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,
)
# This is necessary to pass the PromptLangCLIPTextTransformer
class PromptLangTextModel(CLIPTextModel):
def __init__(self, config: CLIPTextConfig):
super().__init__(config)
self.text_model = PrompLangCLIPTextTransformer(config)
def forward(
self,
input_ids: Optional[List[List[SegOrAction]]] = 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]:
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)
+3 -21
View File
@@ -1,38 +1,20 @@
grammar = r"""
?start: stmt+
?stmt: assign
| item
assign: "$" NAME "=" arg ";"
grammar = """
?start: item+
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]*/
QUOTED_STRING: /"([^"\\\]*(\\\.[^"\\\]*)*)"|'([^'\\\]*(\\\.[^'\\\]*)*)'/
%import common.WS
%ignore WS
%ignore COMMENT
"""
+15 -14
View File
@@ -1,4 +1,4 @@
from typing import List, Union
from typing import Union, List
import torch
from torch import Tensor
@@ -6,25 +6,26 @@ 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 __repr__(self):
return f'"{self.text}"{self.tokens}'
def token_length(self) -> int:
def token_length(self):
return len(self.tokens)
def get_embeddings(self, embedding_module: Embedding) -> Tensor:
"""Look up embeddings for plain int tokens.
tensors = torch.LongTensor(self.tokens).to(embedding_module.weight.device)
unsqueezed_tensors = tensors.unsqueeze(0)
return embedding_module(unsqueezed_tensors)
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))
def depth_repr(self, depth=1):
out = f'"{self.text}"('
cleaned_tokens = list(map(lambda x: str(x) if isinstance(x, int) else "EMBD", self.tokens))
out += ", ".join(cleaned_tokens)
out += ")"
return out
+4 -8
View File
@@ -1,18 +1,14 @@
from typing import Dict, Type
from typing import Type, Dict
from ..actions.base import Action
from custom_nodes.KepPromptLang.lib.action.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
action_registry[str(action.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]
+34 -52
View File
@@ -1,81 +1,63 @@
from typing import List
from lark import Token, Transformer
from lark import Transformer, Token
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
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.KepPromptLang.lib.action.base import Action, ActionArity
from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction
from custom_nodes.KepPromptLang.lib.actions.rand import RandAction
from custom_nodes.KepPromptLang.lib.parser.registration import get_action_by_name
from custom_nodes.KepPromptLang.lib.parser.utils import build_prompt_segment
from custom_nodes.KepPromptLang.lib.actions.neg import NegAction
from custom_nodes.KepPromptLang.lib.actions.norm import NormAction
from custom_nodes.KepPromptLang.lib.actions.sum import SumAction
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class PromptTransformer(Transformer):
"""Maps the Lark parse tree into a flat list of PromptSegments and Actions."""
# def WORD(self, items):
# return items
def __init__(self, tokenizer: SDTokenizer):
def __init__(self, tokenizer: SD1Tokenizer):
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)):
if isinstance(item, Action):
return item
if isinstance(item, PromptSegment):
return item
if item.type == "WORD":
return build_prompt_segment(str(item), self.tokenizer)
if item.type == "QUOTED_STRING":
# Strip surrounding quotes, unescape \" and \'.
elif item.type == "QUOTED_STRING":
# Remove the quotes
unquoted = item[1:-1]
unescaped = unquoted.replace('\\"', '"').replace("\\'", "'")
# Replace escaped quotes with quotes
unescaped = unquoted.replace("\\\"", "\"").replace("\\\'", "\'")
return build_prompt_segment(unescaped, self.tokenizer)
raise ValueError(f"Unknown item type: {item.type}")
elif item.type == "embedding":
return build_prompt_segment(item, self.tokenizer)
elif item.type == "function":
return item
else:
raise Exception("Unknown item type: " + str(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,
)
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")
raise ValueError(f"Action {action.name} should have exactly one argument")
return action(items[1])
if action.arity == ActionArity.MULTI:
return action(items[1:])
raise ValueError(f"Unknown action arity: {action.arity}")
elif action.arity == ActionArity.MULTI:
return action(items[1:][:])
else:
raise ValueError(f"Unknown action arity: {action.arity}")
+12 -15
View File
@@ -1,30 +1,28 @@
from lark import Token
from comfy.sd1_clip import SDTokenizer
from .prompt_segment import PromptSegment
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.KepPromptLang.lib.parser.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], [])
else:
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."""
def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment:
split_text = text.split(" ")
tokens = []
for word in text.split(" "):
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")
embedding, leftover = tokenizer._try_get_embedding(embedding_name)
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")
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)
@@ -35,7 +33,6 @@ def build_prompt_segment(text: str, tokenizer: SDTokenizer) -> PromptSegment:
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)
+54 -108
View File
@@ -1,122 +1,68 @@
"""DSL-aware tokenizers.
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`.
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.
"""
from typing import Dict, Iterable, List, Tuple, Union
from typing import List
from lark import Tree
from comfy.sd1_clip import SD1Tokenizer, SDTokenizer
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
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
from custom_nodes.KepPromptLang.lib.parser import PromptParser
from custom_nodes.KepPromptLang.lib.parser.transformer import PromptTransformer
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
# Side-effect import: registers all built-in actions with the parser.
from . import actions # noqa: F401
class PromptLangTokenizer(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)
TokenEntry = Tuple[Union[int, "Action", object], float]
"""
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[SegOrAction]]:
if self.pad_with_end:
pad_token = self.end_token
else:
pad_token = 0
parsed_prompt = PromptParser.parse(text)
parsed_actions = PromptTransformer(self).transform(parsed_prompt)
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__})")
# reshape token array to CLIP input size
batched_segments = []
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token])]
# batched_segments.append(batch)
batch_size = 1
if isinstance(parsed_actions, Tree):
segments_to_process = parsed_actions.children
else:
segments_to_process = [parsed_actions]
for segment in segments_to_process:
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
# 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)
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)
# start new batch
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token]), segment]
batch_size = num_tokens + 1 # +1 for start token
continue
def _batch_from_tree(self, tree) -> List[List[TokenEntry]]:
pad_token = self.end_token if self.pad_with_end else 0
# Since the segment fits in the current batch, add it
batch.append(segment)
batch_size += num_tokens
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]
# 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)
batches: List[List[TokenEntry]] = []
current: List[TokenEntry] = [(self.start_token, 1.0)]
# for batch in batched_segments:
# batch_size_info(batch)
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
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,
)
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 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)
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),
}
def untokenize(self, token_weight_pair):
return self.clip_g.untokenize(token_weight_pair)
def state_dict(self):
return {}
return batched_segments
+101 -170
View File
@@ -1,24 +1,24 @@
import os
from dataclasses import dataclass
from typing import Any, List, Tuple
import random
from typing import List, Tuple
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.KepPromptLang.lib.clip_model import PromptLangClipModel
from .lib.clip_model import PromptLangSD1ClipModel, PromptLangSDXLClipModel
from .lib.inspect import inspect_prompt
from .lib.tokenizer import PromptLangSD1Tokenizer, PromptLangSDXLTokenizer
from custom_nodes.KepPromptLang.lib.tokenizer import PromptLangTokenizer
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(cls): # type: ignore
return {
"required": {
"source_clip": ("CLIP",),
@@ -27,200 +27,131 @@ 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
clip_target = EmptyClass()
clip_target.params = {}
clip_target.clip = PromptLangClipModel
clip_target.tokenizer = PromptLangTokenizer
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,)
class PromptLangInspect:
"""Shows what a DSL prompt resolves to at the embedding layer: per-slot weight, L2 norm, nearest vocab."""
@classmethod
def INPUT_TYPES(cls): # type: ignore[no-untyped-def]
return {
"required": {
"clip": ("CLIP",),
"text": ("STRING", {"multiline": True}),
"top_k": ("INT", {"default": 3, "min": 1, "max": 10}),
}
}
RETURN_TYPES = ("STRING",)
FUNCTION = "inspect"
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,)}
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,)
def tensor2img(tensor_img) -> Image.Image:
arr = (255.0 * tensor_img.cpu().numpy()).clip(0, 255).astype(np.uint8)
return Image.fromarray(arr)
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()
pass
@classmethod
def INPUT_TYPES(cls): # type: ignore[no-untyped-def]
def INPUT_TYPES(cls): # type: ignore
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_val = 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_val == -1:
split_chunks = 1
chunk_len = len(images)
split_every_val = len(images)
else:
chunk_len = split_requested
split_chunks = len(images) // chunk_len
split_chunks = int(len(images) / split_every_val)
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_val * chunk_idx : split_every_val * (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_val):
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_val * split_idx
split_end = split_every_val * (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"]
+12
View File
@@ -0,0 +1,12 @@
# If the models/checkpoints folder does not have test.txt, then download the model.
import os
from huggingface_hub import hf_hub_download
FILE = "v1-5-pruned-emaonly.safetensors"
REPO_ID = "runwayml/stable-diffusion-v1-5"
if not os.path.exists(f"models/checkpoints/{FILE}"):
print("Downloading model...")
hf_hub_download(repo_id=REPO_ID, filename=FILE, local_dir="models/checkpoints", local_dir_use_symlinks=False)
else:
print("Model already downloaded.")
+109
View File
@@ -0,0 +1,109 @@
#This is an example that uses the websockets api to know when a prompt execution is done
#Once the prompt execution is done it downloads the images using the /history endpoint
import os
import websocket #NOTE: websocket-client (https://github.com/websocket-client/websocket-client)
import uuid
import json
import urllib.request
import urllib.parse
server_address = "127.0.0.1:8188"
client_id = str(uuid.uuid4())
def queue_prompt(prompt):
p = {"prompt": prompt, "client_id": client_id}
data = json.dumps(p).encode('utf-8')
req = urllib.request.Request("http://{}/prompt".format(server_address), data=data)
try:
response = urllib.request.urlopen(req)
return json.loads(response.read())
except urllib.error.HTTPError as e:
print(f"HTTP Error {e.code}: {e.reason}")
error_body = e.read()
# Attempt to read and print the JSON error body
try:
json_error_body = json.loads(error_body)
print(json_error_body)
raise e
except json.JSONDecodeError:
print("Failed to decode error response as JSON.")
print(error_body)
raise e
except Exception as e:
print(f"An unexpected error occurred: {e}")
raise e
def get_image(filename, subfolder, folder_type):
data = {"filename": filename, "subfolder": subfolder, "type": folder_type}
url_values = urllib.parse.urlencode(data)
with urllib.request.urlopen("http://{}/view?{}".format(server_address, url_values)) as response:
return response.read()
def get_history(prompt_id):
with urllib.request.urlopen("http://{}/history/{}".format(server_address, prompt_id)) as response:
return json.loads(response.read())
def get_images(ws, prompt):
prompt_id = queue_prompt(prompt)['prompt_id']
output_images = {}
while True:
out = ws.recv()
if isinstance(out, str):
message = json.loads(out)
if message['type'] == 'executing':
data = message['data']
if data['node'] is None and data['prompt_id'] == prompt_id:
break #Execution is done
else:
continue #previews are binary data
history = get_history(prompt_id)[prompt_id]
for o in history['outputs']:
for node_id in history['outputs']:
node_output = history['outputs'][node_id]
if 'images' in node_output:
images_output = []
for image in node_output['images']:
image_data = get_image(image['filename'], image['subfolder'], image['type'])
images_output.append(image_data)
output_images[node_id] = images_output
return output_images
# Load json from file relative to this script
prompt = json.load(open(os.path.join(os.path.dirname(os.path.realpath(__file__)), "workflow_api.json")))
#set the text prompt for our positive CLIPTextEncode
# prompt["6"]["inputs"]["text"] = "masterpiece best quality man"
#set the seed for our KSampler node
print(queue_prompt(prompt))
ws = websocket.WebSocket()
ws.connect("ws://{}/ws?clientId={}".format(server_address, client_id))
while True:
out = ws.recv()
if isinstance(out, str):
message = json.loads(out)
# print(message)
if message["type"] == "executing" and message["data"]["node"] is None:
print("Execution is done")
break
if message["type"] == "execution_error":
print("Execution error")
print(json.dumps(message["data"], indent=4))
raise Exception("Execution error")
# images = get_images(ws, prompt)
#Commented out code to display the output images:
# for node_id in images:
# for image_data in images[node_id]:
# from PIL import Image
# import io
# image = Image.open(io.BytesIO(image_data))
# image.show()
+84
View File
@@ -0,0 +1,84 @@
{
"1": {
"inputs": {
"text": "A sum(cat|norm(sum(neg(parrot)|rabbit))) outside",
"clip": [
"2",
0
]
},
"class_type": "CLIPTextEncode"
},
"2": {
"inputs": {
"source_clip": [
"4",
1
]
},
"class_type": "Special CLIP Loader"
},
"3": {
"inputs": {
"seed": 556492279461741,
"steps": 1,
"cfg": 8,
"sampler_name": "euler",
"scheduler": "normal",
"denoise": 1,
"model": [
"4",
0
],
"positive": [
"1",
0
],
"negative": [
"1",
0
],
"latent_image": [
"7",
0
]
},
"class_type": "KSampler"
},
"4": {
"inputs": {
"ckpt_name": "v1-5-pruned-emaonly.safetensors"
},
"class_type": "CheckpointLoaderSimple"
},
"5": {
"inputs": {
"samples": [
"3",
0
],
"vae": [
"4",
2
]
},
"class_type": "VAEDecode"
},
"6": {
"inputs": {
"images": [
"5",
0
]
},
"class_type": "PreviewImage"
},
"7": {
"inputs": {
"width": 512,
"height": 512,
"batch_size": 1
},
"class_type": "EmptyLatentImage"
}
}
-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()))