Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e0a15a657c | ||
|
|
dbf1fbf287 | ||
|
|
0413b0c1c8 | ||
|
|
6238776610 | ||
|
|
bb03530d5d | ||
|
|
447d04dffe |
@@ -1,10 +0,0 @@
|
|||||||
name: Test
|
|
||||||
on:
|
|
||||||
workflow_dispatch:
|
|
||||||
jobs:
|
|
||||||
build:
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
steps:
|
|
||||||
- uses: actions/checkout@v2
|
|
||||||
- name: Setup upterm session
|
|
||||||
uses: lhotari/action-upterm@v1
|
|
||||||
@@ -1,85 +0,0 @@
|
|||||||
name: Run Test Workflow
|
|
||||||
on: [workflow_dispatch]
|
|
||||||
|
|
||||||
jobs:
|
|
||||||
Test:
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
strategy:
|
|
||||||
fail-fast: false
|
|
||||||
matrix:
|
|
||||||
python-version: [ "3.7", "3.8", "3.9", "3.10", "3.11" ]
|
|
||||||
steps:
|
|
||||||
- name: Clone Upstream
|
|
||||||
uses: actions/checkout@v3
|
|
||||||
with:
|
|
||||||
repository: comfyanonymous/ComfyUI
|
|
||||||
ref: master
|
|
||||||
fetch-depth: 0
|
|
||||||
|
|
||||||
- name: Clone Node
|
|
||||||
uses: actions/checkout@v3
|
|
||||||
with:
|
|
||||||
ref: master
|
|
||||||
fetch-depth: 0
|
|
||||||
path: custom_nodes/KepPromptLang
|
|
||||||
|
|
||||||
- name: Setup Python
|
|
||||||
uses: actions/setup-python@v4
|
|
||||||
with:
|
|
||||||
# Version range or exact version of Python or PyPy to use, using SemVer's version range syntax. Reads from .python-version if unset.
|
|
||||||
python-version: ${{ matrix.python-version }}
|
|
||||||
|
|
||||||
- name: Cache virtualenv
|
|
||||||
uses: actions/cache@v3
|
|
||||||
id: cache-venv
|
|
||||||
with:
|
|
||||||
path: ./.venv/
|
|
||||||
key: ${{ runner.os }}-venv-${{ matrix.python-version }}-${{ hashFiles('**/requirements.txt') }}
|
|
||||||
restore-keys: |
|
|
||||||
${{ runner.os }}-venv-${{ matrix.python-version }}-
|
|
||||||
|
|
||||||
- name: Install Requirements
|
|
||||||
if: steps.cache-venv.outputs.cache-hit != 'true'
|
|
||||||
run: |
|
|
||||||
python -m venv ./.venv
|
|
||||||
source ./.venv/bin/activate
|
|
||||||
pip install torch --index-url https://download.pytorch.org/whl/cpu
|
|
||||||
pip install -r requirements.txt
|
|
||||||
pip install -r custom_nodes/KepPromptLang/requirements.txt
|
|
||||||
pip install huggingface_hub websocket-client
|
|
||||||
|
|
||||||
# - name: Cache SD Checkpoint
|
|
||||||
# uses: actions/cache@v3
|
|
||||||
# with:
|
|
||||||
# path: |
|
|
||||||
# models/checkpoints
|
|
||||||
# key: ${{ runner.os }}-sd-15-checkpoint
|
|
||||||
|
|
||||||
- name: Check and Download Model
|
|
||||||
run: |
|
|
||||||
source ./.venv/bin/activate
|
|
||||||
python custom_nodes/KepPromptLang/test_files/check_and_download_model.py
|
|
||||||
|
|
||||||
- name: Run in Background
|
|
||||||
env:
|
|
||||||
PYTHONUNBUFFERED: 1
|
|
||||||
run: |
|
|
||||||
source ./.venv/bin/activate
|
|
||||||
python main.py --cpu &> server.log &
|
|
||||||
sleep 10
|
|
||||||
|
|
||||||
# - name: Setup upterm session
|
|
||||||
# uses: lhotari/action-upterm@v1
|
|
||||||
|
|
||||||
- name: Run Workflow
|
|
||||||
run: |
|
|
||||||
source ./.venv/bin/activate
|
|
||||||
python custom_nodes/KepPromptLang/test_files/run_workflow.py
|
|
||||||
|
|
||||||
- name: Upload Comfy Server Log
|
|
||||||
if: always()
|
|
||||||
uses: actions/upload-artifact@v3
|
|
||||||
with:
|
|
||||||
name: comfy-server-log-${{ matrix.python-version }}
|
|
||||||
path: server.log
|
|
||||||
|
|
||||||
@@ -1,100 +1 @@
|
|||||||
# KepPromptLang
|
# ClipStuff
|
||||||
|
|
||||||
A small DSL for ComfyUI that lets you do math on CLIP token embeddings before they're fed into the text transformer.
|
|
||||||
|
|
||||||
```
|
|
||||||
sum(diff(king|man)|woman)
|
|
||||||
norm(sum(cat | dog | horse | parrot))
|
|
||||||
A slerp(cat|dog|0.5) is happy
|
|
||||||
```
|
|
||||||
|
|
||||||
## Install
|
|
||||||
|
|
||||||
Clone into `ComfyUI/custom_nodes/`:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
cd ComfyUI/custom_nodes
|
|
||||||
git clone <repo-url> KepPromptLang
|
|
||||||
pip install -r KepPromptLang/requirements.txt
|
|
||||||
```
|
|
||||||
|
|
||||||
## Usage
|
|
||||||
|
|
||||||
1. Add a **Special CLIP Loader** node and feed it the CLIP output from your **Load Checkpoint**.
|
|
||||||
2. Pass the wrapped CLIP into a standard **CLIP Text Encode** node.
|
|
||||||
3. Use the DSL syntax in your prompt.
|
|
||||||
|
|
||||||
To debug what your DSL is doing, add a **PromptLang Inspect** node — it shows the per-slot weight, L2 norm, and nearest-vocab words for the resolved embeddings.
|
|
||||||
|
|
||||||
See `examples/WIP_Example_workflow.json` for a working workflow.
|
|
||||||
|
|
||||||

|
|
||||||
|
|
||||||
## Syntax
|
|
||||||
|
|
||||||
| Element | Syntax | Example |
|
|
||||||
| --- | --- | --- |
|
|
||||||
| Plain word | alphanumeric (with `,_.-`) | `cat`, `dog_face` |
|
|
||||||
| Quoted string | single or double quotes | `"hello world"`, `'it\'s sunny'` |
|
|
||||||
| Weighted | `(text:weight)` or `emph(text\|weight)` | `(cat:1.3)`, `emph(cat\|1.3)` |
|
|
||||||
| Embedding (textual inversion) | `embedding:NAME` | `embedding:face_vector` |
|
|
||||||
| Function | `name(arg \| arg \| ...)` | `sum(king \| woman)` |
|
|
||||||
|
|
||||||
Arguments inside a function are separated by `|`. Each arg can itself be plain text, an embedding, a quoted string, or another function call.
|
|
||||||
|
|
||||||
### Variables and comments
|
|
||||||
|
|
||||||
```
|
|
||||||
$axis = diff(king|queen); # name an expression
|
|
||||||
sum(actor|$axis) and reject(doctor|$axis)
|
|
||||||
```
|
|
||||||
|
|
||||||
`$NAME = arg;` binds a name; `$NAME` substitutes it. Single-pass: define before use, no reassignment. `#` comments run to end of line. Substitution is structural — multiple refs share the same parsed action object, but actions are evaluated per occurrence (so `$r = rand(3); $r $r` re-rolls each use).
|
|
||||||
|
|
||||||
## Quick examples
|
|
||||||
|
|
||||||
- Average two prompts: `avg(The cat is | The dog is | 0.5)`
|
|
||||||
- Normalize a sum: `norm(sum(cat | dog | horse))`
|
|
||||||
- King − Man + Woman: `sum(diff(king|man)|woman)` (or `sum(king | neg(man) | woman)`)
|
|
||||||
- Negate an embedding: `neg(embedding:body_vector)`
|
|
||||||
|
|
||||||
## Functions
|
|
||||||
|
|
||||||
| Display Name | Action Name | Description | Usage Examples |
|
|
||||||
| --- | --- | --- | --- |
|
|
||||||
| Average | avg | Performs a weighted average between two segments or actions. The recommended weight is 0 - 1. | <ul><li>avg(The cat is\|The dog is\|0.5)</li><li>avg(Cat\|Dog\|0.5)</li></ul> |
|
|
||||||
| Difference | diff | Subtracts the segments in the order they are given. The first segment is subtracted from the second, then the third from the result, and so on. | <ul><li>diff(The cat is\|The dog is)</li><li>diff(Cat\|Dog)</li><li>sum(diff(king\|man)\|woman)</li></ul> |
|
|
||||||
| Multiply | mult | Multiplies the provided segments or actions by the multiplier. | <ul><li>mult(The cat is\|2.5)</li><li>mult(Cat\|-1)</li></ul> |
|
|
||||||
| Nearest Vocab | nearest | Snaps a computed vector to the k nearest real vocabulary tokens (by cosine similarity), returning their embeddings concatenated. The input is mean-pooled before lookup. | <ul><li>nearest(sum(diff(king\|man)\|woman))</li><li>nearest(sum(red\|blue)\|3)</li></ul> |
|
|
||||||
| Negate | neg | Negates the provided segments or actions. | <ul><li>neg(cat)</li><li>sum(king\|neg(man)\|women)</li></ul> |
|
|
||||||
| Noise | noise | Adds Gaussian noise (mean 0, given std) to the embeddings of the first argument. | <ul><li>A noise(cat\|0.05) on a sunny day</li></ul> |
|
|
||||||
| Normalize | norm | Normalizes the provided segments or actions. | <ul><li>norm(cat)</li><li>sum(cat\|norm(sum(tiger\|fish)))</li></ul> |
|
|
||||||
| Positional Embedding Scale | posScale | Scales (multiplies) the positional embeddings of the provided segments or actions by the multiplier. | <ul><li>A posScale(cat\|1.5) on a rainy day</li></ul> |
|
|
||||||
| Ignore Positional Embeddings | postPos | Prevents positional embeddings from being applied to the provided segments or actions. | <ul><li>A postPos(cat) on a rainy day</li></ul> |
|
|
||||||
| Project | proj | Projects the first argument onto the direction of the second (mean, unit-normalized). | <ul><li>proj(king\|gender)</li><li>diff(style\|proj(style\|photorealistic))</li></ul> |
|
|
||||||
| Random Embedding | rand | Returns a random embedding of the specified token length, with the values optionally bounded by the second and third arguments. | <ul><li>A rand(1) cat</li><li>A rand(1\|-1\|1) cat</li></ul> |
|
|
||||||
| Reject | reject | Removes the component of the first argument along the direction of the second (a - proj(a\|b)). | <ul><li>reject(anime girl\|anime)</li></ul> |
|
|
||||||
| Renormalize | renorm | Rescales the first argument so each token's L2 norm matches the (mean) L2 norm of the reference. | <ul><li>renorm(sum(king\|neg(man)\|woman)\|queen)</li></ul> |
|
|
||||||
| Scale Dimensions | scaleDims | Scales the specified dimensions of the input embeddings by the specified amount | <ul><li>The scaleDims(cat\|4,1.5\|76,1.2) is happy</li></ul> |
|
|
||||||
| Set Dimensions | setDims | Sets the specified dimensions of the input embeddings to the specified value | <ul><li>The setDims(cat\|4, -0.01253\|76, 1.2) is happy</li></ul> |
|
|
||||||
| Slerp | slerp | Performs a slerp (interpolation) between two segments or actions, with the given weight. The recommended weight is 0 - 1. | <ul><li>The slerp(cat\|dog\|0.5) is happy</li></ul> |
|
|
||||||
| Sum | sum | Adds the embeddings of the provided segments or actions. | <ul><li>A happy sum(cat\|dog\|shark)</li></ul> |
|
|
||||||
|
|
||||||
`lerp(a|b|t)` is also accepted as an alias for `avg(a|b|t)`.
|
|
||||||
|
|
||||||
Regenerate the table with `python tools/build_docs.py`.
|
|
||||||
|
|
||||||
## Development
|
|
||||||
|
|
||||||
Tests are pytest-based and don't require ComfyUI:
|
|
||||||
|
|
||||||
```bash
|
|
||||||
pip install -e ".[dev]"
|
|
||||||
python -m pytest
|
|
||||||
```
|
|
||||||
|
|
||||||
## Compatibility
|
|
||||||
|
|
||||||
- SD1.x (CLIP-L) and SDXL (CLIP-L + CLIP-G).
|
|
||||||
- SD2 is not supported.
|
|
||||||
- Two pooler-output actions (`_exp-pooler`, `_exp-pooledAvg`) from earlier versions were experimental and have been removed; they relied on direct HuggingFace transformer access that is no longer how ComfyUI structures its CLIP encoders.
|
|
||||||
|
|||||||
+6
-10
@@ -1,15 +1,11 @@
|
|||||||
from .nodes import BuildGif, PromptLangInspect, SpecialClipLoader
|
from .nodes import (
|
||||||
|
KepAdvTextEncode,
|
||||||
|
BuildGif,
|
||||||
|
SpecialClipLoader,
|
||||||
|
)
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
|
"Kep Adv Text Encode": KepAdvTextEncode,
|
||||||
"Build Gif": BuildGif,
|
"Build Gif": BuildGif,
|
||||||
"Special CLIP Loader": SpecialClipLoader,
|
"Special CLIP Loader": SpecialClipLoader,
|
||||||
"PromptLang Inspect": PromptLangInspect,
|
|
||||||
}
|
}
|
||||||
|
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
|
||||||
"Build Gif": "Build GIF (KepPromptLang)",
|
|
||||||
"Special CLIP Loader": "Special CLIP Loader (KepPromptLang)",
|
|
||||||
"PromptLang Inspect": "PromptLang Inspect",
|
|
||||||
}
|
|
||||||
|
|
||||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
|
||||||
|
|||||||
Binary file not shown.
|
Before Width: | Height: | Size: 1.8 MiB |
File diff suppressed because it is too large
Load Diff
+5
-48
@@ -1,49 +1,6 @@
|
|||||||
from ..parser.registration import register_action
|
from custom_nodes.ClipStuff.lib.actions.arith import ArithAction
|
||||||
from .avg import AverageAction
|
from custom_nodes.ClipStuff.lib.actions.nudge import NudgeAction
|
||||||
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 [
|
ALL_ACTIONS = [NudgeAction, ArithAction]
|
||||||
AverageAction,
|
ALL_START_CHARS = [action.START_CHAR for action in ALL_ACTIONS]
|
||||||
DiffAction,
|
ALL_END_CHARS = [action.END_CHAR for action in ALL_ACTIONS]
|
||||||
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)
|
|
||||||
|
|||||||
@@ -1,63 +0,0 @@
|
|||||||
from typing import Callable, List, TypeVar
|
|
||||||
|
|
||||||
import torch
|
|
||||||
from torch import Tensor
|
|
||||||
from torch.nn import Embedding
|
|
||||||
|
|
||||||
from .base import Action
|
|
||||||
from .types import SegOrAction
|
|
||||||
from .weighted import WeightedGroup
|
|
||||||
|
|
||||||
|
|
||||||
def embedding_tensor(seg_or_action: SegOrAction, embedding_module: Embedding) -> Tensor:
|
|
||||||
"""Embeddings for a segment, or get_result() for an action — always a bare tensor."""
|
|
||||||
if isinstance(seg_or_action, Action):
|
|
||||||
result = seg_or_action.get_result(embedding_module)
|
|
||||||
return result[0] if isinstance(result, tuple) else result
|
|
||||||
if isinstance(seg_or_action, WeightedGroup):
|
|
||||||
# Weights apply post-transformer, not to embedding math; recurse and drop the weight.
|
|
||||||
return concat_embeddings(seg_or_action.items, embedding_module)
|
|
||||||
return seg_or_action.get_embeddings(embedding_module)
|
|
||||||
|
|
||||||
|
|
||||||
def get_total_length(args: List[SegOrAction]) -> int:
|
|
||||||
return sum(seg_or_action.token_length() for seg_or_action in args)
|
|
||||||
|
|
||||||
|
|
||||||
def concat_embeddings(args: List[SegOrAction], embedding_module: Embedding) -> Tensor:
|
|
||||||
"""Materialize and concatenate embeddings for a sequence of segments/actions along the seq dim."""
|
|
||||||
return torch.cat([embedding_tensor(x, embedding_module) for x in args], dim=1)
|
|
||||||
|
|
||||||
|
|
||||||
def add_with_broadcast(result: Tensor, arg_embedding: Tensor, op: str) -> Tensor:
|
|
||||||
"""Add or subtract arg_embedding into result, averaging arg over the seq dim if shapes mismatch."""
|
|
||||||
matched = arg_embedding.shape[-2] == 1 or result.shape[-2] == arg_embedding.shape[-2]
|
|
||||||
if not matched:
|
|
||||||
print(f"WARNING: shape mismatch when trying to apply {op}, arg will be averaged")
|
|
||||||
arg_embedding = torch.mean(arg_embedding, dim=1, keepdim=True)
|
|
||||||
return result.add(arg_embedding) if op == "add" else result.sub(arg_embedding)
|
|
||||||
|
|
||||||
|
|
||||||
T = TypeVar("T")
|
|
||||||
|
|
||||||
|
|
||||||
def parse_numeric_arg(
|
|
||||||
arg: List[SegOrAction],
|
|
||||||
*,
|
|
||||||
action_name: str,
|
|
||||||
role: str,
|
|
||||||
cast: Callable[[str], T] = float,
|
|
||||||
) -> T:
|
|
||||||
"""Pull a single numeric value out of a one-segment arg, with helpful errors.
|
|
||||||
|
|
||||||
Used by every action that takes a scalar weight/multiplier/length.
|
|
||||||
"""
|
|
||||||
if len(arg) != 1:
|
|
||||||
raise ValueError(f"{action_name} {role} should have exactly one segment")
|
|
||||||
item = arg[0]
|
|
||||||
if isinstance(item, (Action, WeightedGroup)):
|
|
||||||
raise ValueError(f"{action_name} {role} must be a plain numeric segment")
|
|
||||||
try:
|
|
||||||
return cast(item.text)
|
|
||||||
except ValueError:
|
|
||||||
raise ValueError(f"{action_name} {role} should be a {cast.__name__}")
|
|
||||||
@@ -0,0 +1,131 @@
|
|||||||
|
from typing import Callable, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch.nn import Embedding
|
||||||
|
|
||||||
|
from comfy.sd1_clip import SD1Tokenizer
|
||||||
|
from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment
|
||||||
|
|
||||||
|
|
||||||
|
class ArithAction(Action):
|
||||||
|
START_CHAR = "<"
|
||||||
|
END_CHAR = ">"
|
||||||
|
|
||||||
|
def __init__(self, base_segment: PromptSegment | Action, ops: dict[str, list[PromptSegment | Action]]):
|
||||||
|
self.base_segment = base_segment
|
||||||
|
self.ops = ops
|
||||||
|
|
||||||
|
def __repr__(self):
|
||||||
|
return f"ArithAction(\n\tbase_segment={self.base_segment},\n\tops={self.ops}\n)"
|
||||||
|
|
||||||
|
def depth_repr(self, depth=1):
|
||||||
|
out = "ArithAction(\n"
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
base_segment_repr = self.base_segment.depth_repr(depth + 1)
|
||||||
|
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
|
||||||
|
elif isinstance(self.base_segment, PromptSegment):
|
||||||
|
out += "\t" * depth + f'base_segment={self.base_segment.depth_repr(depth)}'
|
||||||
|
else:
|
||||||
|
out += "\t" * depth + f'base_segment="{self.base_segment}",'
|
||||||
|
|
||||||
|
for op_key, ops in self.ops.items():
|
||||||
|
for op in ops:
|
||||||
|
out += "\n" + "\t" * depth + f'"{op_key}":[\n'
|
||||||
|
if isinstance(op, Action):
|
||||||
|
op_repr = op.depth_repr(depth + 2)
|
||||||
|
out += "\t" * (depth + 1) + f"{op_repr}\n"
|
||||||
|
else:
|
||||||
|
out += "\t" * (depth + 1) + f'{op.depth_repr()},\n'
|
||||||
|
out += "\t" * depth + "],"
|
||||||
|
out += "\n" + "\t" * (depth - 1) + ")"
|
||||||
|
return out
|
||||||
|
|
||||||
|
def token_length(self):
|
||||||
|
# ArithAction modifies the embeddings of the base segment, so the length is the length of the base segment
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
return self.base_segment.token_length()
|
||||||
|
|
||||||
|
return len(self.base_segment.tokens)
|
||||||
|
|
||||||
|
|
||||||
|
def get_all_segments(self):
|
||||||
|
segments = []
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
segments += self.base_segment.get_all_segments()
|
||||||
|
else:
|
||||||
|
segments.append(self.base_segment)
|
||||||
|
|
||||||
|
for op_key, ops in self.ops.items():
|
||||||
|
for op in ops:
|
||||||
|
if isinstance(op, Action):
|
||||||
|
segments += op.get_all_segments()
|
||||||
|
else:
|
||||||
|
segments.append(op)
|
||||||
|
|
||||||
|
return segments
|
||||||
|
|
||||||
|
def get_result(self, embedding_module: Embedding):
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
base_segment_result = self.base_segment.get_result(embedding_module)
|
||||||
|
else:
|
||||||
|
base_segment_result = self.base_segment.get_embeddings(embedding_module)
|
||||||
|
|
||||||
|
for op_key, ops in self.ops.items():
|
||||||
|
for op in ops:
|
||||||
|
if isinstance(op, Action):
|
||||||
|
op_result = op.get_result(embedding_module)
|
||||||
|
else:
|
||||||
|
op_result = op.get_embeddings(embedding_module)
|
||||||
|
|
||||||
|
|
||||||
|
if op_result.shape[1] > base_segment_result.shape[1]:
|
||||||
|
print('[WARN] ArithAction: op_result.shape[1] > base_segment_result.shape[1] - averaging op_result')
|
||||||
|
op_result = torch.mean(op_result, dim=1, keepdim=True)
|
||||||
|
|
||||||
|
if op_key == "+":
|
||||||
|
base_segment_result.add(op_result)
|
||||||
|
elif op_key == "-":
|
||||||
|
base_segment_result.subtract(op_result)
|
||||||
|
|
||||||
|
return base_segment_result
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def parse_segment(
|
||||||
|
cls,
|
||||||
|
tokens: list[str],
|
||||||
|
start_chars: list[str],
|
||||||
|
end_chars: list[str],
|
||||||
|
parent_parser: Callable[[list[str], SD1Tokenizer], Union[PromptSegment, 'Action']],
|
||||||
|
tokenizer: SD1Tokenizer,
|
||||||
|
) -> Action:
|
||||||
|
"""
|
||||||
|
Parse an arithmetic action from a list of tokens
|
||||||
|
Supported formats:
|
||||||
|
<base_segment:+op1-op2-op3>
|
||||||
|
|
||||||
|
:param tokens: List of tokens, will be modified
|
||||||
|
:param start_chars: List of start chars for all actions
|
||||||
|
:param end_chars: List of end chars for all actions
|
||||||
|
:param parent_parser: Function to parse segments to allow for nested actions
|
||||||
|
:return:
|
||||||
|
"""
|
||||||
|
token = tokens.pop(0)
|
||||||
|
assert token == cls.START_CHAR, "ArithAction must start with " + cls.START_CHAR + " but got " + token
|
||||||
|
|
||||||
|
# Parse base segment
|
||||||
|
base_segment = parent_parser(tokens, tokenizer)
|
||||||
|
|
||||||
|
token = tokens.pop(0)
|
||||||
|
assert token == ":", "ArithAction must have a ':' after the base segment" + " but got " + token
|
||||||
|
|
||||||
|
# Parse ops string
|
||||||
|
ops = {'+': [], '-': []}
|
||||||
|
while tokens[0] != cls.END_CHAR:
|
||||||
|
op_char = tokens.pop(0)
|
||||||
|
assert op_char in ["+", "-"], "ArithAction must have a '+' or '-' as an op char but got " + op_char
|
||||||
|
ops[op_char].append(parent_parser(tokens, tokenizer))
|
||||||
|
|
||||||
|
token = tokens.pop(0)
|
||||||
|
assert token == cls.END_CHAR, "ArithAction must end with " + cls.END_CHAR + " but got " + token
|
||||||
|
|
||||||
|
return cls(base_segment, ops)
|
||||||
@@ -1,49 +0,0 @@
|
|||||||
from typing import List
|
|
||||||
|
|
||||||
from torch import Tensor
|
|
||||||
from torch.nn import Embedding
|
|
||||||
|
|
||||||
from .action_utils import concat_embeddings, get_total_length, parse_numeric_arg
|
|
||||||
from .base import MultiArgAction
|
|
||||||
from .types import SegOrAction
|
|
||||||
|
|
||||||
|
|
||||||
class AverageAction(MultiArgAction):
|
|
||||||
grammar = 'avg(" arg "|" arg "|" arg ")"'
|
|
||||||
|
|
||||||
display_name = "Average"
|
|
||||||
action_name = "avg"
|
|
||||||
description = "Performs a weighted average between two segments or actions. The recommended weight is 0 - 1."
|
|
||||||
usage_examples = [
|
|
||||||
"avg(The cat is|The dog is|0.5)",
|
|
||||||
"avg(Cat|Dog|0.5)",
|
|
||||||
]
|
|
||||||
|
|
||||||
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
|
||||||
super().__init__(args)
|
|
||||||
if len(args) != 3:
|
|
||||||
raise ValueError("Average action expects exactly three arguments (2 vectors and a weight)")
|
|
||||||
|
|
||||||
self.first_arg = args[0]
|
|
||||||
self.second_arg = args[1]
|
|
||||||
self.parsed_weight = parse_numeric_arg(
|
|
||||||
args[2], action_name="Average", role="weight", cast=float
|
|
||||||
)
|
|
||||||
|
|
||||||
first_len = get_total_length(self.first_arg)
|
|
||||||
second_len = get_total_length(self.second_arg)
|
|
||||||
if first_len != second_len:
|
|
||||||
raise ValueError(
|
|
||||||
f"Average start and end arguments should have the same length. Got {first_len} and {second_len}"
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.parsed_weight < 0 or self.parsed_weight > 1:
|
|
||||||
print(f"WARNING: Average weight should be between 0 and 1. Got {self.parsed_weight}")
|
|
||||||
|
|
||||||
def token_length(self) -> int:
|
|
||||||
return get_total_length(self.first_arg)
|
|
||||||
|
|
||||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
|
||||||
start = concat_embeddings(self.first_arg, embedding_module)
|
|
||||||
end = concat_embeddings(self.second_arg, embedding_module)
|
|
||||||
return start * (1 - self.parsed_weight) + end * self.parsed_weight
|
|
||||||
+73
-51
@@ -1,73 +1,95 @@
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from dataclasses import dataclass
|
from typing import Callable, Union
|
||||||
from enum import Enum
|
|
||||||
from typing import List, Optional, Tuple, Union
|
|
||||||
|
|
||||||
|
import torch
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
from torch.nn import Embedding
|
from torch.nn import Embedding
|
||||||
|
|
||||||
|
from comfy.sd1_clip import SD1Tokenizer
|
||||||
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):
|
class Action(ABC):
|
||||||
arity: ActionArity = ActionArity.NONE
|
@property
|
||||||
display_name: str = ""
|
@abstractmethod
|
||||||
action_name: str = ""
|
def START_CHAR(self):
|
||||||
description: str = ""
|
pass
|
||||||
grammar: str = ""
|
|
||||||
usage_examples: List[str] = []
|
@property
|
||||||
|
@abstractmethod
|
||||||
|
def END_CHAR(self):
|
||||||
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def __init__(self, *args, **kwargs) -> None: ...
|
def token_length(self):
|
||||||
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def token_length(self) -> int: ...
|
def get_all_segments(self):
|
||||||
|
pass
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def get_result(self, embedding_module: Embedding) -> ActionResult: ...
|
def get_result(self, embedding_module: Embedding):
|
||||||
|
pass
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
@abstractmethod
|
||||||
|
def parse_segment(
|
||||||
|
cls,
|
||||||
|
tokens: list[str],
|
||||||
|
start_chars: list[str],
|
||||||
|
end_chars: list[str],
|
||||||
|
parent_parser: Callable[[list[str], SD1Tokenizer], Union[str, 'Action']],
|
||||||
|
tokenizer: SD1Tokenizer,
|
||||||
|
) -> 'Action':
|
||||||
|
pass
|
||||||
|
|
||||||
class SingleArgAction(Action, ABC):
|
def depth_repr(self, depth=1):
|
||||||
arity = ActionArity.SINGLE
|
raise NotImplementedError()
|
||||||
|
|
||||||
def __init__(self, arg: List):
|
class PromptSegment:
|
||||||
self.arg = arg
|
def __init__(self, text: str, tokens: list[Union[int, Tensor]]):
|
||||||
|
self.text = text
|
||||||
|
self.tokens = tokens
|
||||||
|
|
||||||
def __repr__(self) -> str:
|
def token_length(self):
|
||||||
return f"{self.action_name}({self.arg})"
|
return len(self.tokens)
|
||||||
|
|
||||||
|
def get_embeddings(self, embedding_module: Embedding):
|
||||||
|
tensors = torch.LongTensor(self.tokens).to(torch.device('cpu'))
|
||||||
|
unsqueezed_tensors = tensors.unsqueeze(0)
|
||||||
|
return embedding_module(unsqueezed_tensors)
|
||||||
|
|
||||||
class MultiArgAction(Action, ABC):
|
def depth_repr(self, depth=1):
|
||||||
arity = ActionArity.MULTI
|
out = f'"{self.text}"('
|
||||||
|
|
||||||
def __init__(self, args: List[List]):
|
cleaned_tokens = list(map(lambda x: str(x) if isinstance(x, int) else "EMBD", self.tokens))
|
||||||
self.all_args = args
|
out += ", ".join(cleaned_tokens)
|
||||||
|
|
||||||
def __repr__(self) -> str:
|
out += ")"
|
||||||
joined = " | ".join(str(a) for a in self.all_args)
|
return out
|
||||||
return f"{self.action_name}({joined})"
|
|
||||||
|
def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment:
|
||||||
|
split_text = text.split(" ")
|
||||||
|
tokens = []
|
||||||
|
for word in split_text:
|
||||||
|
if word.startswith(tokenizer.embedding_identifier) and tokenizer.embedding_directory is not None:
|
||||||
|
embedding_name = word[len(tokenizer.embedding_identifier):].strip('\n')
|
||||||
|
|
||||||
|
get_embed_ret = tokenizer._try_get_embedding(embedding_name)
|
||||||
|
embedding = get_embed_ret[0]
|
||||||
|
leftover = get_embed_ret[1]
|
||||||
|
if embedding is None:
|
||||||
|
print(f"warning, embedding:{embedding_name} does not exist, ignoring")
|
||||||
|
else:
|
||||||
|
if len(embedding.shape) == 1:
|
||||||
|
tokens.append(embedding)
|
||||||
|
else:
|
||||||
|
tokens.extend(embedding)
|
||||||
|
|
||||||
|
if leftover != "":
|
||||||
|
word = leftover
|
||||||
|
else:
|
||||||
|
continue
|
||||||
|
tokens.extend(tokenizer.tokenizer(word)["input_ids"][1:-1])
|
||||||
|
|
||||||
|
return PromptSegment(text, tokens)
|
||||||
|
|||||||
@@ -1,39 +0,0 @@
|
|||||||
from typing import List
|
|
||||||
|
|
||||||
from torch import Tensor
|
|
||||||
from torch.nn import Embedding
|
|
||||||
|
|
||||||
from .action_utils import add_with_broadcast, concat_embeddings
|
|
||||||
from .base import MultiArgAction
|
|
||||||
from .types import SegOrAction
|
|
||||||
|
|
||||||
|
|
||||||
class DiffAction(MultiArgAction):
|
|
||||||
grammar = 'diff(" arg ("|" arg)* ")"'
|
|
||||||
|
|
||||||
display_name = "Difference"
|
|
||||||
action_name = "diff"
|
|
||||||
description = (
|
|
||||||
"Subtracts the segments in the order they are given. "
|
|
||||||
"The first segment is subtracted from the second, then the third from the result, and so on."
|
|
||||||
)
|
|
||||||
usage_examples = [
|
|
||||||
"diff(The cat is|The dog is)",
|
|
||||||
"diff(Cat|Dog)",
|
|
||||||
"sum(diff(king|man)|woman)",
|
|
||||||
]
|
|
||||||
|
|
||||||
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
|
||||||
super().__init__(args)
|
|
||||||
self.base_arg = args[0]
|
|
||||||
self.additional_args = args[1:]
|
|
||||||
|
|
||||||
def token_length(self) -> int:
|
|
||||||
return sum(s.token_length() for s in self.base_arg)
|
|
||||||
|
|
||||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
|
||||||
result = concat_embeddings(self.base_arg, embedding_module)
|
|
||||||
for arg in self.additional_args:
|
|
||||||
arg_embedding = concat_embeddings(arg, embedding_module)
|
|
||||||
result = add_with_broadcast(result, arg_embedding, op="sub")
|
|
||||||
return result
|
|
||||||
@@ -0,0 +1,17 @@
|
|||||||
|
from custom_nodes.ClipStuff.lib.actions import ALL_START_CHARS, ALL_END_CHARS
|
||||||
|
from custom_nodes.ClipStuff.lib.actions.base import Action
|
||||||
|
|
||||||
|
|
||||||
|
def is_action_segment(action_class: Action.__class__, segment: str):
|
||||||
|
if not issubclass(action_class, Action):
|
||||||
|
raise Exception(
|
||||||
|
f"action_class must be a subclass of Action, got {action_class}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return (
|
||||||
|
segment[0] == action_class.START_CHAR and segment[-1] == action_class.END_CHAR
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def is_any_action_segment(segment: str):
|
||||||
|
return segment[0] in ALL_START_CHARS and segment[-1] in ALL_END_CHARS
|
||||||
@@ -1,36 +0,0 @@
|
|||||||
from typing import List
|
|
||||||
|
|
||||||
from torch import Tensor
|
|
||||||
from torch.nn import Embedding
|
|
||||||
|
|
||||||
from .action_utils import concat_embeddings, get_total_length, parse_numeric_arg
|
|
||||||
from .base import MultiArgAction
|
|
||||||
from .types import SegOrAction
|
|
||||||
|
|
||||||
|
|
||||||
class MultiplyAction(MultiArgAction):
|
|
||||||
grammar = 'mult(" arg+ ")"'
|
|
||||||
|
|
||||||
display_name = "Multiply"
|
|
||||||
action_name = "mult"
|
|
||||||
description = "Multiplies the provided segments or actions by the multiplier."
|
|
||||||
usage_examples = [
|
|
||||||
"mult(The cat is|2.5)",
|
|
||||||
"mult(Cat|-1)",
|
|
||||||
]
|
|
||||||
|
|
||||||
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
|
||||||
super().__init__(args)
|
|
||||||
if len(args) != 2:
|
|
||||||
raise ValueError("Multiply action expects exactly two arguments")
|
|
||||||
|
|
||||||
self.target_arg = args[0]
|
|
||||||
self.parsed_multiplier = parse_numeric_arg(
|
|
||||||
args[1], action_name="Multiply", role="multiplier", cast=float
|
|
||||||
)
|
|
||||||
|
|
||||||
def token_length(self) -> int:
|
|
||||||
return get_total_length(self.target_arg)
|
|
||||||
|
|
||||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
|
||||||
return concat_embeddings(self.target_arg, embedding_module) * self.parsed_multiplier
|
|
||||||
@@ -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)
|
|
||||||
@@ -1,23 +0,0 @@
|
|||||||
from torch import Tensor
|
|
||||||
from torch.nn import Embedding
|
|
||||||
|
|
||||||
from .action_utils import concat_embeddings, get_total_length
|
|
||||||
from .base import SingleArgAction
|
|
||||||
|
|
||||||
|
|
||||||
class NegAction(SingleArgAction):
|
|
||||||
grammar = 'neg(" arg+ ")"'
|
|
||||||
|
|
||||||
display_name = "Negate"
|
|
||||||
action_name = "neg"
|
|
||||||
description = "Negates the provided segments or actions."
|
|
||||||
usage_examples = [
|
|
||||||
"neg(cat)",
|
|
||||||
"sum(king|neg(man)|women)",
|
|
||||||
]
|
|
||||||
|
|
||||||
def token_length(self) -> int:
|
|
||||||
return get_total_length(self.arg)
|
|
||||||
|
|
||||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
|
||||||
return concat_embeddings(self.arg, embedding_module) * -1
|
|
||||||
@@ -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
|
|
||||||
@@ -1,25 +0,0 @@
|
|||||||
import torch
|
|
||||||
from torch import Tensor
|
|
||||||
from torch.nn import Embedding
|
|
||||||
|
|
||||||
from .action_utils import concat_embeddings, get_total_length
|
|
||||||
from .base import SingleArgAction
|
|
||||||
|
|
||||||
|
|
||||||
class NormAction(SingleArgAction):
|
|
||||||
grammar = 'norm(" arg+ ")"'
|
|
||||||
|
|
||||||
display_name = "Normalize"
|
|
||||||
action_name = "norm"
|
|
||||||
description = "Normalizes the provided segments or actions."
|
|
||||||
usage_examples = [
|
|
||||||
"norm(cat)",
|
|
||||||
"sum(cat|norm(sum(tiger|fish)))",
|
|
||||||
]
|
|
||||||
|
|
||||||
def token_length(self) -> int:
|
|
||||||
return get_total_length(self.arg)
|
|
||||||
|
|
||||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
|
||||||
embeddings = concat_embeddings(self.arg, embedding_module)
|
|
||||||
return torch.div(embeddings, torch.norm(embeddings, dim=-1, keepdim=True))
|
|
||||||
@@ -0,0 +1,128 @@
|
|||||||
|
from typing import Optional, Union, Callable
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch.nn import Embedding
|
||||||
|
|
||||||
|
from comfy.sd1_clip import SD1Tokenizer
|
||||||
|
from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment
|
||||||
|
|
||||||
|
|
||||||
|
class NudgeAction(Action):
|
||||||
|
def token_length(self):
|
||||||
|
# Nudge nudges the embeddings of the base segment, so the length is the length of the base segment
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
return self.base_segment.token_length()
|
||||||
|
|
||||||
|
return len(self.base_segment.tokens)
|
||||||
|
|
||||||
|
def get_all_segments(self):
|
||||||
|
segments = []
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
segments += self.base_segment.get_all_segments()
|
||||||
|
else:
|
||||||
|
segments.append(self.base_segment)
|
||||||
|
|
||||||
|
if isinstance(self.target, Action):
|
||||||
|
segments += self.target.get_all_segments()
|
||||||
|
else:
|
||||||
|
segments.append(self.target)
|
||||||
|
|
||||||
|
return segments
|
||||||
|
|
||||||
|
def get_result(self, embedding_module: Embedding):
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
base_segment_result = self.base_segment.get_result(embedding_module)
|
||||||
|
else:
|
||||||
|
base_segment_result = self.base_segment.get_embeddings(embedding_module)
|
||||||
|
|
||||||
|
if isinstance(self.target, Action):
|
||||||
|
target_segment_result = self.target.get_result(embedding_module)
|
||||||
|
else:
|
||||||
|
target_segment_result = self.target.get_embeddings(embedding_module)
|
||||||
|
|
||||||
|
base_mean = torch.mean(base_segment_result, dim=1, keepdim=True)
|
||||||
|
if target_segment_result.shape[1] == 1:
|
||||||
|
translation_vector = target_segment_result - base_mean
|
||||||
|
else:
|
||||||
|
translation_vector = torch.mean(target_segment_result, dim=1, keepdim=True) - base_mean
|
||||||
|
|
||||||
|
return base_segment_result.add(translation_vector, alpha=self.weight)
|
||||||
|
|
||||||
|
START_CHAR = "["
|
||||||
|
END_CHAR = "]"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
base_segment: PromptSegment | Action,
|
||||||
|
target: Union[PromptSegment, Action],
|
||||||
|
weight: Optional[float] = None,
|
||||||
|
):
|
||||||
|
self.base_segment = base_segment
|
||||||
|
self.weight = weight
|
||||||
|
self.target = target
|
||||||
|
|
||||||
|
def __repr__(self):
|
||||||
|
return f"NudgeAction(\n\tbase_segment={self.base_segment},\n\ttarget={self.target},\n\tweight={self.weight}\n)"
|
||||||
|
|
||||||
|
def depth_repr(self, depth=1):
|
||||||
|
out = "NudgeAction(\n"
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
base_segment_repr = self.base_segment.depth_repr(depth + 1)
|
||||||
|
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
|
||||||
|
else:
|
||||||
|
out += "\t" * depth + f'base_segment={self.base_segment.depth_repr()},\n'
|
||||||
|
|
||||||
|
if isinstance(self.target, Action):
|
||||||
|
target_repr = self.target.depth_repr(depth + 1)
|
||||||
|
out += "\t" * depth + f"target={target_repr},\n"
|
||||||
|
else:
|
||||||
|
out += "\t" * depth + f"target={self.target.depth_repr()},\n"
|
||||||
|
out += "\t" * depth + f"weight={self.weight},\n"
|
||||||
|
out += "\t" * (depth - 1) + ")"
|
||||||
|
return out
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def parse_segment(
|
||||||
|
cls,
|
||||||
|
tokens: list[str],
|
||||||
|
start_chars: list[str],
|
||||||
|
end_chars: list[str],
|
||||||
|
parent_parser: Callable[[list[str], SD1Tokenizer], PromptSegment | Action],
|
||||||
|
tokenizer: SD1Tokenizer,
|
||||||
|
) -> Action:
|
||||||
|
"""
|
||||||
|
Parse a nudge action from a list of tokens
|
||||||
|
Supported formats:
|
||||||
|
[base_segment:target_segment]
|
||||||
|
[base_segment:target_segment:weight]
|
||||||
|
|
||||||
|
Weight is optional, if not provided it will be None
|
||||||
|
:param tokens: List of tokens, will be modified
|
||||||
|
:param start_chars: List of start chars for all actions
|
||||||
|
:param end_chars: List of end chars for all actions
|
||||||
|
:param parent_parser: Function to parse segments to allow for nested actions
|
||||||
|
:return:
|
||||||
|
"""
|
||||||
|
token = tokens.pop(0)
|
||||||
|
assert token == cls.START_CHAR, "NudgeAction must start with " + cls.START_CHAR + " got " + token
|
||||||
|
|
||||||
|
# Parse base segment
|
||||||
|
base_segment = parent_parser(tokens, tokenizer)
|
||||||
|
|
||||||
|
token = tokens.pop(0)
|
||||||
|
assert token == ":", "NudgeAction must have a ':' after the base segment" + " but got " + token
|
||||||
|
|
||||||
|
# Parse target segment
|
||||||
|
target_segment = parent_parser(tokens, tokenizer)
|
||||||
|
|
||||||
|
# Parse weight if it exists
|
||||||
|
weight = None
|
||||||
|
if tokens[0] == ":":
|
||||||
|
# Parse weight
|
||||||
|
tokens.pop(0)
|
||||||
|
weight = float(tokens.pop(0))
|
||||||
|
|
||||||
|
token = tokens.pop(0)
|
||||||
|
assert token == cls.END_CHAR, "NudgeAction must end with " + cls.END_CHAR + " got " + token
|
||||||
|
|
||||||
|
return cls(base_segment, target_segment, weight)
|
||||||
@@ -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)
|
|
||||||
@@ -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)
|
|
||||||
@@ -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))
|
|
||||||
@@ -1,53 +0,0 @@
|
|||||||
from typing import List
|
|
||||||
|
|
||||||
import torch
|
|
||||||
from torch import Tensor
|
|
||||||
from torch.nn import Embedding
|
|
||||||
|
|
||||||
from .action_utils import parse_numeric_arg
|
|
||||||
from .base import MultiArgAction
|
|
||||||
from .types import SegOrAction
|
|
||||||
|
|
||||||
|
|
||||||
class RandAction(MultiArgAction):
|
|
||||||
grammar = 'rand(" arg ")"'
|
|
||||||
|
|
||||||
display_name = "Random Embedding"
|
|
||||||
action_name = "rand"
|
|
||||||
description = (
|
|
||||||
"Returns a random embedding of the specified token length, "
|
|
||||||
"with the values optionally bounded by the second and third arguments."
|
|
||||||
)
|
|
||||||
usage_examples = [
|
|
||||||
"A rand(1) cat",
|
|
||||||
"A rand(1|-1|1) cat",
|
|
||||||
]
|
|
||||||
|
|
||||||
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
|
||||||
super().__init__(args)
|
|
||||||
if len(args) not in (1, 3):
|
|
||||||
raise ValueError("Random action expects exactly one or three arguments")
|
|
||||||
|
|
||||||
self.parsed_token_length = parse_numeric_arg(
|
|
||||||
args[0], action_name="Random", role="first argument (token length)", cast=int
|
|
||||||
)
|
|
||||||
if len(args) == 3:
|
|
||||||
self.range_min = parse_numeric_arg(
|
|
||||||
args[1], action_name="Random", role="second argument (min)", cast=int
|
|
||||||
)
|
|
||||||
self.range_max = parse_numeric_arg(
|
|
||||||
args[2], action_name="Random", role="third argument (max)", cast=int
|
|
||||||
)
|
|
||||||
if self.range_min > self.range_max:
|
|
||||||
raise ValueError("Random action min must be <= max")
|
|
||||||
else:
|
|
||||||
self.range_min = 0
|
|
||||||
self.range_max = 1
|
|
||||||
|
|
||||||
def token_length(self) -> int:
|
|
||||||
return self.parsed_token_length
|
|
||||||
|
|
||||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
|
||||||
return torch.empty(
|
|
||||||
1, self.parsed_token_length, embedding_module.embedding_dim
|
|
||||||
).uniform_(self.range_min, self.range_max)
|
|
||||||
@@ -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)
|
|
||||||
@@ -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
|
|
||||||
@@ -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
|
|
||||||
@@ -1,52 +0,0 @@
|
|||||||
from typing import List
|
|
||||||
|
|
||||||
from torch import Tensor
|
|
||||||
from torch.nn import Embedding
|
|
||||||
|
|
||||||
from .action_utils import concat_embeddings, get_total_length, parse_numeric_arg
|
|
||||||
from .base import MultiArgAction
|
|
||||||
from .types import SegOrAction
|
|
||||||
from .utils import slerp
|
|
||||||
|
|
||||||
|
|
||||||
class SlerpAction(MultiArgAction):
|
|
||||||
grammar = 'slerp(" arg "|" arg "|" arg ")"'
|
|
||||||
|
|
||||||
display_name = "Slerp"
|
|
||||||
action_name = "slerp"
|
|
||||||
description = (
|
|
||||||
"Performs a slerp (interpolation) between two segments or actions, with the given weight. "
|
|
||||||
"The recommended weight is 0 - 1."
|
|
||||||
)
|
|
||||||
usage_examples = [
|
|
||||||
"The slerp(cat|dog|0.5) is happy",
|
|
||||||
]
|
|
||||||
|
|
||||||
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
|
||||||
super().__init__(args)
|
|
||||||
if len(args) != 3:
|
|
||||||
raise ValueError("Slerp action expects exactly three arguments (2 vectors and a weight)")
|
|
||||||
|
|
||||||
self.start_argument = args[0]
|
|
||||||
self.end_argument = args[1]
|
|
||||||
self.parsed_weight = parse_numeric_arg(
|
|
||||||
args[2], action_name="Slerp", role="weight", cast=float
|
|
||||||
)
|
|
||||||
|
|
||||||
start_len = get_total_length(self.start_argument)
|
|
||||||
end_len = get_total_length(self.end_argument)
|
|
||||||
if start_len != end_len:
|
|
||||||
raise ValueError(
|
|
||||||
f"Slerp start and end arguments should have the same length. Got {start_len} and {end_len}"
|
|
||||||
)
|
|
||||||
|
|
||||||
if self.parsed_weight < 0 or self.parsed_weight > 1:
|
|
||||||
print(f"WARNING: Slerp weight should be between 0 and 1. Got {self.parsed_weight}")
|
|
||||||
|
|
||||||
def token_length(self) -> int:
|
|
||||||
return get_total_length(self.start_argument)
|
|
||||||
|
|
||||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
|
||||||
start = concat_embeddings(self.start_argument, embedding_module)
|
|
||||||
end = concat_embeddings(self.end_argument, embedding_module)
|
|
||||||
return slerp(self.parsed_weight, start, end)
|
|
||||||
@@ -1,34 +0,0 @@
|
|||||||
from typing import List
|
|
||||||
|
|
||||||
from torch import Tensor
|
|
||||||
from torch.nn import Embedding
|
|
||||||
|
|
||||||
from .action_utils import add_with_broadcast, concat_embeddings
|
|
||||||
from .base import MultiArgAction
|
|
||||||
from .types import SegOrAction
|
|
||||||
|
|
||||||
|
|
||||||
class SumAction(MultiArgAction):
|
|
||||||
grammar = 'sum(" arg ("|" arg)+ ")"'
|
|
||||||
|
|
||||||
display_name = "Sum"
|
|
||||||
action_name = "sum"
|
|
||||||
description = "Adds the embeddings of the provided segments or actions."
|
|
||||||
usage_examples = [
|
|
||||||
"A happy sum(cat|dog|shark)",
|
|
||||||
]
|
|
||||||
|
|
||||||
def __init__(self, args: List[List[SegOrAction]]) -> None:
|
|
||||||
super().__init__(args)
|
|
||||||
self.base_arg = args[0]
|
|
||||||
self.additional_args = args[1:]
|
|
||||||
|
|
||||||
def token_length(self) -> int:
|
|
||||||
return sum(s.token_length() for s in self.base_arg)
|
|
||||||
|
|
||||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
|
||||||
result = concat_embeddings(self.base_arg, embedding_module)
|
|
||||||
for arg in self.additional_args:
|
|
||||||
arg_embedding = concat_embeddings(arg, embedding_module)
|
|
||||||
result = add_with_broadcast(result, arg_embedding, op="add")
|
|
||||||
return result
|
|
||||||
@@ -1,7 +0,0 @@
|
|||||||
from typing import Union
|
|
||||||
|
|
||||||
from ..parser.prompt_segment import PromptSegment
|
|
||||||
from .base import Action
|
|
||||||
from .weighted import WeightedGroup
|
|
||||||
|
|
||||||
SegOrAction = Union[PromptSegment, Action, WeightedGroup]
|
|
||||||
@@ -1,23 +0,0 @@
|
|||||||
import torch
|
|
||||||
|
|
||||||
|
|
||||||
def slerp(val: float, low: torch.Tensor, high: torch.Tensor, epsilon: float = 1e-5) -> torch.Tensor:
|
|
||||||
"""Spherical linear interpolation between two tensors along the last dim."""
|
|
||||||
val_t = torch.tensor(val, dtype=torch.float32, device=low.device).clamp(0, 1)
|
|
||||||
|
|
||||||
low_norm = low / torch.norm(low, dim=-1, keepdim=True)
|
|
||||||
high_norm = high / torch.norm(high, dim=-1, keepdim=True)
|
|
||||||
|
|
||||||
dot = (low_norm * high_norm).sum(-1, keepdim=True).clamp(-1, 1)
|
|
||||||
omega = torch.acos(dot)
|
|
||||||
sin_omega = torch.sin(omega)
|
|
||||||
|
|
||||||
scale_low = torch.sin((1.0 - val_t) * omega) / (sin_omega + epsilon)
|
|
||||||
scale_high = torch.sin(val_t * omega) / (sin_omega + epsilon)
|
|
||||||
|
|
||||||
# Fall back to linear interp where the angle is too small for stable slerp.
|
|
||||||
close = sin_omega < epsilon
|
|
||||||
scale_low = torch.where(close, 1.0 - val_t, scale_low)
|
|
||||||
scale_high = torch.where(close, val_t, scale_high)
|
|
||||||
|
|
||||||
return scale_low * low + scale_high * high
|
|
||||||
@@ -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})"
|
|
||||||
@@ -0,0 +1,25 @@
|
|||||||
|
{
|
||||||
|
"_name_or_path": "openai/clip-vit-large-patch14",
|
||||||
|
"architectures": [
|
||||||
|
"CLIPTextModel"
|
||||||
|
],
|
||||||
|
"attention_dropout": 0.0,
|
||||||
|
"bos_token_id": 0,
|
||||||
|
"dropout": 0.0,
|
||||||
|
"eos_token_id": 2,
|
||||||
|
"hidden_act": "quick_gelu",
|
||||||
|
"hidden_size": 768,
|
||||||
|
"initializer_factor": 1.0,
|
||||||
|
"initializer_range": 0.02,
|
||||||
|
"intermediate_size": 3072,
|
||||||
|
"layer_norm_eps": 1e-05,
|
||||||
|
"max_position_embeddings": 77,
|
||||||
|
"model_type": "clip_text_model",
|
||||||
|
"num_attention_heads": 12,
|
||||||
|
"num_hidden_layers": 12,
|
||||||
|
"pad_token_id": 1,
|
||||||
|
"projection_dim": 768,
|
||||||
|
"torch_dtype": "float32",
|
||||||
|
"transformers_version": "4.24.0",
|
||||||
|
"vocab_size": 49408
|
||||||
|
}
|
||||||
+175
-104
@@ -1,129 +1,200 @@
|
|||||||
"""DSL-aware CLIP text encoders.
|
import contextlib
|
||||||
|
import os
|
||||||
The tokenizer emits ComfyUI's native `(token, weight)` format with one twist:
|
from typing import Union
|
||||||
a `token` can also be a lazily-evaluated `Action`. We resolve those to tensors
|
|
||||||
here in `process_tokens` (where the embedding module is available) and delegate
|
|
||||||
everything else — embedding lookup, mask building, splice — to the stock
|
|
||||||
`SDClipModel.process_tokens`.
|
|
||||||
|
|
||||||
`posScale` / `postPos` actions return a `PostModifiers` alongside their tensor.
|
|
||||||
ComfyUI's `CLIPTextModel_.forward` adds the position embedding inline whenever
|
|
||||||
`embeds` is supplied, so we pre-bake `(modified - default)` into `embeds` such
|
|
||||||
that the transformer's add nets to `+ modified`.
|
|
||||||
"""
|
|
||||||
|
|
||||||
import dataclasses
|
|
||||||
from typing import List
|
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import Tensor
|
from transformers import CLIPTextConfig, modeling_utils
|
||||||
from torch.nn import Embedding
|
|
||||||
|
|
||||||
from comfy import sd1_clip, sdxl_clip
|
from comfy import model_management
|
||||||
|
import comfy.ops
|
||||||
|
from comfy.sd import CLIP
|
||||||
|
from custom_nodes.ClipStuff.lib.actions.base import PromptSegment, Action
|
||||||
|
from custom_nodes.ClipStuff.lib.fun_clip_stuff import MyCLIPTextModel
|
||||||
|
from custom_nodes.ClipStuff.lib.tokenizer import TokenDict
|
||||||
|
|
||||||
from .actions.base import ACTION_CONTINUATION, Action, PostModifiers
|
class SD1FunClipModel(torch.nn.Module):
|
||||||
|
"""Uses the CLIP transformer encoder for text (from huggingface)"""
|
||||||
|
LAYERS = [
|
||||||
|
"last",
|
||||||
|
"pooled",
|
||||||
|
"hidden"
|
||||||
|
]
|
||||||
|
|
||||||
|
def __init__(self, version="openai/clip-vit-large-patch14", device="cpu", max_length=77,
|
||||||
|
freeze=True, layer="last", layer_idx=None, textmodel_json_config=None,
|
||||||
|
textmodel_path=None): # clip-vit-base-patch32
|
||||||
|
super().__init__()
|
||||||
|
assert layer in self.LAYERS
|
||||||
|
self.num_layers = 12
|
||||||
|
if textmodel_path is not None:
|
||||||
|
self.transformer = MyCLIPTextModel.from_pretrained(textmodel_path)
|
||||||
|
else:
|
||||||
|
if textmodel_json_config is None:
|
||||||
|
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "clip_config.json")
|
||||||
|
config = CLIPTextConfig.from_json_file(textmodel_json_config)
|
||||||
|
self.num_layers = config.num_hidden_layers
|
||||||
|
with comfy.ops.use_comfy_ops():
|
||||||
|
with modeling_utils.no_init_weights():
|
||||||
|
self.transformer = MyCLIPTextModel(config)
|
||||||
|
|
||||||
class PromptLangSDClipModel(sd1_clip.SDClipModel):
|
self.max_length = max_length
|
||||||
def process_tokens(self, tokens, device): # type: ignore[override]
|
if freeze:
|
||||||
embedding_module = self.transformer.get_input_embeddings()
|
self.freeze()
|
||||||
|
self.layer = layer
|
||||||
|
self.layer_idx = None
|
||||||
|
self.empty_tokens = [[49406] + [49407] * 76]
|
||||||
|
self.text_projection = None
|
||||||
|
self.layer_norm_hidden_state = True
|
||||||
|
if layer == "hidden":
|
||||||
|
assert layer_idx is not None
|
||||||
|
assert abs(layer_idx) <= self.num_layers
|
||||||
|
self.clip_layer(layer_idx)
|
||||||
|
self.layer_default = (self.layer, self.layer_idx)
|
||||||
|
|
||||||
|
def freeze(self):
|
||||||
|
self.transformer = self.transformer.eval()
|
||||||
|
# self.train = disabled_train
|
||||||
|
for param in self.parameters():
|
||||||
|
param.requires_grad = False
|
||||||
|
|
||||||
|
def clip_layer(self, layer_idx):
|
||||||
|
if abs(layer_idx) >= self.num_layers:
|
||||||
|
self.layer = "last"
|
||||||
|
else:
|
||||||
|
self.layer = "hidden"
|
||||||
|
self.layer_idx = layer_idx
|
||||||
|
|
||||||
|
def reset_clip_layer(self):
|
||||||
|
self.layer = self.layer_default[0]
|
||||||
|
self.layer_idx = self.layer_default[1]
|
||||||
|
|
||||||
|
def set_up_textual_embeddings(self, tokens: list[list[PromptSegment | Action]], current_embeds):
|
||||||
|
next_new_token = token_dict_size = current_embeds.weight.shape[0] - 1
|
||||||
|
embedding_weights = []
|
||||||
|
|
||||||
|
# For each batch
|
||||||
|
for batch in tokens:
|
||||||
|
for seg_or_action in batch:
|
||||||
|
if isinstance(seg_or_action, Action):
|
||||||
|
segments = seg_or_action.get_all_segments()
|
||||||
|
else:
|
||||||
|
segments = [seg_or_action]
|
||||||
|
|
||||||
|
for segment in segments:
|
||||||
|
tokens_temp = []
|
||||||
|
segment_length = segment.token_length()
|
||||||
|
for tid_or_tensor in segment.tokens:
|
||||||
|
if isinstance(tid_or_tensor, int):
|
||||||
|
if tid_or_tensor == token_dict_size: # Is EOS token
|
||||||
|
tid_or_tensor = -1 # Set to -1 so that it can be replaced with the EOS token later
|
||||||
|
tokens_temp += [tid_or_tensor]
|
||||||
|
else:
|
||||||
|
if tid_or_tensor.shape[0] == current_embeds.weight.shape[1]:
|
||||||
|
embedding_weights += [tid_or_tensor]
|
||||||
|
tokens_temp += [next_new_token]
|
||||||
|
next_new_token += 1
|
||||||
|
else:
|
||||||
|
print("WARNING: shape mismatch when trying to apply embedding, embedding will be ignored",
|
||||||
|
tid_or_tensor.shape[0], current_embeds.weight.shape[1])
|
||||||
|
if len(tokens_temp) < segment_length:
|
||||||
|
# Pretty sure this is only needed if the embedding is not the same size as the CLIP embedding
|
||||||
|
print("WARNING: segment length mismatch, padding with EOS token")
|
||||||
|
tokens_temp.extend([self.empty_tokens[0][-1] * (segment_length - len(tokens_temp))])
|
||||||
|
segment.tokens = tokens_temp
|
||||||
|
|
||||||
|
n = token_dict_size
|
||||||
|
if len(embedding_weights) > 0:
|
||||||
|
# Create new embedding, with size of current embedding + number of new embeddings
|
||||||
|
new_embedding = torch.nn.Embedding(next_new_token + 1, current_embeds.weight.shape[1],
|
||||||
|
device=current_embeds.weight.device, dtype=current_embeds.weight.dtype)
|
||||||
|
# Copy current embedding weights to new embedding
|
||||||
|
new_embedding.weight[:token_dict_size] = current_embeds.weight[:-1]
|
||||||
|
# Add new embeddings
|
||||||
|
for embed in embedding_weights:
|
||||||
|
new_embedding.weight[n] = embed
|
||||||
|
n += 1
|
||||||
|
|
||||||
|
# Set re-add the EOS token
|
||||||
|
new_embedding.weight[n] = current_embeds.weight[-1] # EOS embedding
|
||||||
|
self.transformer.set_input_embeddings(new_embedding)
|
||||||
|
|
||||||
resolved: List[list] = []
|
|
||||||
pos_modifiers_per_batch: List[List[PostModifiers]] = []
|
|
||||||
|
|
||||||
for batch in tokens:
|
for batch in tokens:
|
||||||
row: list = []
|
for seg_or_action in batch:
|
||||||
modifiers: List[PostModifiers] = []
|
if isinstance(seg_or_action, Action):
|
||||||
position = 0
|
segments = seg_or_action.get_all_segments()
|
||||||
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
|
|
||||||
else:
|
else:
|
||||||
row.append(entry)
|
segments = [seg_or_action]
|
||||||
position += 1
|
|
||||||
resolved.append(row)
|
|
||||||
pos_modifiers_per_batch.append(modifiers)
|
|
||||||
|
|
||||||
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):
|
def forward(self, tokens, **kwargs):
|
||||||
embeds = _apply_pos_modifiers(
|
backup_embeds = self.transformer.get_input_embeddings()
|
||||||
embeds, pos_modifiers_per_batch, self._get_position_embedding()
|
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
|
||||||
def _get_position_embedding(self) -> Embedding:
|
else:
|
||||||
"""Isolated so a ComfyUI internal layout change only needs one fix."""
|
precision_scope = contextlib.nullcontext
|
||||||
return self.transformer.text_model.embeddings.position_embedding
|
|
||||||
|
|
||||||
|
|
||||||
class PromptLangSDXLClipG(sdxl_clip.SDXLClipG, PromptLangSDClipModel):
|
if (kwargs.get("position_ids", None) is not None):
|
||||||
"""SDXL's larger CLIP-G text encoder, with our DSL-aware process_tokens."""
|
position_ids = torch.LongTensor(kwargs["position_ids"]).to(device)
|
||||||
|
else:
|
||||||
|
position_ids = None
|
||||||
|
|
||||||
|
|
||||||
def _apply_pos_modifiers(
|
with precision_scope(model_management.get_autocast_device(device)):
|
||||||
embeds: Tensor,
|
outputs = self.transformer(input_ids=tokens, output_hidden_states=self.layer == "hidden",
|
||||||
pos_modifiers_per_batch: List[List[PostModifiers]],
|
position_ids=position_ids)
|
||||||
position_embedding: Embedding,
|
self.transformer.set_input_embeddings(backup_embeds)
|
||||||
) -> Tensor:
|
|
||||||
seq_len = embeds.shape[1]
|
|
||||||
pos_weights = position_embedding.weight[:seq_len].to(device=embeds.device, dtype=embeds.dtype)
|
|
||||||
|
|
||||||
out = embeds.clone()
|
if self.layer == "last":
|
||||||
for batch_idx, modifiers in enumerate(pos_modifiers_per_batch):
|
z = outputs.last_hidden_state
|
||||||
for mod in modifiers:
|
elif self.layer == "pooled":
|
||||||
default_slice = pos_weights[mod.start_idx:mod.end_idx]
|
z = outputs.pooler_output[:, None, :]
|
||||||
|
|
||||||
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)
|
|
||||||
else:
|
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`.
|
pooled_output = outputs.pooler_output
|
||||||
out[batch_idx, mod.start_idx:mod.end_idx] += modified_slice - default_slice
|
if self.text_projection is not None:
|
||||||
|
pooled_output = pooled_output.to(self.text_projection.device) @ self.text_projection
|
||||||
|
return z.float(), pooled_output.float()
|
||||||
|
|
||||||
return out
|
def encode(self, tokens, **kwargs):
|
||||||
|
return self(tokens, **kwargs)
|
||||||
|
|
||||||
|
def load_sd(self, sd):
|
||||||
|
return self.transformer.load_state_dict(sd, strict=False)
|
||||||
|
|
||||||
class PromptLangSD1ClipModel(sd1_clip.SD1ClipModel):
|
def encode_token_weights(self, prompt_segments: list[list[Union[PromptSegment | Action]]], **kwargs):
|
||||||
def __init__(self, device="cpu", dtype=None, model_options=None, **kwargs):
|
to_encode = [[PromptSegment(text="_Empty Batch_", tokens=self.empty_tokens[0])]]
|
||||||
super().__init__(
|
for batch in prompt_segments:
|
||||||
device=device,
|
to_encode.append(batch)
|
||||||
dtype=dtype,
|
|
||||||
model_options=model_options or {},
|
|
||||||
clip_name="l",
|
|
||||||
clip_model=PromptLangSDClipModel,
|
|
||||||
**kwargs,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
out, pooled = self.encode(to_encode, **kwargs)
|
||||||
|
z_empty = out[0:1]
|
||||||
|
if pooled.shape[0] > 1:
|
||||||
|
first_pooled = pooled[1:2]
|
||||||
|
else:
|
||||||
|
first_pooled = pooled[0:1]
|
||||||
|
|
||||||
class PromptLangSDXLClipModel(sdxl_clip.SDXLClipModel):
|
output = []
|
||||||
def __init__(self, device="cpu", dtype=None, model_options=None) -> None:
|
for k in range(1, out.shape[0]):
|
||||||
torch.nn.Module.__init__(self)
|
z = out[k:k + 1]
|
||||||
opts = model_options or {}
|
# for i in range(len(z)):
|
||||||
self.clip_l = PromptLangSDClipModel(
|
# for j in range(len(z[i])):
|
||||||
layer="hidden",
|
# weight = token_dicts[k - 1][j][0].weight
|
||||||
layer_idx=-2,
|
# z[i][j] = (z[i][j] - z_empty[0][j]) * weight + z_empty[0][j]
|
||||||
device=device,
|
output.append(z)
|
||||||
dtype=dtype,
|
|
||||||
layer_norm_hidden_state=False,
|
if (len(output) == 0):
|
||||||
model_options=opts,
|
return z_empty.cpu(), first_pooled.cpu()
|
||||||
)
|
return torch.cat(output, dim=-2).cpu(), first_pooled.cpu()
|
||||||
self.clip_g = PromptLangSDXLClipG(device=device, dtype=dtype, model_options=opts)
|
|
||||||
self.dtypes = {dtype} if dtype is not None else set()
|
|
||||||
|
|||||||
@@ -0,0 +1,217 @@
|
|||||||
|
from typing import Optional, Tuple, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch import device
|
||||||
|
from transformers import CLIPTextConfig
|
||||||
|
from transformers.modeling_outputs import BaseModelOutputWithPooling
|
||||||
|
from transformers.models.clip.modeling_clip import _expand_mask, CLIPTextEmbeddings, CLIPTextTransformer, \
|
||||||
|
CLIPTextModel
|
||||||
|
|
||||||
|
from custom_nodes.ClipStuff.lib.actions.base import PromptSegment, Action
|
||||||
|
from custom_nodes.ClipStuff.lib.tokenizer import TokenDict
|
||||||
|
|
||||||
|
def slerp(val, low, high):
|
||||||
|
low = low.unsqueeze(0)
|
||||||
|
high = high.unsqueeze(0)
|
||||||
|
low_norm = low/torch.norm(low, dim=1, keepdim=True)
|
||||||
|
high_norm = high/torch.norm(high, dim=1, keepdim=True)
|
||||||
|
omega = torch.acos((low_norm*high_norm).sum(1))
|
||||||
|
so = torch.sin(omega)
|
||||||
|
res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high
|
||||||
|
return res
|
||||||
|
|
||||||
|
class MyCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||||
|
def __init__(self, config: CLIPTextConfig):
|
||||||
|
super().__init__(config)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
input_dicts: Optional[list[list[tuple[TokenDict]]]] = None,
|
||||||
|
input_ids: Optional[torch.LongTensor] = None,
|
||||||
|
position_ids: Optional[torch.LongTensor] = None,
|
||||||
|
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
|
||||||
|
|
||||||
|
batches = []
|
||||||
|
for batch_idx, batch in enumerate(input_dicts):
|
||||||
|
results = []
|
||||||
|
for seg_or_action in batch:
|
||||||
|
if isinstance(seg_or_action, Action):
|
||||||
|
results.append(seg_or_action.get_result(self.token_embedding))
|
||||||
|
else:
|
||||||
|
results.append(seg_or_action.get_embeddings(self.token_embedding))
|
||||||
|
batches.append(results)
|
||||||
|
|
||||||
|
seq_length = batches[0][0].shape[-2]
|
||||||
|
|
||||||
|
if position_ids is None:
|
||||||
|
position_ids = self.position_ids[:, :seq_length]
|
||||||
|
|
||||||
|
# if inputs_embeds is None:
|
||||||
|
# inputs_embeds = self.token_embedding(input_ids)
|
||||||
|
|
||||||
|
# for batch_idx, batch in enumerate(input_dicts):
|
||||||
|
# for token_idx, token in enumerate(batch):
|
||||||
|
# if token[0].nudge_id is not None:
|
||||||
|
# nudged_embed = inputs_embeds[batch_idx, token_idx][:] + self.token_embedding(torch.LongTensor([token[0].nudge_id]).to(torch.device('cpu')))[0]
|
||||||
|
# if token[0].nudge_index_start is not None and token[0].nudge_index_stop is not None:
|
||||||
|
# nudge_start = token[0].nudge_index_start
|
||||||
|
# nudge_end = token[0].nudge_index_stop
|
||||||
|
# else:
|
||||||
|
# nudge_start = 0
|
||||||
|
# nudge_end = 768
|
||||||
|
# inputs_embeds[batch_idx, token_idx][nudge_start:nudge_end] = (slerp(token[0].nudge_weight, inputs_embeds[batch_idx, token_idx][:], nudged_embed)[0][nudge_start:nudge_end])
|
||||||
|
# elif token[0].arith_ops is not None:
|
||||||
|
# for op, id_list in token[0].arith_ops.items():
|
||||||
|
# if op == '+':
|
||||||
|
# for this_id in id_list:
|
||||||
|
# inputs_embeds[batch_idx, token_idx] += self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0]
|
||||||
|
# elif op == '-':
|
||||||
|
# for this_id in id_list:
|
||||||
|
# inputs_embeds[batch_idx, token_idx] -= self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0]
|
||||||
|
|
||||||
|
embeds = []
|
||||||
|
for batch in batches:
|
||||||
|
if len(batch) == 1:
|
||||||
|
embeds.append(batch[0])
|
||||||
|
else:
|
||||||
|
embeds.append(torch.cat(batch, dim=-2))
|
||||||
|
|
||||||
|
position_embeddings = self.position_embedding(position_ids)
|
||||||
|
embeddings = torch.cat(embeds, dim=0) + position_embeddings
|
||||||
|
|
||||||
|
return embeddings
|
||||||
|
|
||||||
|
|
||||||
|
class MyCLIPTextTransformer(CLIPTextTransformer):
|
||||||
|
def __init__(self, config: CLIPTextConfig):
|
||||||
|
super().__init__(config)
|
||||||
|
self.embeddings = MyCLIPTextEmbeddings(config)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
input_ids: Optional[list[list[PromptSegment | Action]]] = None,
|
||||||
|
attention_mask: Optional[torch.Tensor] = None,
|
||||||
|
position_ids: Optional[torch.Tensor] = None,
|
||||||
|
output_attentions: Optional[bool] = None,
|
||||||
|
output_hidden_states: Optional[bool] = None,
|
||||||
|
return_dict: Optional[bool] = None,
|
||||||
|
) -> Union[Tuple, BaseModelOutputWithPooling]:
|
||||||
|
r"""
|
||||||
|
Returns:
|
||||||
|
|
||||||
|
"""
|
||||||
|
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
|
||||||
|
output_hidden_states = (
|
||||||
|
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
|
||||||
|
)
|
||||||
|
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||||
|
|
||||||
|
if input_ids is None:
|
||||||
|
raise ValueError("You have to specify input_ids")
|
||||||
|
|
||||||
|
# input_shape = input_ids.size()
|
||||||
|
# input_ids = input_ids.view(-1, input_shape[-1])
|
||||||
|
|
||||||
|
hidden_states = self.embeddings(input_dicts=input_ids)
|
||||||
|
|
||||||
|
bsz = len(input_ids)
|
||||||
|
# TODO: Properly gather this
|
||||||
|
seq_len = 77
|
||||||
|
# bsz, seq_len = input_shape
|
||||||
|
# CLIP's text model uses causal mask, prepare it here.
|
||||||
|
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
|
||||||
|
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(
|
||||||
|
hidden_states.device
|
||||||
|
)
|
||||||
|
# expand attention_mask
|
||||||
|
if attention_mask is not None:
|
||||||
|
# [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
|
||||||
|
attention_mask = _expand_mask(attention_mask, hidden_states.dtype)
|
||||||
|
|
||||||
|
encoder_outputs = self.encoder(
|
||||||
|
inputs_embeds=hidden_states,
|
||||||
|
attention_mask=attention_mask,
|
||||||
|
causal_attention_mask=causal_attention_mask,
|
||||||
|
output_attentions=output_attentions,
|
||||||
|
output_hidden_states=output_hidden_states,
|
||||||
|
return_dict=return_dict,
|
||||||
|
)
|
||||||
|
|
||||||
|
last_hidden_state = encoder_outputs[0]
|
||||||
|
last_hidden_state = self.final_layer_norm(last_hidden_state)
|
||||||
|
|
||||||
|
|
||||||
|
# Hacky way to get idx of first EOT token
|
||||||
|
eot_idx = [1]
|
||||||
|
for batch in input_ids[1:]:
|
||||||
|
idx = 0
|
||||||
|
for seg_or_action in batch:
|
||||||
|
if isinstance(seg_or_action, Action):
|
||||||
|
idx += seg_or_action.token_length()
|
||||||
|
else:
|
||||||
|
if seg_or_action.text == '__PAD__':
|
||||||
|
break
|
||||||
|
eot_idx.append(idx)
|
||||||
|
# text_embeds.shape = [batch_size, sequence_length, transformer.width]
|
||||||
|
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
||||||
|
# casting to torch.int for onnx compatibility: argmax doesn't support int64 inputs with opset 14
|
||||||
|
# TODO: Get the index of the first EOT token
|
||||||
|
pooled_output = last_hidden_state[
|
||||||
|
torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device),
|
||||||
|
eot_idx
|
||||||
|
]
|
||||||
|
|
||||||
|
if not return_dict:
|
||||||
|
return (last_hidden_state, pooled_output) + encoder_outputs[1:]
|
||||||
|
|
||||||
|
return BaseModelOutputWithPooling(
|
||||||
|
last_hidden_state=last_hidden_state,
|
||||||
|
pooler_output=pooled_output,
|
||||||
|
hidden_states=encoder_outputs.hidden_states,
|
||||||
|
attentions=encoder_outputs.attentions,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class MyCLIPTextModel(CLIPTextModel):
|
||||||
|
def __init__(self, config: CLIPTextConfig):
|
||||||
|
super().__init__(config)
|
||||||
|
self.text_model = MyCLIPTextTransformer(config)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
input_ids: Optional[list[list[tuple[TokenDict]]]] = None,
|
||||||
|
attention_mask: Optional[torch.Tensor] = None,
|
||||||
|
position_ids: Optional[torch.Tensor] = None,
|
||||||
|
output_attentions: Optional[bool] = None,
|
||||||
|
output_hidden_states: Optional[bool] = None,
|
||||||
|
return_dict: Optional[bool] = None,
|
||||||
|
) -> Union[Tuple, BaseModelOutputWithPooling]:
|
||||||
|
r"""
|
||||||
|
Returns:
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
|
||||||
|
```python
|
||||||
|
>>> from transformers import AutoTokenizer, CLIPTextModel
|
||||||
|
|
||||||
|
>>> model = CLIPTextModel.from_pretrained("openai/clip-vit-base-patch32")
|
||||||
|
>>> tokenizer = AutoTokenizer.from_pretrained("openai/clip-vit-base-patch32")
|
||||||
|
|
||||||
|
>>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="pt")
|
||||||
|
|
||||||
|
>>> outputs = model(**inputs)
|
||||||
|
>>> last_hidden_state = outputs.last_hidden_state
|
||||||
|
>>> pooled_output = outputs.pooler_output # pooled (EOS token) states
|
||||||
|
```"""
|
||||||
|
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||||
|
|
||||||
|
return self.text_model(
|
||||||
|
input_ids=input_ids,
|
||||||
|
attention_mask=attention_mask,
|
||||||
|
position_ids=position_ids,
|
||||||
|
output_attentions=output_attentions,
|
||||||
|
output_hidden_states=output_hidden_states,
|
||||||
|
return_dict=return_dict,
|
||||||
|
)
|
||||||
@@ -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)
|
|
||||||
@@ -1,5 +0,0 @@
|
|||||||
from lark import Lark
|
|
||||||
|
|
||||||
from .grammar import grammar
|
|
||||||
|
|
||||||
PromptParser = Lark(grammar, start="start", parser="earley")
|
|
||||||
@@ -1,38 +0,0 @@
|
|||||||
grammar = r"""
|
|
||||||
?start: stmt+
|
|
||||||
|
|
||||||
?stmt: assign
|
|
||||||
| item
|
|
||||||
|
|
||||||
assign: "$" NAME "=" arg ";"
|
|
||||||
|
|
||||||
item: embedding
|
|
||||||
| WORD
|
|
||||||
| generic_function
|
|
||||||
| QUOTED_STRING
|
|
||||||
| weighted
|
|
||||||
| ref
|
|
||||||
|
|
||||||
generic_function: FUNC_NAME "(" arg ("|" arg)* ")"
|
|
||||||
|
|
||||||
weighted: "(" arg ":" SIGNED_NUMBER ")"
|
|
||||||
|
|
||||||
ref: "$" NAME
|
|
||||||
|
|
||||||
arg: item+
|
|
||||||
|
|
||||||
embedding: "embedding:" WORD
|
|
||||||
// NAME and WORD overlap on bare identifiers; the earley parser's dynamic lexer
|
|
||||||
// disambiguates by grammar context (the leading "$" forces NAME). This breaks
|
|
||||||
// under a basic/contextual lexer, so keep parser="earley" in __init__.py.
|
|
||||||
FUNC_NAME: /[A-Za-z_-]+/
|
|
||||||
NAME: /[A-Za-z_][A-Za-z0-9_]*/
|
|
||||||
WORD: /[A-Za-z0-9,_\.-]+/
|
|
||||||
QUOTED_STRING: /"([^"\\]*(\\.[^"\\]*)*)"|'([^'\\]*(\\.[^'\\]*)*)'/
|
|
||||||
SIGNED_NUMBER: /-?\d+(\.\d+)?/
|
|
||||||
COMMENT: /#[^\n]*/
|
|
||||||
|
|
||||||
%import common.WS
|
|
||||||
%ignore WS
|
|
||||||
%ignore COMMENT
|
|
||||||
"""
|
|
||||||
@@ -1,30 +0,0 @@
|
|||||||
from typing import List, Union
|
|
||||||
|
|
||||||
import torch
|
|
||||||
from torch import Tensor
|
|
||||||
from torch.nn import Embedding
|
|
||||||
|
|
||||||
|
|
||||||
class PromptSegment:
|
|
||||||
"""A run of contiguous tokens from the user's prompt, possibly with inline TI tensor entries."""
|
|
||||||
|
|
||||||
def __init__(self, text: str, tokens: List[Union[int, Tensor]]):
|
|
||||||
self.text = text
|
|
||||||
self.tokens = tokens
|
|
||||||
|
|
||||||
def __repr__(self) -> str:
|
|
||||||
cleaned = ", ".join(str(t) if isinstance(t, int) else "EMBD" for t in self.tokens)
|
|
||||||
return f'"{self.text}"({cleaned})'
|
|
||||||
|
|
||||||
def token_length(self) -> int:
|
|
||||||
return len(self.tokens)
|
|
||||||
|
|
||||||
def get_embeddings(self, embedding_module: Embedding) -> Tensor:
|
|
||||||
"""Look up embeddings for plain int tokens.
|
|
||||||
|
|
||||||
Inline TI tensors aren't handled here; the encoder splices them in at a higher level.
|
|
||||||
"""
|
|
||||||
ids = torch.LongTensor([t for t in self.tokens if isinstance(t, int)]).to(
|
|
||||||
embedding_module.weight.device
|
|
||||||
)
|
|
||||||
return embedding_module(ids.unsqueeze(0))
|
|
||||||
@@ -1,18 +0,0 @@
|
|||||||
from typing import Dict, Type
|
|
||||||
|
|
||||||
from ..actions.base import Action
|
|
||||||
|
|
||||||
action_registry: Dict[str, Type[Action]] = {}
|
|
||||||
|
|
||||||
|
|
||||||
def register_action(action: Type[Action]) -> None:
|
|
||||||
name = str(action.action_name)
|
|
||||||
if name in action_registry:
|
|
||||||
raise ValueError(f"Action {name} already registered")
|
|
||||||
action_registry[name] = action
|
|
||||||
|
|
||||||
|
|
||||||
def get_action_by_name(name: str) -> Type[Action]:
|
|
||||||
if name not in action_registry:
|
|
||||||
raise ValueError(f"Action {name} not found in registry")
|
|
||||||
return action_registry[name]
|
|
||||||
@@ -1,81 +0,0 @@
|
|||||||
from typing import List
|
|
||||||
|
|
||||||
from lark import Token, Transformer
|
|
||||||
|
|
||||||
from comfy.sd1_clip import SDTokenizer
|
|
||||||
|
|
||||||
from ..actions.action_utils import parse_numeric_arg
|
|
||||||
from ..actions.base import Action, ActionArity
|
|
||||||
from ..actions.weighted import WeightedGroup
|
|
||||||
from .prompt_segment import PromptSegment
|
|
||||||
from .registration import get_action_by_name
|
|
||||||
from .utils import build_prompt_segment
|
|
||||||
|
|
||||||
|
|
||||||
class PromptTransformer(Transformer):
|
|
||||||
"""Maps the Lark parse tree into a flat list of PromptSegments and Actions."""
|
|
||||||
|
|
||||||
def __init__(self, tokenizer: SDTokenizer):
|
|
||||||
super().__init__()
|
|
||||||
self.tokenizer = tokenizer
|
|
||||||
self.vars: dict = {}
|
|
||||||
|
|
||||||
def assign(self, items):
|
|
||||||
name = str(items[0])
|
|
||||||
if name in self.vars:
|
|
||||||
raise ValueError(f"Variable ${name} is already defined")
|
|
||||||
self.vars[name] = items[1]
|
|
||||||
return None
|
|
||||||
|
|
||||||
def ref(self, items):
|
|
||||||
name = str(items[0])
|
|
||||||
if name not in self.vars:
|
|
||||||
raise ValueError(f"Variable ${name} referenced before assignment")
|
|
||||||
# Weight 1.0 makes the group transparent: _flatten and embedding_tensor
|
|
||||||
# already recurse through WeightedGroup, so no new container type needed.
|
|
||||||
return WeightedGroup(self.vars[name], weight=1.0)
|
|
||||||
|
|
||||||
def item(self, items: List[Token]):
|
|
||||||
for item in items:
|
|
||||||
if isinstance(item, (Action, PromptSegment, WeightedGroup)):
|
|
||||||
return item
|
|
||||||
|
|
||||||
if item.type == "WORD":
|
|
||||||
return build_prompt_segment(str(item), self.tokenizer)
|
|
||||||
if item.type == "QUOTED_STRING":
|
|
||||||
# Strip surrounding quotes, unescape \" and \'.
|
|
||||||
unquoted = item[1:-1]
|
|
||||||
unescaped = unquoted.replace('\\"', '"').replace("\\'", "'")
|
|
||||||
return build_prompt_segment(unescaped, self.tokenizer)
|
|
||||||
raise ValueError(f"Unknown item type: {item.type}")
|
|
||||||
|
|
||||||
def arg(self, items):
|
|
||||||
return items
|
|
||||||
|
|
||||||
def weighted(self, items):
|
|
||||||
arg_items, weight_token = items
|
|
||||||
return WeightedGroup(arg_items, float(weight_token))
|
|
||||||
|
|
||||||
def embedding(self, items):
|
|
||||||
return build_prompt_segment(
|
|
||||||
f"{self.tokenizer.embedding_identifier}{items[0]}",
|
|
||||||
self.tokenizer,
|
|
||||||
)
|
|
||||||
|
|
||||||
def generic_function(self, items):
|
|
||||||
# `emph(text|w)` is sugar for `(text:w)`; handled here so it doesn't need
|
|
||||||
# to fit the Action ABC (it changes weights, not embeddings).
|
|
||||||
if str(items[0]) == "emph":
|
|
||||||
if len(items) != 3:
|
|
||||||
raise ValueError("emph expects exactly two arguments: emph(text|weight)")
|
|
||||||
weight = parse_numeric_arg(items[2], action_name="emph", role="weight", cast=float)
|
|
||||||
return WeightedGroup(items[1], weight)
|
|
||||||
|
|
||||||
action = get_action_by_name(items[0])
|
|
||||||
if action.arity == ActionArity.SINGLE:
|
|
||||||
if len(items) != 2:
|
|
||||||
raise ValueError(f"Action {action.action_name} expects exactly one argument")
|
|
||||||
return action(items[1])
|
|
||||||
if action.arity == ActionArity.MULTI:
|
|
||||||
return action(items[1:])
|
|
||||||
raise ValueError(f"Unknown action arity: {action.arity}")
|
|
||||||
@@ -1,41 +0,0 @@
|
|||||||
from lark import Token
|
|
||||||
|
|
||||||
from comfy.sd1_clip import SDTokenizer
|
|
||||||
|
|
||||||
from .prompt_segment import PromptSegment
|
|
||||||
|
|
||||||
|
|
||||||
def flatten_tree(tree):
|
|
||||||
if isinstance(tree, Token):
|
|
||||||
return [str(tree)]
|
|
||||||
return [str(tree.data)] + sum([flatten_tree(child) for child in tree.children], [])
|
|
||||||
|
|
||||||
|
|
||||||
def build_prompt_segment(text: str, tokenizer: SDTokenizer) -> PromptSegment:
|
|
||||||
"""Tokenize a chunk of plain text into a PromptSegment, expanding `embedding:NAME` refs to tensors."""
|
|
||||||
tokens = []
|
|
||||||
for word in text.split(" "):
|
|
||||||
if word.startswith(tokenizer.embedding_identifier) and tokenizer.embedding_directory is not None:
|
|
||||||
embedding_name = word[len(tokenizer.embedding_identifier):].strip("\n")
|
|
||||||
embedding, leftover = tokenizer._try_get_embedding(embedding_name)
|
|
||||||
if embedding is None:
|
|
||||||
print(f"warning, embedding:{embedding_name} does not exist, ignoring")
|
|
||||||
elif embedding.shape[1] != tokenizer.embedding_size:
|
|
||||||
print(
|
|
||||||
f"warning, embedding:{embedding_name} has size {embedding.shape[1]}, "
|
|
||||||
f"expected {tokenizer.embedding_size}, ignoring"
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
if len(embedding.shape) == 1:
|
|
||||||
tokens.append(embedding)
|
|
||||||
else:
|
|
||||||
tokens.extend(embedding)
|
|
||||||
|
|
||||||
if leftover != "":
|
|
||||||
word = leftover
|
|
||||||
else:
|
|
||||||
continue
|
|
||||||
# Strip the SOT/EOT bracketing tokens added by the underlying CLIP tokenizer.
|
|
||||||
tokens.extend(tokenizer.tokenizer(word)["input_ids"][1:-1])
|
|
||||||
|
|
||||||
return PromptSegment(text, tokens)
|
|
||||||
+161
-105
@@ -1,122 +1,178 @@
|
|||||||
"""DSL-aware tokenizers.
|
import re
|
||||||
|
from typing import Union
|
||||||
|
|
||||||
Override `tokenize_with_weights` to parse our DSL and emit ComfyUI's native
|
from comfy.sd1_clip import SD1Tokenizer
|
||||||
`List[List[(token, weight)]]` format, where `token` is an int id, an inline
|
from custom_nodes.ClipStuff.lib.actions import (
|
||||||
TI tensor, a lazily-evaluated `Action`, or `ACTION_CONTINUATION`.
|
NudgeAction,
|
||||||
|
ArithAction,
|
||||||
|
ALL_START_CHARS,
|
||||||
|
ALL_END_CHARS,
|
||||||
|
ALL_ACTIONS,
|
||||||
|
)
|
||||||
|
from custom_nodes.ClipStuff.lib.actions.base import (
|
||||||
|
Action,
|
||||||
|
PromptSegment,
|
||||||
|
build_prompt_segment,
|
||||||
|
)
|
||||||
|
from custom_nodes.ClipStuff.lib.actions.lib import (
|
||||||
|
is_any_action_segment,
|
||||||
|
is_action_segment,
|
||||||
|
)
|
||||||
|
from custom_nodes.ClipStuff.lib.actions.utils import batch_size_info
|
||||||
|
|
||||||
Row alignment matters: comfy's stock `encode_token_weights` indexes weights by
|
arith_action = r'(<[a-zA-Z0-9\-_]+:[a-zA-Z0-9\-_]+>)'
|
||||||
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
|
# TODO: Get embedding identifier from tokenizer
|
||||||
|
tokenizer_regex = re.compile(
|
||||||
from lark import Tree
|
fr"""
|
||||||
|
\d+\.\d+ # Capture decimals
|
||||||
from comfy.sd1_clip import SD1Tokenizer, SDTokenizer
|
|
|
||||||
|
(?:(?!embedding:)[\w\s]|embedding:[a-zA-Z0-9_]+)+ # Capture sequences of characters, including "embedding:"
|
||||||
from .actions.base import ACTION_CONTINUATION, Action
|
|
|
||||||
from .actions.weighted import WeightedGroup
|
\d+ # Capture whole numbers
|
||||||
from .parser import PromptParser
|
|
|
||||||
from .parser.prompt_segment import PromptSegment
|
[:+-{re.escape("".join(ALL_START_CHARS))}{re.escape("".join(ALL_END_CHARS))}] # Capture special characters including start and end characters
|
||||||
from .parser.transformer import PromptTransformer
|
""",
|
||||||
|
re.VERBOSE
|
||||||
# Side-effect import: registers all built-in actions with the parser.
|
)
|
||||||
from . import actions # noqa: F401
|
def tokenize(text: str) -> list[str]:
|
||||||
|
# Captures:
|
||||||
TokenEntry = Tuple[Union[int, "Action", object], float]
|
# 1. Words
|
||||||
|
# 2. Numbers(1.0, 1)
|
||||||
|
# 3. Special characters(ALL_START_CHARS, ALL_END_CHARS, :, +, -)
|
||||||
|
tokens = re.findall(tokenizer_regex, text)
|
||||||
|
print(tokens)
|
||||||
|
return [token.strip() for token in tokens]
|
||||||
|
|
||||||
|
|
||||||
def _flatten(item, weight: float) -> Iterable[TokenEntry]:
|
|
||||||
"""Walk the parsed item tree, yielding one (token, weight) entry per output slot."""
|
def parse_segment(tokens: list[str], tokenizer: SD1Tokenizer) -> PromptSegment | Action:
|
||||||
if isinstance(item, WeightedGroup):
|
print("Parse segment: Checking token: " + tokens[0])
|
||||||
for sub in item.items:
|
for action in ALL_ACTIONS:
|
||||||
yield from _flatten(sub, weight * item.weight)
|
if tokens[0] == action.START_CHAR:
|
||||||
elif isinstance(item, Action):
|
return action.parse_segment(tokens, ALL_START_CHARS, ALL_END_CHARS, parse_segment, tokenizer)
|
||||||
yield (item, weight)
|
# If we get here, it's a text segment
|
||||||
for _ in range(item.token_length() - 1):
|
return build_prompt_segment(tokens.pop(0), tokenizer)
|
||||||
yield (ACTION_CONTINUATION, weight)
|
|
||||||
elif isinstance(item, PromptSegment):
|
def parse(tokens: list[str], tokenizer: SD1Tokenizer) -> list[PromptSegment | Action]:
|
||||||
for tok in item.tokens:
|
parsed = []
|
||||||
yield (tok, weight)
|
while tokens:
|
||||||
else:
|
if tokens[0] == '':
|
||||||
raise TypeError(f"Unexpected parse item {item!r} ({type(item).__name__})")
|
tokens.pop(0)
|
||||||
|
continue
|
||||||
|
print("Parse: Checking token: " + tokens[0])
|
||||||
|
if tokens[0] in ALL_START_CHARS:
|
||||||
|
parsed.append(parse_segment(tokens, tokenizer))
|
||||||
|
else:
|
||||||
|
parsed.append(build_prompt_segment(tokens.pop(0), tokenizer))
|
||||||
|
return parsed
|
||||||
|
|
||||||
|
|
||||||
class PromptLangSDTokenizer(SDTokenizer):
|
def parse_special_tokens(string) -> list[str]:
|
||||||
def tokenize_with_weights( # type: ignore[override]
|
out = []
|
||||||
self, text: str, return_word_ids: bool = False, **kwargs
|
current = ""
|
||||||
) -> List[List[TokenEntry]]:
|
|
||||||
# SDXL passes a pre-parsed tree to avoid re-running Lark per sub-tokenizer.
|
|
||||||
tree = kwargs.pop("_parsed_tree", None) or PromptParser.parse(text)
|
|
||||||
return self._batch_from_tree(tree)
|
|
||||||
|
|
||||||
def _batch_from_tree(self, tree) -> List[List[TokenEntry]]:
|
for char in string:
|
||||||
pad_token = self.end_token if self.pad_with_end else 0
|
if char in ALL_START_CHARS:
|
||||||
|
out += [current]
|
||||||
parsed = PromptTransformer(self).transform(tree)
|
current = char
|
||||||
items = parsed.children if isinstance(parsed, Tree) else [parsed]
|
elif char in ALL_END_CHARS:
|
||||||
# assign stmts return None (they only populate the transformer's var table).
|
out += [current + char]
|
||||||
items = [i for i in items if i is not None]
|
current = ""
|
||||||
|
else:
|
||||||
batches: List[List[TokenEntry]] = []
|
current += char
|
||||||
current: List[TokenEntry] = [(self.start_token, 1.0)]
|
out += [current]
|
||||||
|
return out
|
||||||
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 parse_segment_actions(string, tokenizer: SD1Tokenizer) -> list[PromptSegment | NudgeAction | ArithAction]:
|
||||||
def __init__(self, embedding_directory=None, tokenizer_data=None, clip_name="l", tokenizer=PromptLangSDTokenizer):
|
tokens = tokenize(string)
|
||||||
super().__init__(
|
parsed = parse(tokens, tokenizer)
|
||||||
embedding_directory=embedding_directory,
|
return parsed
|
||||||
tokenizer_data=tokenizer_data or {},
|
|
||||||
clip_name=clip_name,
|
class TokenDict:
|
||||||
tokenizer=tokenizer,
|
def __init__(self,
|
||||||
)
|
token_id: int,
|
||||||
|
weight: float = None,
|
||||||
|
nudge_id=None, nudge_weight=None, nudge_start: int = None, nudge_end: int = None,
|
||||||
|
arith_ops: dict[str, list[str]] = None):
|
||||||
|
if weight is None:
|
||||||
|
self.weight = 1.0
|
||||||
|
else:
|
||||||
|
self.weight = weight
|
||||||
|
|
||||||
|
self.token_id = token_id
|
||||||
|
self.nudge_id = nudge_id
|
||||||
|
self.nudge_weight = nudge_weight
|
||||||
|
self.nudge_index_start = nudge_start
|
||||||
|
self.nudge_index_stop = nudge_end
|
||||||
|
|
||||||
|
self.arith_ops = arith_ops
|
||||||
|
|
||||||
|
|
||||||
class PromptLangSDXLClipGTokenizer(PromptLangSDTokenizer):
|
class MyTokenizer(SD1Tokenizer):
|
||||||
def __init__(self, tokenizer_path=None, embedding_directory=None, tokenizer_data=None):
|
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__(
|
super().__init__(tokenizer_path, max_length, pad_with_end, embedding_directory, embedding_size, embedding_key)
|
||||||
tokenizer_path=tokenizer_path,
|
|
||||||
pad_with_end=False,
|
|
||||||
embedding_directory=embedding_directory,
|
|
||||||
embedding_size=1280,
|
|
||||||
embedding_key="clip_g",
|
|
||||||
tokenizer_data=tokenizer_data or {},
|
|
||||||
)
|
|
||||||
|
|
||||||
|
"""
|
||||||
|
Doesn't actually tokenize...
|
||||||
|
Returns batches of segments and actions
|
||||||
|
:return: List of list(batches) of segments and actions
|
||||||
|
"""
|
||||||
|
def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs) -> list[list[PromptSegment | Action]]:
|
||||||
|
if self.pad_with_end:
|
||||||
|
pad_token = self.end_token
|
||||||
|
else:
|
||||||
|
pad_token = 0
|
||||||
|
|
||||||
class PromptLangSDXLTokenizer:
|
parsed_actions = parse_segment_actions(text, self)
|
||||||
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]]]:
|
# nudge_start = kwargs.get("nudge_start")
|
||||||
tree = PromptParser.parse(text)
|
# nudge_end = kwargs.get("nudge_end")
|
||||||
return {
|
#
|
||||||
"g": self.clip_g.tokenize_with_weights(text, return_word_ids, _parsed_tree=tree, **kwargs),
|
# if nudge_start is not None and nudge_end is not None:
|
||||||
"l": self.clip_l.tokenize_with_weights(text, return_word_ids, _parsed_tree=tree, **kwargs),
|
# nudge_start = int(nudge_start)
|
||||||
}
|
# nudge_end = int(nudge_end)
|
||||||
|
#
|
||||||
|
# # tokenize words
|
||||||
|
for segment in parsed_actions:
|
||||||
|
if isinstance(segment, Action):
|
||||||
|
print(segment.depth_repr())
|
||||||
|
else:
|
||||||
|
print(segment.depth_repr())
|
||||||
|
|
||||||
def untokenize(self, token_weight_pair):
|
# reshape token array to CLIP input size
|
||||||
return self.clip_g.untokenize(token_weight_pair)
|
batched_segments = []
|
||||||
|
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token])]
|
||||||
|
# batched_segments.append(batch)
|
||||||
|
batch_size = 1
|
||||||
|
for segment in parsed_actions:
|
||||||
|
num_tokens = segment.token_length()
|
||||||
|
# determine if we're going to try and keep the tokens in a single batch
|
||||||
|
is_large = num_tokens >= self.max_word_length
|
||||||
|
|
||||||
def state_dict(self):
|
# If the segment is too large to fit in a single batch, pad the current batch and start a new one
|
||||||
return {}
|
if num_tokens + batch_size > self.max_length - 1:
|
||||||
|
remaining_length = self.max_length - batch_size - 1 # -1 for end token
|
||||||
|
# Pad batch
|
||||||
|
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length - 1))
|
||||||
|
batched_segments.append(batch)
|
||||||
|
|
||||||
|
# start new batch
|
||||||
|
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token]), segment]
|
||||||
|
batch_size = num_tokens + 1 # +1 for start token
|
||||||
|
continue
|
||||||
|
|
||||||
|
# If the segment is small enough to fit in the current batch, add it
|
||||||
|
batch.append(segment)
|
||||||
|
batch_size += num_tokens
|
||||||
|
|
||||||
|
# Pad the last batch
|
||||||
|
remaining_length = self.max_length - batch_size - 1 # -1 for end token
|
||||||
|
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length))
|
||||||
|
batched_segments.append(batch)
|
||||||
|
|
||||||
|
for batch in batched_segments:
|
||||||
|
batch_size_info(batch)
|
||||||
|
|
||||||
|
return batched_segments
|
||||||
|
|||||||
@@ -1,24 +1,23 @@
|
|||||||
import os
|
import random
|
||||||
from dataclasses import dataclass
|
|
||||||
from typing import Any, List, Tuple
|
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
import comfy.sd
|
|
||||||
import folder_paths
|
import folder_paths
|
||||||
from comfy.supported_models_base import ClipTarget
|
import comfy.sd
|
||||||
|
import comfy.ops
|
||||||
|
from custom_nodes.ClipStuff.lib.clip_model import SD1FunClipModel
|
||||||
|
|
||||||
from .lib.clip_model import PromptLangSD1ClipModel, PromptLangSDXLClipModel
|
from custom_nodes.ClipStuff.lib.tokenizer import MyTokenizer
|
||||||
from .lib.inspect import inspect_prompt
|
|
||||||
from .lib.tokenizer import PromptLangSD1Tokenizer, PromptLangSDXLTokenizer
|
|
||||||
|
class EmptyClass:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
class SpecialClipLoader:
|
class SpecialClipLoader:
|
||||||
"""Wraps a loaded CLIP with our DSL-aware tokenizer + text encoder."""
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls): # type: ignore[no-untyped-def]
|
def INPUT_TYPES(s):
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"source_clip": ("CLIP",),
|
"source_clip": ("CLIP",),
|
||||||
@@ -27,200 +26,175 @@ class SpecialClipLoader:
|
|||||||
|
|
||||||
RETURN_TYPES = ("CLIP",)
|
RETURN_TYPES = ("CLIP",)
|
||||||
FUNCTION = "load_clip"
|
FUNCTION = "load_clip"
|
||||||
|
OUTPUT_IS_LIST = (False,)
|
||||||
CATEGORY = "conditioning"
|
CATEGORY = "conditioning"
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def load_clip(source_clip: comfy.sd.CLIP) -> Tuple[comfy.sd.CLIP]:
|
def load_clip(source_clip):
|
||||||
is_sdxl = hasattr(source_clip.cond_stage_model, "clip_g") and hasattr(source_clip.cond_stage_model, "clip_l")
|
clip_target = EmptyClass()
|
||||||
embedding_directory = source_clip.tokenizer.clip_l.embedding_directory
|
clip_target.params = {}
|
||||||
|
clip_target.clip = SD1FunClipModel
|
||||||
|
clip_target.tokenizer = MyTokenizer
|
||||||
|
|
||||||
if is_sdxl:
|
# TODO: Extract embedding directory from source_clip
|
||||||
target = ClipTarget(PromptLangSDXLTokenizer, PromptLangSDXLClipModel)
|
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.embedding_directory)
|
||||||
else:
|
comfy.sd.load_clip_weights(
|
||||||
target = ClipTarget(PromptLangSD1Tokenizer, PromptLangSD1ClipModel)
|
clip.cond_stage_model, source_clip.cond_stage_model.state_dict()
|
||||||
|
)
|
||||||
new_clip = comfy.sd.CLIP(target=target, embedding_directory=embedding_directory)
|
return (clip,)
|
||||||
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:
|
class KepAdvTextEncode:
|
||||||
"""Shows what a DSL prompt resolves to at the embedding layer: per-slot weight, L2 norm, nearest vocab."""
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls): # type: ignore[no-untyped-def]
|
def INPUT_TYPES(s):
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"clip": ("CLIP",),
|
|
||||||
"text": ("STRING", {"multiline": True}),
|
"text": ("STRING", {"multiline": True}),
|
||||||
"top_k": ("INT", {"default": 3, "min": 1, "max": 10}),
|
"clip": ("CLIP",),
|
||||||
|
"nudge_start": ("INT", {}),
|
||||||
|
"nudge_end": ("INT", {}),
|
||||||
|
"split_newlines": ("BOOL", {"default": True}),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("STRING",)
|
RETURN_TYPES = ("CONDITIONING",)
|
||||||
FUNCTION = "inspect"
|
FUNCTION = "encode"
|
||||||
|
OUTPUT_IS_LIST = (True,)
|
||||||
CATEGORY = "conditioning"
|
CATEGORY = "conditioning"
|
||||||
OUTPUT_NODE = True
|
|
||||||
|
|
||||||
def inspect(self, clip, text: str, top_k: int):
|
@staticmethod
|
||||||
report = inspect_prompt(clip, text, top_k=top_k)
|
def encode(clip, text, nudge_start, nudge_end, split_newlines):
|
||||||
return {"ui": {"text": [report]}, "result": (report,)}
|
ret = []
|
||||||
|
if split_newlines:
|
||||||
|
prompts = text.split("\n")
|
||||||
|
else:
|
||||||
|
prompts = [text]
|
||||||
|
|
||||||
|
for prompt in prompts:
|
||||||
|
if prompt.strip() == "":
|
||||||
|
continue
|
||||||
|
tokens = clip.tokenizer.tokenize_with_weights(
|
||||||
|
text,
|
||||||
|
return_word_ids=False,
|
||||||
|
nudge_start=nudge_start,
|
||||||
|
nudge_end=nudge_end,
|
||||||
|
)
|
||||||
|
cond, pooled = clip.encode_from_tokens(
|
||||||
|
tokens, return_pooled=True, position_ids=[0] * 77
|
||||||
|
)
|
||||||
|
cond = [[cond, {"pooled_output": pooled}]]
|
||||||
|
ret.append(cond)
|
||||||
|
return (ret,)
|
||||||
|
|
||||||
|
|
||||||
def tensor2img(tensor_img) -> Image.Image:
|
def tensor2img(tensor_img):
|
||||||
arr = (255.0 * tensor_img.cpu().numpy()).clip(0, 255).astype(np.uint8)
|
i = 255.0 * tensor_img.cpu().numpy()
|
||||||
return Image.fromarray(arr)
|
i_np_arr = np.clip(i, 0, 255, out=i).astype(np.uint8, copy=False)
|
||||||
|
return Image.fromarray(i_np_arr)
|
||||||
|
|
||||||
|
|
||||||
class BuildGif:
|
class BuildGif:
|
||||||
"""Builds an animated webp from a list of image batches.
|
def __init__(self):
|
||||||
|
pass
|
||||||
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()
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls): # type: ignore[no-untyped-def]
|
def INPUT_TYPES(cls):
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"images": ("IMAGE",),
|
"images": ("IMAGE",),
|
||||||
"split_every": ("INT", {"default": -1}),
|
"split_every": ("INT", {"default": -1}),
|
||||||
"frame_duration": ("INT", {"default": 125}),
|
"output_mode": (
|
||||||
"output_mode": (["One Per Split", "Big Grid"], {"default": "Big Grid"}),
|
["One Per Split", "Big Grid"],
|
||||||
|
{"default": "Big Grid"},
|
||||||
|
),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ()
|
RELOAD_INST = True
|
||||||
|
RETURN_TYPES = ("IMAGE",)
|
||||||
|
RETURN_NAMES = ("Gifs",)
|
||||||
INPUT_IS_LIST = True
|
INPUT_IS_LIST = True
|
||||||
FUNCTION = "build_gif"
|
FUNCTION = "build_gif"
|
||||||
OUTPUT_NODE = True
|
OUTPUT_IS_LIST = (True,)
|
||||||
|
# OUTPUT_NODE = False
|
||||||
|
|
||||||
CATEGORY = "List Stuff"
|
CATEGORY = "List Stuff"
|
||||||
|
|
||||||
def build_gif(
|
@staticmethod
|
||||||
self,
|
def build_gif(images: list, split_every: list[int], output_mode: str):
|
||||||
images: List[Any],
|
print("Build GIF called!")
|
||||||
split_every: List[int],
|
print(f"{type(images)}")
|
||||||
frame_duration: List[int],
|
|
||||||
output_mode: List[str],
|
|
||||||
):
|
|
||||||
if len(split_every) > 1:
|
if len(split_every) > 1:
|
||||||
raise ValueError("List input for split_every is not supported.")
|
raise Exception("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
|
|
||||||
)
|
|
||||||
|
|
||||||
|
split_every = split_every[0]
|
||||||
batch_size = images[0].size()[0]
|
batch_size = images[0].size()[0]
|
||||||
# split_every=-1 means "don't split": one chunk containing everything.
|
if split_every == -1:
|
||||||
if split_requested == -1:
|
|
||||||
split_chunks = 1
|
split_chunks = 1
|
||||||
chunk_len = len(images)
|
split_every = len(images)
|
||||||
else:
|
else:
|
||||||
chunk_len = split_requested
|
split_chunks = int(len(images) / split_every)
|
||||||
split_chunks = len(images) // chunk_len
|
|
||||||
|
out = []
|
||||||
|
|
||||||
|
num_wide = batch_size
|
||||||
|
num_tall = split_chunks
|
||||||
|
|
||||||
chunked_batches = [
|
chunked_batches = [
|
||||||
images[chunk_len * i : chunk_len * (i + 1)]
|
images[split_every * chunk_idx : split_every * (chunk_idx + 1)]
|
||||||
for i in range(split_chunks)
|
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 = []
|
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}_"
|
if output_mode == "Big Grid":
|
||||||
save_path = os.path.join(ctx.full_output_folder, file)
|
# For every image in gif
|
||||||
frames[0].save(
|
for idx_in_chunk in range(split_every):
|
||||||
f"{save_path}.webp",
|
img_shape = images[0][0].shape
|
||||||
lossless=True,
|
img_frame = Image.new(
|
||||||
save_all=True,
|
"RGB", size=(num_wide * img_shape[0], num_tall * img_shape[1])
|
||||||
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,
|
|
||||||
)
|
)
|
||||||
results.append({"filename": f"{file}.webp", "subfolder": ctx.subfolder, "type": "output"})
|
# For every chunk of images
|
||||||
return results
|
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)
|
||||||
|
|
||||||
|
save_path = (
|
||||||
@dataclass
|
f"{folder_paths.get_output_directory()}/{random.randint(1, 100)}"
|
||||||
class _SaveContext:
|
)
|
||||||
images: Any
|
frames[0].save(
|
||||||
chunked_batches: Any
|
f"{save_path}.webp",
|
||||||
chunk_len: int
|
# quality=100,
|
||||||
batch_size: int
|
# method=6,
|
||||||
split_chunks: int
|
lossless=True,
|
||||||
full_output_folder: str
|
save_all=True,
|
||||||
filename: str
|
append_images=frames[1:],
|
||||||
counter: int
|
optimize=False,
|
||||||
subfolder: str
|
duration=125,
|
||||||
duration: int
|
loop=0,
|
||||||
|
)
|
||||||
|
elif output_mode == "One Per Split":
|
||||||
|
for split_idx in range(int(split_chunks)):
|
||||||
|
split_start = split_every * split_idx
|
||||||
|
split_end = split_every * (split_idx + 1)
|
||||||
|
for batch_idx in range(batch_size):
|
||||||
|
save_path = f"{folder_paths.get_output_directory()}/-{batch_idx}-{random.randint(1, 100)}"
|
||||||
|
print(save_path)
|
||||||
|
tensor2img(images[split_start][batch_idx]).save(
|
||||||
|
f"{save_path}.webp",
|
||||||
|
save_all=True,
|
||||||
|
append_images=[
|
||||||
|
tensor2img(nested_batch[batch_idx])
|
||||||
|
for nested_batch in images[split_start + 1 : split_end]
|
||||||
|
],
|
||||||
|
optimize=False,
|
||||||
|
duration=125,
|
||||||
|
loop=0,
|
||||||
|
)
|
||||||
|
return (out,)
|
||||||
|
|||||||
@@ -1,26 +0,0 @@
|
|||||||
[project]
|
|
||||||
name = "keppromptlang"
|
|
||||||
version = "0.2.0"
|
|
||||||
description = "A small DSL for ComfyUI that lets you do math on CLIP token embeddings before they're fed into the text transformer."
|
|
||||||
readme = "README.md"
|
|
||||||
license = { text = "MIT" }
|
|
||||||
requires-python = ">=3.10"
|
|
||||||
dependencies = [
|
|
||||||
"lark",
|
|
||||||
]
|
|
||||||
|
|
||||||
[project.optional-dependencies]
|
|
||||||
dev = [
|
|
||||||
"pytest",
|
|
||||||
"torch",
|
|
||||||
]
|
|
||||||
|
|
||||||
[project.urls]
|
|
||||||
Repository = "https://github.com/M1kep/KepPromptLang"
|
|
||||||
|
|
||||||
[tool.comfy]
|
|
||||||
PublisherId = "m1kep"
|
|
||||||
DisplayName = "KepPromptLang"
|
|
||||||
|
|
||||||
[tool.pytest.ini_options]
|
|
||||||
testpaths = ["tests"]
|
|
||||||
@@ -1 +0,0 @@
|
|||||||
lark
|
|
||||||
@@ -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()
|
|
||||||
@@ -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
|
|
||||||
@@ -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
|
|
||||||
@@ -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)
|
|
||||||
@@ -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
|
|
||||||
@@ -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) == []
|
|
||||||
@@ -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()))
|
|
||||||
Reference in New Issue
Block a user