Author SHA1 Message Date
Claude f4c862ec56 Add $var bindings and # comments
$NAME = arg; binds; $NAME substitutes. Single-pass: define before
use, no reassignment. Substitution is structural — refs share the
parsed action object, but actions are evaluated per occurrence (so
$r = rand(3); $r $r re-rolls each use).

Implementation: assign stores into PromptTransformer.vars and
returns None (filtered by tokenizer); ref returns the stored arg
wrapped in a transparent WeightedGroup(items, 1.0), so existing
_flatten and embedding_tensor code paths handle it without a new
container type.

NAME and WORD terminals overlap on bare identifiers; the earley
parser's dynamic lexer disambiguates by grammar context (the $
prefix forces NAME). Noted in grammar.py since this would break
under a basic/contextual lexer.

# comments run to end-of-line and are lexer-ignored.

Nine new tests cover top-level/arg-level/weighted substitution,
chaining, define-before-use error, reassignment error, comments,
and assign-only prompts.
2026-04-13 02:58:56 +00:00
Claude 3dddbe1671 Add proj, reject, renorm, noise, nearest actions and lerp alias
proj(a|b)    - project a onto the (mean, unit) direction of b
reject(a|b)  - a minus that projection (orthogonal component)
renorm(a|ref)- rescale a so each token's L2 norm matches ref's mean L2
noise(a|std) - a + N(0, std)
nearest(e|k) - snap mean(e) to its k nearest vocab tokens (cosine sim)
lerp(a|b|t)  - alias for avg(a|b|t)

All follow the existing MultiArgAction pattern; nearest reuses the
embedding-weight cosine machinery the Inspect node uses. Nine new
tests cover the math invariants (proj+reject reconstructs input,
reject is orthogonal to b, renorm matches ref norm, noise(_, 0) is
identity, nearest of a single token returns that token).

Relax MultiArgAction/SingleArgAction __init__ type hints to List
since SegOrAction now includes WeightedGroup and the Union isn't
importable in base.py without a cycle.
2026-04-12 06:02:00 +00:00
Claude c6910ab775 Add (word:1.2) weight syntax and PromptLang Inspect node
Weight syntax:
- Grammar: `(arg:SIGNED_NUMBER)` parses to a WeightedGroup
- emph(text|w) is a function-style alias for the same
- Nested weights multiply: ((cat:1.2):0.5) -> 0.6, matching A1111
- Weights are emitted in the (token, weight) tuples that comfy's
  stock encode_token_weights already applies post-transformer

To make per-position weight indexing work, rows are now always
exactly max_length entries: a multi-slot Action emits as one
(action, w) entry followed by (ACTION_CONTINUATION, w) placeholders
that process_tokens drops before handing to super().process_tokens.
This replaces the previous variable-length-row + padding-math.

PromptLang Inspect:
- New node: CLIP + text -> STRING report of per-slot weight, L2
  norm, and top-k nearest vocab tokens for the resolved embedding
- lib/inspect.py uses the wrapper's stored attr name (cond.clip /
  tok.clip) rather than hardcoding 'clip_l', and uses the
  tokenizer's inv_vocab for decode

Fix: parse_numeric_arg now rejects WeightedGroup (was only Action).
SegOrAction widened to include WeightedGroup.
2026-04-12 05:52:06 +00:00
Claude 36c0a43395 Lower DSL into ComfyUI's native token format
The tokenizer now emits comfy's standard List[List[(token, weight)]]
format with one extension: a token entry may also be a lazy Action
instance. PromptLangSDClipModel shrinks to a single process_tokens
override that resolves Actions to tensors and otherwise defers to
the stock SDClipModel.process_tokens for embedding lookup, mask
construction, TI splice, and num_tokens accounting.

This deletes _build_embeddings, _build_attention_mask, and the
custom encode_token_weights override (which were re-implementing
what comfy already does), and means stock encode_token_weights /
forward handle everything else.

posScale/postPos still work via the (modified - default) pre-bake
in _apply_pos_modifiers; that's unchanged.

Also:
- SDXL parses the Lark tree once and transforms per-CLIP, instead
  of re-parsing the same text for both clip_l and clip_g
- Drop Action.get_all_segments (was only needed by the old
  set_up_textual_embeddings path)
- Inline get_embedding into embedding_tensor (no external callers)
- SpecialClipLoader: load full cond_stage_model state_dict instead
  of per-transformer copies
- Add pyproject.toml with [tool.comfy] metadata; drop packaging
  dep (was only for HF version sniffing in fun_clip_stuff.py)
- Add tokenizer tests covering native-format emission and slot
  accounting (logical post-splice length == max_length)
2026-04-12 05:24:23 +00:00
Claude 6fcd3b8a86 Port to current ComfyUI CLIP API and refactor
Rebase the DSL text encoder on comfy.sd1_clip.SDClipModel and
comfy.clip_model.CLIPTextModel_, replacing the old HuggingFace
transformers CLIPTextModel/CLIPTextTransformer/CLIPTextEmbeddings
subclasses (which are no longer how ComfyUI implements CLIP).

Structure:
- Merge lib/action/ into lib/actions/; drop lib/fun_clip_stuff.py
  and the bundled clip_config*.json (comfy ships its own)
- Convert all custom_nodes.KepPromptLang.* absolute imports to
  relative imports so installs via Manager work regardless of
  install directory name

Encoder:
- PromptLangSDClipModel overrides encode_token_weights to walk
  the segment/action tree, assemble [B, seq, hidden] embeds, and
  call self.transformer(None, mask, embeds=..., num_tokens=...)
- posScale / postPos are supported without patching the transformer
  by pre-baking the delta (modified - default) into embeds, so the
  transformer's inline add yields the modified position embedding
- Drop the unused empty-baseline batch that was prepended and then
  sliced off; halves the per-encode forward pass for single prompts

Actions:
- Fix copy-pasted broken __repr__ / depth_repr across sum/diff/avg/
  slerp that referenced fields that didn't exist
- Fix NameError in AverageAction._validate_args (start_arg_token_length)
- Fix class-level mutable state in RandAction
- Unify _parse_scalar / _parse_scalar_weight / _parse_int into a
  single parse_numeric_arg helper in action_utils.py
- Share add_with_broadcast between SumAction and DiffAction
- Convert PostModifiers from TypedDict to dataclass (attribute
  access catches typos that .get() on string keys hides)

Drop the two _exp-pooler / _exp-pooledAvg actions: they recursively
invoked the HF CLIPTextTransformer and would need a rework to fit
the current CLIPTextModel_ interface. They were experimental and
not documented as stable.

BuildGif node:
- Collapse the 10-positional-arg _save_* helpers onto a small
  _SaveContext dataclass
- Stop mutating the input arg semantics (split_every_val was
  reassigned to len(images) when -1)

Tests:
- Add pytest suite covering the parser (grammar, nesting, errors)
  and every action (embedding math, shape, modifier payloads)
- conftest stubs ComfyUI at collection time so tests don't need a
  real ComfyUI install

Drop the broken test_files/ scripts (CI helpers, not a test suite)
and regenerate the README from tools/build_docs.py.
2026-04-12 05:05:58 +00:00
Michael Poutre a0b3695800 refactor(Docs): Add table of all actions, and tiny cleanup 2023-11-19 00:26:57 -08:00
Michael Poutre c299354c9c feat(Docs): Add script for generating docs 2023-11-19 00:26:33 -08:00
Michael Poutre e002d0ad64 refactor(Action): Update property names to allow automatic doc gen 2023-11-19 00:26:19 -08:00
Michael Poutre b5a728a997 feat(Action): Add experimental pooledAvg, pooler actions(_exp- prefix) 2023-11-18 23:41:27 -08:00
Michael Poutre 72f46ad938 feat(Action): Add PostPost action 2023-11-18 23:07:01 -08:00
Michael Poutre f74728267d feat: Add call to process_with_transformers for custom actions 2023-11-18 23:05:42 -08:00
Michael Poutre 0de2ab3f9a fix: Properly determine seq_len for TextTransformer processing 2023-11-18 23:04:56 -08:00
Michael Poutre 2fa1b95045 fix: Handle EOT as well as __PAD__ segments for pooler output EOT calc 2023-11-18 23:04:06 -08:00
Michael Poutre dba41a2c4d fix: Drop invalid embeddings during prompt parsing to avoid it later 2023-11-18 22:48:10 -08:00
Michael Poutre 3cbfc8b78c feat(Actions): Add bypass_pos_embed PostModifier 2023-11-18 22:46:39 -08:00
Michael Poutre fe74c556f5 refactor(PromptTransformer): Simple cleanup and consolidation 2023-11-18 21:36:28 -08:00
Michael Poutre d042fc572f fix: Update importlib.metadata import 2023-11-13 19:54:15 -08:00
Michael Poutre 46ae926911 fix: Update to work with transformers changes to attention masks 2023-11-13 19:47:56 -08:00
Michael Poutre 2823a80078 feat(PosScale: Add node 2023-11-13 19:47:56 -08:00
Michael Poutre fccffdbdf0 feat: Add support for post embedding modifiers 2023-11-13 19:47:56 -08:00
Michael Poutre 7e138268f9 chore: Cleanup imports 2023-11-13 19:47:56 -08:00
Michael Poutre ef3f66ef91 refactor/fix: More fixes to align with clip changes in base 2023-11-13 19:47:56 -08:00
Michael Poutre 2f115352ec refactor/fix(SD1): Update for changes to clip handling
Fix end token calculation for long prompts
2023-11-13 19:47:56 -08:00
Michael Poutre a0ebcdaf4d fix(Clip): Fix bug in EOT token detection 2023-11-13 19:47:56 -08:00
Michael Poutre 55cc09b0be feat(SDXL): Initial stab at SDXL 2023-11-13 19:47:56 -08:00
Michael Poutre 2e731a480d refactor(Node): Add checks for clip model type 2023-11-13 19:47:56 -08:00
Michael Poutre 65bc2c0327 refactor(Node): Use new comfy.supported_models_base.ClipTarget 2023-11-13 19:47:56 -08:00
Michael Poutre d67d8600ab Merge branch 'func/setDims' 2023-09-05 00:14:13 -07:00
Michael Poutre 6444e2890a feat(func/setDims): Add func 2023-09-04 23:35:24 -07:00
Michael Poutre 28141f2cbe refactor(nodes): Update Build Gif to show preview of Gif 2023-09-04 18:32:53 -07:00
Michael Poutre 70704f5e68 feat(func/scaleDims): Add scaleDims 2023-09-04 18:15:30 -07:00
Michael Poutre e59aaa499d fix(action reg): Don't allow multiple actions with same name 2023-09-04 18:00:35 -07:00
Michael Poutre c98df8d289 feat(fun/average): Add function 2023-08-31 23:50:45 -07:00
Michael Poutre a7c4bbe332 feat(nodes): Update saving method for build_gif 2023-08-31 23:05:41 -07:00
Michael Poutre ce795c52bc refactor(func/mult): Update error messages 2023-08-31 22:51:27 -07:00
Michael Poutre ef9693ec73 feat(nodes): Add frame_duration to build_gif 2023-08-31 22:51:16 -07:00
Michael Poutre 5362ac75fa feat(func/slerp): Add function 2023-08-31 22:50:47 -07:00
Michael Poutre 4fcdee837b fix(func/sum): Allow for passing a single argument 2023-08-31 21:47:22 -07:00
Michael Poutre 6b2bf7485e fix(promptsegment): Put tokens on same device as embeddingmodule weights 2023-08-31 14:19:49 -07:00
Michael Poutre d330663624 fix(promtsegment): Put tokens on the GPU 2023-08-31 14:13:39 -07:00
Michael Poutre 5de8d6821c fix(typing): tuple -> Tuple for <=python3.8 2023-08-31 00:33:58 -07:00
Michael Poutre 0f36544b0e feat(func): Add Mult 2023-08-31 00:29:58 -07:00
Michael Poutre 05920e0392 refactor(nodes): Cleanup some Mypy type errors 2023-08-30 23:42:34 -07:00
Michael Poutre 9ccf5d2583 feat: Use a more generic grammar to allow for easier interoperability 2023-08-30 23:40:39 -07:00
Michael Poutre f185b39f06 feat(func/rand): Add support for defining range for rand 2023-08-30 21:06:23 -07:00
Michael Poutre af761a5620 feat(func): Add rand(<token_length>) 2023-08-30 20:29:46 -07:00
Michael Poutre fee9c56abd refactor: Rename repo to match github 2023-08-30 18:44:11 -07:00
Michael Poutre 6879267391 refactor(Tests): Only run 1 step in sampler 2023-08-30 18:33:59 -07:00
Michael Poutre a40ed34eac fix(CI): Source venv for all python runs 2023-08-30 18:26:11 -07:00
Michael Poutre 94b2347d01 feat(CI): Cache venv instead of pip cache 2023-08-30 18:23:34 -07:00
Michael Poutre bcc083f2fa fix: Use typing.List for backwards compatibility 2023-08-30 18:09:10 -07:00
Michael Poutre f91d3aa233 fix(Typing): Replace | syntax with Union 2023-08-29 16:23:05 -07:00
Michael Poutre b497a81b43 fix(CI): PYTHONBUFFERED to correct step 2023-08-29 16:13:05 -07:00
Michael Poutre 4076d2beb1 fix(CI): Try sleeping for 30s maybe? 2023-08-29 16:03:06 -07:00
Michael Poutre e644f219dd fix(CI): Set PYTHONUNBUFFERED=1 for running ComfyUI server 2023-08-29 15:50:30 -07:00
Michael Poutre 9e7e79a1a1 fix(tests): Read error body, then attempt to load JSON 2023-08-29 15:30:10 -07:00
Michael Poutre f18ef4b292 feat(CI): Archive server.log 2023-08-29 15:18:49 -07:00
Michael Poutre 78c14ce95a feat(CI): Add all supported python version 2023-08-29 15:18:39 -07:00
Michael Poutre 60bc4e5a4d fix(tests): Log error when JSON decode fails 2023-08-29 15:18:22 -07:00
Michael Poutre c11a6c5178 fix(CI): Don't fail fast on multi-python test 2023-08-29 15:11:50 -07:00
Michael Poutre c0adb16e30 feat(CI): Test python 3.9-11 2023-08-29 15:08:27 -07:00
Michael Poutre df98b07ffc fix(ClipTransformer): Fix causal_map change between transformer versions 2023-08-28 22:28:10 -07:00
Michael Poutre 79f17aac60 workflows: Install correct websocket library 2023-08-28 22:09:18 -07:00
Michael Poutre 7804426fa2 workflows: Better output and pass error to actions 2023-08-28 22:01:35 -07:00
Michael Poutre 5be4787a12 fix(ClipTransformer): Update with transformers library 2023-08-28 22:01:16 -07:00
Michael Poutre fd2316c8fa workflows: Fix double ext... 2023-08-28 21:38:44 -07:00
Michael Poutre 137a3f0f24 workflows: Error handling on script 2023-08-28 21:35:08 -07:00
Michael Poutre c422e6000c Move up Debug actino 2023-08-28 21:26:12 -07:00
Michael Poutre 5054bfdb82 actions: Fix again 2023-08-28 21:23:33 -07:00
Michael Poutre 7a854c7fea workflow: Open json relative to script file 2023-08-28 21:19:46 -07:00
Michael Poutre 9f12d2f16d workflow: Don't use cache for SD - To slow 2023-08-28 21:15:00 -07:00
Michael Poutre bd369930bd Fix run_workflow.py 2023-08-28 21:14:39 -07:00
Michael Poutre 9bc7f54bcf Workflow: Sleep longer 2023-08-28 21:08:21 -07:00
Michael Poutre 337dad1cc9 Add more files for workflow testing 2023-08-28 21:00:18 -07:00
Michael Poutre ade09bf806 update(clip_model): Sync with upstream 2023-08-28 20:19:52 -07:00
Michael Poutre 88c3804446 New test node 2023-08-28 19:54:14 -07:00
Michael Poutre acbaf7cefe First one didn't show up for some reason.. 2023-08-28 19:07:49 -07:00
Michael Poutre 1f1e74cd30 Merge branch 'gh-workflow' 2023-08-28 19:05:22 -07:00
47 changed files with 2303 additions and 1001 deletions
+85
View File
@@ -0,0 +1,85 @@
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
+87 -68
View File
@@ -1,81 +1,100 @@
# ClipStuff
## Basic Instructions.
Clone repo into custom_nodes folder.
Install the requirements.txt file via pip.
# KepPromptLang
Pass CLIP output from Load Checkpoint into SpecialClipLoader node, then use the outputted clip with standard Clip Text Encode.
A small DSL for ComfyUI that lets you do math on CLIP token embeddings before they're fed into the text transformer.
See example workflow in examples folder.
```
sum(diff(king|man)|woman)
norm(sum(cat | dog | horse | parrot))
A slerp(cat|dog|0.5) is happy
```
### Example Photo
![Example Photo](assets/first_example.png)
## Install
Clone into `ComfyUI/custom_nodes/`:
```bash
cd ComfyUI/custom_nodes
git clone <repo-url> KepPromptLang
pip install -r KepPromptLang/requirements.txt
```
## Usage
1. Add a **Special CLIP Loader** node and feed it the CLIP output from your **Load Checkpoint**.
2. Pass the wrapped CLIP into a standard **CLIP Text Encode** node.
3. Use the DSL syntax in your prompt.
To debug what your DSL is doing, add a **PromptLang Inspect** node — it shows the per-slot weight, L2 norm, and nearest-vocab words for the resolved embeddings.
See `examples/WIP_Example_workflow.json` for a working workflow.
![Example](assets/first_example.png)
## Syntax
| Element | Syntax | Example |
| --- | --- | --- |
| Plain word | alphanumeric (with `,_.-`) | `cat`, `dog_face` |
| Quoted string | single or double quotes | `"hello world"`, `'it\'s sunny'` |
| Weighted | `(text:weight)` or `emph(text\|weight)` | `(cat:1.3)`, `emph(cat\|1.3)` |
| Embedding (textual inversion) | `embedding:NAME` | `embedding:face_vector` |
| Function | `name(arg \| arg \| ...)` | `sum(king \| woman)` |
Arguments inside a function are separated by `|`. Each arg can itself be plain text, an embedding, a quoted string, or another function call.
### Variables and comments
```
$axis = diff(king|queen); # name an expression
sum(actor|$axis) and reject(doctor|$axis)
```
`$NAME = arg;` binds a name; `$NAME` substitutes it. Single-pass: define before use, no reassignment. `#` comments run to end of line. Substitution is structural — multiple refs share the same parsed action object, but actions are evaluated per occurrence (so `$r = rand(3); $r $r` re-rolls each use).
## Quick examples
- Average two prompts: `avg(The cat is | The dog is | 0.5)`
- Normalize a sum: `norm(sum(cat | dog | horse))`
- King − Man + Woman: `sum(diff(king|man)|woman)` (or `sum(king | neg(man) | woman)`)
- Negate an embedding: `neg(embedding:body_vector)`
## Functions
## Syntax Elements
| 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> |
1. **Embedding**:
- Syntax: `embedding:WORD`
- Example: `embedding:face_vector`
- Represents a named vector embedding(Textual Inversion).
`lerp(a|b|t)` is also accepted as an alias for `avg(a|b|t)`.
2. **Word**:
- Syntax: Any alphanumeric word including characters such as `,`, `_`, and `-`.
- Example: `cat, dog_face, id_123`
- Represents simple words or identifiers.
Regenerate the table with `python tools/build_docs.py`.
3. **Quoted String**:
- Syntax: A string enclosed within double or single quotes. You can escape quotes inside the string using a backslash (`\`).
- Example: `"Hello World"`, `'It\'s a sunny day'`
- Represents string literals.
## Development
## Functions
Tests are pytest-based and don't require ComfyUI:
Here are the available functions and their usage:
```bash
pip install -e ".[dev]"
python -m pytest
```
1. **Sum Function**:
- Syntax: `sum(arg1 | arg2 | ... | argN)`
- Adds together multiple embeddings.
- Example: `sum(embedding:face1 | dog)`
## Compatibility
2. **Negation Function**:
- Syntax: `neg(arg)`
- Negates the output.
- Example: `neg(A embedding:happycats outside)`
3. **Normalization Function**:
- Syntax: `norm(arg)`
- Normalizes the given vector embedding.
- Example: `norm(sum(embedding:face1 | embedding:face2))`
4. **Difference Function**:
- Syntax: `diff(arg1 | arg2 | ... | argN)`
- Computes the difference between multiple vector embeddings.
- Example: `diff(embedding:face1 | embedding:face2)`
### Notes on Arguments:
- Each function takes one or more arguments.
- An argument (`arg`) can be an embedding, a word, another function, or a quoted string.
- For functions that accept multiple arguments, they are separated by the `|` symbol.
## Examples
1. Add two embeddings and normalize the result:
```
norm(sum(cat | dog | horse | parrot))
```
2. Negate an embedding:
```
neg(embedding:body_vector)
```
3. King - Man + Woman = Queen:
```
sum(diff(king|man)|woman)
```
or
```
sum(king|neg(man)|woman)
```
```
- 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.
+10 -4
View File
@@ -1,9 +1,15 @@
from .nodes import (
BuildGif,
SpecialClipLoader,
)
from .nodes import BuildGif, PromptLangInspect, SpecialClipLoader
NODE_CLASS_MAPPINGS = {
"Build Gif": BuildGif,
"Special CLIP Loader": SpecialClipLoader,
"PromptLang Inspect": PromptLangInspect,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"Build Gif": "Build GIF (KepPromptLang)",
"Special CLIP Loader": "Special CLIP Loader (KepPromptLang)",
"PromptLang Inspect": "PromptLang Inspect",
}
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
View File
-100
View File
@@ -1,100 +0,0 @@
from abc import ABC, abstractmethod
from typing import Union
from torch import Tensor
from torch.nn import Embedding
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
class Action(ABC):
@property
@abstractmethod
def chars(self) -> list[str] | None:
pass
@property
@abstractmethod
def name(self) -> str:
pass
@property
@abstractmethod
def grammar(self) -> str:
"""
The grammar for this action. This is used to parse the action from the prompt.
:return:
"""
pass
@abstractmethod
def token_length(self) -> int:
"""
The length of the tokens that this action will add to the prompt.
:return:
"""
pass
@abstractmethod
def get_all_segments(self) -> list[PromptSegment]:
"""
Get all segments, including nested segments.
:return:
"""
pass
@abstractmethod
def get_result(self, embedding_module: Embedding) -> Tensor:
"""
Get the result of this action. This is called when the embeddings are being calculated.
:param embedding_module: The embedding module to use to get the base embeddings for tokens.
:return:
"""
pass
def depth_repr(self, depth: int = 1) -> str:
raise NotImplementedError()
class SingleArgAction(Action, ABC):
def get_all_segments(self) -> list[PromptSegment]:
segments = []
for seg_or_action in self.arg:
if isinstance(seg_or_action, Action):
segments.extend(seg_or_action.get_all_segments())
else:
segments.append(seg_or_action)
return segments
def __init__(self, arg: list[PromptSegment | Action]):
# TODO: Target is a list now... what does this mean for us..
self.arg = arg
def __repr__(self) -> str:
return f"{self.name}({self.arg})"
class MultiArgAction(Action, ABC):
def get_all_segments(self) -> list[PromptSegment]:
segments = []
for seg_or_action in self.base_segment:
if isinstance(seg_or_action, Action):
segments.extend(seg_or_action.get_all_segments())
else:
segments.append(seg_or_action)
for arg in self.args:
for seg_or_action in arg:
if isinstance(seg_or_action, Action):
segments.extend(seg_or_action.get_all_segments())
else:
segments.append(seg_or_action)
return segments
def __init__(
self,
base_segment: list[PromptSegment | Action],
args: list[list[Union[PromptSegment, Action]]],
):
self.base_segment = base_segment
self.args = args
+49
View File
@@ -0,0 +1,49 @@
from ..parser.registration import register_action
from .avg import AverageAction
from .diff import DiffAction
from .mult import MultiplyAction
from .nearest import NearestAction
from .neg import NegAction
from .noise import NoiseAction
from .norm import NormAction
from .pos_scale import PosScaleAction
from .post_pos import PostPosAction
from .project import ProjectAction, RejectAction
from .rand import RandAction
from .renorm import RenormAction
from .scale_dims import ScaleDims
from .set_dims import SetDims
from .slerp import SlerpAction
from .sum import SumAction
for _action in [
AverageAction,
DiffAction,
MultiplyAction,
NearestAction,
NegAction,
NoiseAction,
NormAction,
PosScaleAction,
PostPosAction,
ProjectAction,
RandAction,
RejectAction,
RenormAction,
ScaleDims,
SetDims,
SlerpAction,
SumAction,
]:
register_action(_action)
class _LerpAlias(AverageAction):
"""`lerp(a|b|t)` is sugar for `avg(a|b|t)`."""
display_name = "Lerp"
action_name = "lerp"
usage_examples = ["lerp(cat|dog|0.5)"]
register_action(_LerpAlias)
+56 -4
View File
@@ -1,11 +1,63 @@
from typing import Callable, List, TypeVar
import torch
from torch import Tensor
from torch.nn import Embedding
from custom_nodes.ClipStuff.lib.action.base import Action
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
from .base import Action
from .types import SegOrAction
from .weighted import WeightedGroup
def get_embedding(seg_or_action: SegOrAction, embedding_module: Embedding) -> Tensor:
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):
return seg_or_action.get_result(embedding_module)
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__}")
+49
View File
@@ -0,0 +1,49 @@
from typing import List
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length, parse_numeric_arg
from .base import MultiArgAction
from .types import SegOrAction
class AverageAction(MultiArgAction):
grammar = 'avg(" arg "|" arg "|" arg ")"'
display_name = "Average"
action_name = "avg"
description = "Performs a weighted average between two segments or actions. The recommended weight is 0 - 1."
usage_examples = [
"avg(The cat is|The dog is|0.5)",
"avg(Cat|Dog|0.5)",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 3:
raise ValueError("Average action expects exactly three arguments (2 vectors and a weight)")
self.first_arg = args[0]
self.second_arg = args[1]
self.parsed_weight = parse_numeric_arg(
args[2], action_name="Average", role="weight", cast=float
)
first_len = get_total_length(self.first_arg)
second_len = get_total_length(self.second_arg)
if first_len != second_len:
raise ValueError(
f"Average start and end arguments should have the same length. Got {first_len} and {second_len}"
)
if self.parsed_weight < 0 or self.parsed_weight > 1:
print(f"WARNING: Average weight should be between 0 and 1. Got {self.parsed_weight}")
def token_length(self) -> int:
return get_total_length(self.first_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
start = concat_embeddings(self.first_arg, embedding_module)
end = concat_embeddings(self.second_arg, embedding_module)
return start * (1 - self.parsed_weight) + end * self.parsed_weight
+73
View File
@@ -0,0 +1,73 @@
from abc import ABC, abstractmethod
from dataclasses import dataclass
from enum import Enum
from typing import List, Optional, Tuple, Union
from torch import Tensor
from torch.nn import Embedding
class ActionArity(Enum):
NONE = 0
SINGLE = 1
MULTI = 2
@dataclass
class PostModifiers:
"""Optional position-embedding tweaks an action can request for its token range.
`start_idx` / `end_idx` are filled in by the encoder once the action's position in
the final token stream is known.
"""
position_embed_scale: Optional[float] = None
bypass_pos_embed: bool = False
start_idx: int = 0
end_idx: int = 0
ActionResult = Union[Tensor, Tuple[Tensor, PostModifiers]]
# Tokenizer placeholder for the 2nd..Nth slots of a multi-token Action, so each
# row stays exactly max_length entries (required for comfy's per-position weight
# indexing). process_tokens drops these; the Action's tensor fills the slots.
ACTION_CONTINUATION = object()
class Action(ABC):
arity: ActionArity = ActionArity.NONE
display_name: str = ""
action_name: str = ""
description: str = ""
grammar: str = ""
usage_examples: List[str] = []
@abstractmethod
def __init__(self, *args, **kwargs) -> None: ...
@abstractmethod
def token_length(self) -> int: ...
@abstractmethod
def get_result(self, embedding_module: Embedding) -> ActionResult: ...
class SingleArgAction(Action, ABC):
arity = ActionArity.SINGLE
def __init__(self, arg: List):
self.arg = arg
def __repr__(self) -> str:
return f"{self.action_name}({self.arg})"
class MultiArgAction(Action, ABC):
arity = ActionArity.MULTI
def __init__(self, args: List[List]):
self.all_args = args
def __repr__(self) -> str:
joined = " | ".join(str(a) for a in self.all_args)
return f"{self.action_name}({joined})"
+29 -56
View File
@@ -1,66 +1,39 @@
import torch
from typing import List
from torch import Tensor
from torch.nn import Embedding
from custom_nodes.ClipStuff.lib.action.base import Action, MultiArgAction
from custom_nodes.ClipStuff.lib.actions.action_utils import get_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)* ")"'
name = "diff"
chars = ["-", "-"]
display_name = "Difference"
action_name = "diff"
description = (
"Subtracts the segments in the order they are given. "
"The first segment is subtracted from the second, then the third from the result, and so on."
)
usage_examples = [
"diff(The cat is|The dog is)",
"diff(Cat|Dog)",
"sum(diff(king|man)|woman)",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
self.base_arg = args[0]
self.additional_args = args[1:]
def token_length(self) -> int:
# Sum adds to the embeddings of the base segment, so the length is the length of the base segment
return sum(seg_or_action.token_length() for seg_or_action in self.base_segment)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the base segment
all_base_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.base_segment
]
result = torch.cat(all_base_embeddings, dim=1)
for arg in self.args:
all_arg_embeddings = [
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
]
arg_embedding = torch.cat(all_arg_embeddings, dim=1)
if (
arg_embedding.shape[-2] == 1
or result.shape[-2] == arg_embedding.shape[-2]
):
result = result.sub(arg_embedding)
else:
print(
"WARNING: shape mismatch when trying to apply sum, arg will be averaged"
)
result = result.sub(torch.mean(arg_embedding, dim=1, keepdim=True))
return 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
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.args))})"
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.args, Action):
target_repr = self.args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
+36
View File
@@ -0,0 +1,36 @@
from typing import List
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length, parse_numeric_arg
from .base import MultiArgAction
from .types import SegOrAction
class MultiplyAction(MultiArgAction):
grammar = 'mult(" arg+ ")"'
display_name = "Multiply"
action_name = "mult"
description = "Multiplies the provided segments or actions by the multiplier."
usage_examples = [
"mult(The cat is|2.5)",
"mult(Cat|-1)",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 2:
raise ValueError("Multiply action expects exactly two arguments")
self.target_arg = args[0]
self.parsed_multiplier = parse_numeric_arg(
args[1], action_name="Multiply", role="multiplier", cast=float
)
def token_length(self) -> int:
return get_total_length(self.target_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
return concat_embeddings(self.target_arg, embedding_module) * self.parsed_multiplier
+45
View File
@@ -0,0 +1,45 @@
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)
+14 -23
View File
@@ -1,32 +1,23 @@
import torch
from torch import Tensor
from torch.nn import Embedding
from custom_nodes.ClipStuff.lib.action.base import Action, SingleArgAction
from .action_utils import concat_embeddings, get_total_length
from .base import SingleArgAction
class NegAction(SingleArgAction):
grammar = 'neg(" arg+ ")"'
name = "neg"
chars = ["[", "]"]
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:
"""
Neg negates the embeddings of the base segment, so the length is the length of the base segment
:return:
"""
total_length = 0
for seg_or_action in self.arg:
total_length += seg_or_action.token_length()
return get_total_length(self.arg)
return total_length
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
all_embeddings = []
for seg_or_action in self.arg:
if isinstance(seg_or_action, Action):
all_embeddings.append(seg_or_action.get_result(embedding_module))
else:
all_embeddings.append(seg_or_action.get_embeddings(embedding_module))
target_embeddings = torch.cat(all_embeddings, dim=1)
return target_embeddings * -1
def get_result(self, embedding_module: Embedding) -> Tensor:
return concat_embeddings(self.arg, embedding_module) * -1
+34
View File
@@ -0,0 +1,34 @@
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
+15 -27
View File
@@ -1,37 +1,25 @@
import torch
from torch import Tensor
from torch.nn import Embedding
from custom_nodes.ClipStuff.lib.action.base import (
Action,
SingleArgAction,
)
from .action_utils import concat_embeddings, get_total_length
from .base import SingleArgAction
class NormAction(SingleArgAction):
grammar = 'norm(" arg+ ")"'
name = "norm"
chars = None
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:
"""
Norm normalizes the embeddings of the base segment, so the length is the length of the base segment
:return:
"""
total_length = 0
for seg_or_action in self.arg:
total_length += seg_or_action.token_length()
return total_length
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
all_embeddings = []
for seg_or_action in self.arg:
if isinstance(seg_or_action, Action):
target_embeddings = seg_or_action.get_result(embedding_module)
else:
target_embeddings = seg_or_action.get_embeddings(embedding_module)
all_embeddings.append(target_embeddings)
target_embeddings = torch.cat(all_embeddings, dim=1)
return torch.div(target_embeddings, torch.norm(target_embeddings, dim=-1, keepdim=True))
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))
+37
View File
@@ -0,0 +1,37 @@
from typing import List, Tuple
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length, parse_numeric_arg
from .base import MultiArgAction, PostModifiers
from .types import SegOrAction
class PosScaleAction(MultiArgAction):
grammar = 'posScale(" arg+ ")"'
display_name = "Positional Embedding Scale"
action_name = "posScale"
description = (
"Scales (multiplies) the positional embeddings of the provided segments or actions by the multiplier."
)
usage_examples = [
"A posScale(cat|1.5) on a rainy day",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 2:
raise ValueError("PosScale action expects exactly two arguments")
self.target_arg = args[0]
self.parsed_multiplier = parse_numeric_arg(
args[1], action_name="PosScale", role="multiplier", cast=float
)
def token_length(self) -> int:
return get_total_length(self.target_arg)
def get_result(self, embedding_module: Embedding) -> Tuple[Tensor, PostModifiers]:
target_embeddings = concat_embeddings(self.target_arg, embedding_module)
return target_embeddings, PostModifiers(position_embed_scale=self.parsed_multiplier)
+24
View File
@@ -0,0 +1,24 @@
from typing import Tuple
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length
from .base import PostModifiers, SingleArgAction
class PostPosAction(SingleArgAction):
grammar = 'postPos(" arg+ ")"'
display_name = "Ignore Positional Embeddings"
action_name = "postPos"
description = "Prevents positional embeddings from being applied to the provided segments or actions."
usage_examples = [
"A postPos(cat) on a rainy day",
]
def token_length(self) -> int:
return get_total_length(self.arg)
def get_result(self, embedding_module: Embedding) -> Tuple[Tensor, PostModifiers]:
return concat_embeddings(self.arg, embedding_module), PostModifiers(bypass_pos_embed=True)
+72
View File
@@ -0,0 +1,72 @@
from typing import List
import torch
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length
from .base import MultiArgAction
from .types import SegOrAction
def _direction(args: List[SegOrAction], embedding_module: Embedding) -> Tensor:
"""Mean unit direction of an arg's embeddings: [1, 1, hidden]."""
emb = concat_embeddings(args, embedding_module)
mean = emb.mean(dim=1, keepdim=True)
return torch.nn.functional.normalize(mean, dim=-1)
def _project(a: Tensor, b_hat: Tensor) -> Tensor:
coeff = (a * b_hat).sum(dim=-1, keepdim=True)
return coeff * b_hat
class ProjectAction(MultiArgAction):
grammar = 'proj(" arg "|" arg ")"'
display_name = "Project"
action_name = "proj"
description = "Projects the first argument onto the direction of the second (mean, unit-normalized)."
usage_examples = [
"proj(king|gender)",
"diff(style|proj(style|photorealistic))",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 2:
raise ValueError("proj expects exactly two arguments: proj(a|b)")
self.a_arg = args[0]
self.b_arg = args[1]
def token_length(self) -> int:
return get_total_length(self.a_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
a = concat_embeddings(self.a_arg, embedding_module)
return _project(a, _direction(self.b_arg, embedding_module))
class RejectAction(MultiArgAction):
grammar = 'reject(" arg "|" arg ")"'
display_name = "Reject"
action_name = "reject"
description = "Removes the component of the first argument along the direction of the second (a - proj(a|b))."
usage_examples = [
"reject(anime girl|anime)",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 2:
raise ValueError("reject expects exactly two arguments: reject(a|b)")
self.a_arg = args[0]
self.b_arg = args[1]
def token_length(self) -> int:
return get_total_length(self.a_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
a = concat_embeddings(self.a_arg, embedding_module)
return a - _project(a, _direction(self.b_arg, embedding_module))
+53
View File
@@ -0,0 +1,53 @@
from typing import List
import torch
from torch import Tensor
from torch.nn import Embedding
from .action_utils import parse_numeric_arg
from .base import MultiArgAction
from .types import SegOrAction
class RandAction(MultiArgAction):
grammar = 'rand(" arg ")"'
display_name = "Random Embedding"
action_name = "rand"
description = (
"Returns a random embedding of the specified token length, "
"with the values optionally bounded by the second and third arguments."
)
usage_examples = [
"A rand(1) cat",
"A rand(1|-1|1) cat",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) not in (1, 3):
raise ValueError("Random action expects exactly one or three arguments")
self.parsed_token_length = parse_numeric_arg(
args[0], action_name="Random", role="first argument (token length)", cast=int
)
if len(args) == 3:
self.range_min = parse_numeric_arg(
args[1], action_name="Random", role="second argument (min)", cast=int
)
self.range_max = parse_numeric_arg(
args[2], action_name="Random", role="third argument (max)", cast=int
)
if self.range_min > self.range_max:
raise ValueError("Random action min must be <= max")
else:
self.range_min = 0
self.range_max = 1
def token_length(self) -> int:
return self.parsed_token_length
def get_result(self, embedding_module: Embedding) -> Tensor:
return torch.empty(
1, self.parsed_token_length, embedding_module.embedding_dim
).uniform_(self.range_min, self.range_max)
+37
View File
@@ -0,0 +1,37 @@
from typing import List
import torch
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length
from .base import MultiArgAction
from .types import SegOrAction
class RenormAction(MultiArgAction):
grammar = 'renorm(" arg "|" arg ")"'
display_name = "Renormalize"
action_name = "renorm"
description = "Rescales the first argument so each token's L2 norm matches the (mean) L2 norm of the reference."
usage_examples = [
"renorm(sum(king|neg(man)|woman)|queen)",
]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 2:
raise ValueError("renorm expects exactly two arguments: renorm(a|ref)")
self.a_arg = args[0]
self.ref_arg = args[1]
def token_length(self) -> int:
return get_total_length(self.a_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
a = concat_embeddings(self.a_arg, embedding_module)
ref = concat_embeddings(self.ref_arg, embedding_module)
a_norm = torch.norm(a, dim=-1, keepdim=True).clamp(min=1e-8)
ref_norm = torch.norm(ref, dim=-1, keepdim=True).mean()
return a * (ref_norm / a_norm)
+68
View File
@@ -0,0 +1,68 @@
from typing import List, Tuple
from torch import Tensor
from torch.nn import Embedding
from ..parser.prompt_segment import PromptSegment
from .action_utils import concat_embeddings, get_total_length
from .base import Action, MultiArgAction
from .types import SegOrAction
class ScaleDims(MultiArgAction):
grammar = 'scaleDims(" arg ("|" arg)* ")"'
display_name = "Scale Dimensions"
action_name = "scaleDims"
description = "Scales the specified dimensions of the input embeddings by the specified amount"
usage_examples = [
"The scaleDims(cat|4,1.5|76,1.2) is happy",
]
def __init__(self, args: List[List[SegOrAction]]):
super().__init__(args)
self.base_arg = args[0]
self.scale_args: List[Tuple[int, float]] = _parse_dim_value_pairs(args[1:], action_name="ScaleDims")
def token_length(self) -> int:
return get_total_length(self.base_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
embeddings = concat_embeddings(self.base_arg, embedding_module)
for dim, scale in self.scale_args:
embeddings[0, :, dim] *= scale
return embeddings
def _parse_dim_value_pairs(
args: List[List[SegOrAction]],
*,
action_name: str,
) -> List[Tuple[int, float]]:
"""Parse args of the form `<dim>,<value>` into `(int, float)` pairs.
Used by both scaleDims and setDims.
"""
pairs: List[Tuple[int, float]] = []
for arg in args:
if isinstance(arg, Action):
raise ValueError(f"{action_name} args must be in the format <dim>,<value> but got an action")
if len(arg) != 1:
raise ValueError(f"{action_name} args must be a single segment of <dim>,<value>")
seg = arg[0]
assert isinstance(seg, PromptSegment)
if "," not in seg.text:
raise ValueError(f"{action_name} args must be <dim>,<value> but got: {seg.text!r}")
dim_str, value_str = seg.text.split(",", 1)
try:
dim = int(dim_str)
except ValueError:
raise ValueError(f"{action_name} dim must be an integer; got {dim_str!r}")
try:
value = float(value_str)
except ValueError:
raise ValueError(f"{action_name} value must be a float; got {value_str!r}")
pairs.append((dim, value))
return pairs
+34
View File
@@ -0,0 +1,34 @@
from typing import List, Tuple
from torch import Tensor
from torch.nn import Embedding
from .action_utils import concat_embeddings, get_total_length
from .base import MultiArgAction
from .scale_dims import _parse_dim_value_pairs
from .types import SegOrAction
class SetDims(MultiArgAction):
grammar = 'setDims(" arg ("|" arg)* ")"'
display_name = "Set Dimensions"
action_name = "setDims"
description = "Sets the specified dimensions of the input embeddings to the specified value"
usage_examples = [
"The setDims(cat|4, -0.01253|76, 1.2) is happy",
]
def __init__(self, args: List[List[SegOrAction]]):
super().__init__(args)
self.base_arg = args[0]
self.value_args: List[Tuple[int, float]] = _parse_dim_value_pairs(args[1:], action_name="SetDims")
def token_length(self) -> int:
return get_total_length(self.base_arg)
def get_result(self, embedding_module: Embedding) -> Tensor:
embeddings = concat_embeddings(self.base_arg, embedding_module)
for dim, value in self.value_args:
embeddings[0, :, dim] = value
return embeddings
+52
View File
@@ -0,0 +1,52 @@
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)
+24 -86
View File
@@ -1,96 +1,34 @@
from typing import Union
from typing import List
import torch
from torch import Tensor
from torch.nn import Embedding
from custom_nodes.ClipStuff.lib.actions.action_utils import get_embedding
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
from custom_nodes.ClipStuff.lib.action.base import Action
from .action_utils import add_with_broadcast, concat_embeddings
from .base import MultiArgAction
from .types import SegOrAction
class SumAction(Action):
grammar = 'sum(" arg ("|" arg)* ")"'
name = "sum"
chars = ["+", "+"]
class SumAction(MultiArgAction):
grammar = 'sum(" arg ("|" arg)+ ")"'
def __init__(
self,
base_segment: list[PromptSegment | Action],
args: list[list[Union[PromptSegment, Action]]],
):
self.base_segment = base_segment
self.args = args
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:
# Sum adds to the embeddings of the base segment, so the length is the length of the base segment
return sum(seg_or_action.token_length() for seg_or_action in self.base_segment)
def get_all_segments(self) -> list[PromptSegment]:
segments = []
for seg_or_action in self.base_segment:
if isinstance(seg_or_action, Action):
segments.extend(seg_or_action.get_all_segments())
else:
segments.append(seg_or_action)
for arg in self.args:
for seg_or_action in arg:
if isinstance(seg_or_action, Action):
segments.extend(seg_or_action.get_all_segments())
else:
segments.append(seg_or_action)
return segments
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the base segment
all_base_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.base_segment
]
result = torch.cat(all_base_embeddings, dim=1)
for arg in self.args:
all_arg_embeddings = [
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
]
arg_embedding = torch.cat(all_arg_embeddings, dim=1)
if (
arg_embedding.shape[-2] == 1
or result.shape[-2] == arg_embedding.shape[-2]
):
result = result.add(arg_embedding)
else:
print(
"WARNING: shape mismatch when trying to apply sum, arg will be averaged"
)
result = result.add(
torch.mean(arg_embedding, dim=1, keepdim=True)
)
return 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
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.args))})"
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.args, Action):
target_repr = self.args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
+4 -3
View File
@@ -1,6 +1,7 @@
from typing import Union
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
from custom_nodes.ClipStuff.lib.action.base import Action
from ..parser.prompt_segment import PromptSegment
from .base import Action
from .weighted import WeightedGroup
SegOrAction = Union[PromptSegment, Action]
SegOrAction = Union[PromptSegment, Action, WeightedGroup]
+22 -5
View File
@@ -1,6 +1,23 @@
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
import torch
def batch_size_info(batch: list[SegOrAction]):
for segment in batch:
print("Token Len: " + str(segment.token_length()))
print(segment.depth_repr())
def slerp(val: float, low: torch.Tensor, high: torch.Tensor, epsilon: float = 1e-5) -> torch.Tensor:
"""Spherical linear interpolation between two tensors along the last dim."""
val_t = torch.tensor(val, dtype=torch.float32, device=low.device).clamp(0, 1)
low_norm = low / torch.norm(low, dim=-1, keepdim=True)
high_norm = high / torch.norm(high, dim=-1, keepdim=True)
dot = (low_norm * high_norm).sum(-1, keepdim=True).clamp(-1, 1)
omega = torch.acos(dot)
sin_omega = torch.sin(omega)
scale_low = torch.sin((1.0 - val_t) * omega) / (sin_omega + epsilon)
scale_high = torch.sin(val_t * omega) / (sin_omega + epsilon)
# Fall back to linear interp where the angle is too small for stable slerp.
close = sin_omega < epsilon
scale_low = torch.where(close, 1.0 - val_t, scale_low)
scale_high = torch.where(close, val_t, scale_high)
return scale_low * low + scale_high * high
+15
View File
@@ -0,0 +1,15 @@
from typing import List
class WeightedGroup:
"""A group of segments/actions sharing an attention weight (the `(text:1.2)` syntax)."""
def __init__(self, items: List, weight: float):
self.items = items
self.weight = weight
def token_length(self) -> int:
return sum(item.token_length() for item in self.items)
def __repr__(self) -> str:
return f"({self.items}:{self.weight})"
-25
View File
@@ -1,25 +0,0 @@
{
"_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
}
+108 -181
View File
@@ -1,202 +1,129 @@
import contextlib
import os
"""DSL-aware CLIP text encoders.
The tokenizer emits ComfyUI's native `(token, weight)` format with one twist:
a `token` can also be a lazily-evaluated `Action`. We resolve those to tensors
here in `process_tokens` (where the embedding module is available) and delegate
everything else — embedding lookup, mask building, splice — to the stock
`SDClipModel.process_tokens`.
`posScale` / `postPos` actions return a `PostModifiers` alongside their tensor.
ComfyUI's `CLIPTextModel_.forward` adds the position embedding inline whenever
`embeds` is supplied, so we pre-bake `(modified - default)` into `embeds` such
that the transformer's add nets to `+ modified`.
"""
import dataclasses
from typing import List
import torch
from transformers import CLIPTextConfig, modeling_utils
from torch import Tensor
from torch.nn import Embedding
from comfy import model_management
import comfy.ops
from custom_nodes.ClipStuff.lib.action.base import Action
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
from custom_nodes.ClipStuff.lib.fun_clip_stuff import PromptLangTextModel
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
from comfy import sd1_clip, sdxl_clip
from .actions.base import ACTION_CONTINUATION, Action, PostModifiers
# Methods with no comment can be assumed to be the same as comfy.sd1_clip.SD1ClipModel
class PromptLangClipModel(torch.nn.Module):
"""Uses the CLIP transformer encoder for text (from huggingface)"""
LAYERS = [
"last",
"pooled",
"hidden"
]
def __init__(self, version="openai/clip-vit-large-patch14", device="cpu", max_length=77,
freeze=True, layer="last", layer_idx=None, textmodel_json_config=None,
textmodel_path=None): # clip-vit-base-patch32
super().__init__()
assert layer in self.LAYERS
self.num_layers = 12
if textmodel_path is not None:
# Our transformer
self.transformer = PromptLangTextModel.from_pretrained(textmodel_path)
else:
if textmodel_json_config is None:
# TODO: Maybe re-use clip config?
# Config could come from cond_stage_model.transformer.config
# Copied clip_config
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "clip_config.json")
config = CLIPTextConfig.from_json_file(textmodel_json_config)
self.num_layers = config.num_hidden_layers
with comfy.ops.use_comfy_ops():
with modeling_utils.no_init_weights():
# Our transformer
self.transformer = PromptLangTextModel(config)
self.max_length = max_length
if freeze:
self.freeze()
self.layer = layer
self.layer_idx = None
self.empty_tokens = [[49406] + [49407] * 76]
self.text_projection = None
self.layer_norm_hidden_state = True
if layer == "hidden":
assert layer_idx is not None
assert abs(layer_idx) <= self.num_layers
self.clip_layer(layer_idx)
self.layer_default = (self.layer, self.layer_idx)
def freeze(self):
self.transformer = self.transformer.eval()
# self.train = disabled_train
for param in self.parameters():
param.requires_grad = False
def clip_layer(self, layer_idx):
if abs(layer_idx) >= self.num_layers:
self.layer = "last"
else:
self.layer = "hidden"
self.layer_idx = layer_idx
def reset_clip_layer(self):
self.layer = self.layer_default[0]
self.layer_idx = self.layer_default[1]
# Completely changed to support Segments and actions
def set_up_textual_embeddings(self, tokens: list[list[SegOrAction]], current_embeds):
next_new_token = token_dict_size = current_embeds.weight.shape[0] - 1
embedding_weights = []
# For each batch
for batch in tokens:
for seg_or_action in batch:
if isinstance(seg_or_action, Action):
segments = seg_or_action.get_all_segments()
else:
segments = [seg_or_action]
for segment in segments:
tokens_temp = []
segment_length = segment.token_length()
for tid_or_tensor in segment.tokens:
if isinstance(tid_or_tensor, int):
if tid_or_tensor == token_dict_size: # Is EOS token
tid_or_tensor = -1 # Set to -1 so that it can be replaced with the EOS token later
tokens_temp += [tid_or_tensor]
else:
if tid_or_tensor.shape[0] == current_embeds.weight.shape[1]:
embedding_weights += [tid_or_tensor]
tokens_temp += [next_new_token]
next_new_token += 1
else:
print("WARNING: shape mismatch when trying to apply embedding, embedding will be ignored",
tid_or_tensor.shape[0], current_embeds.weight.shape[1])
if len(tokens_temp) < segment_length:
# Pretty sure this is only needed if the embedding is not the same size as the CLIP embedding
print("WARNING: segment length mismatch, padding with EOS token")
tokens_temp.extend([self.empty_tokens[0][-1] * (segment_length - len(tokens_temp))])
segment.tokens = tokens_temp
n = token_dict_size
if len(embedding_weights) > 0:
# Create new embedding, with size of current embedding + number of new embeddings
new_embedding = torch.nn.Embedding(next_new_token + 1, current_embeds.weight.shape[1],
device=current_embeds.weight.device, dtype=current_embeds.weight.dtype)
# Copy current embedding weights to new embedding
new_embedding.weight[:token_dict_size] = current_embeds.weight[:-1]
# Add new embeddings
for embed in embedding_weights:
new_embedding.weight[n] = embed
n += 1
# Set re-add the EOS token
new_embedding.weight[n] = current_embeds.weight[-1] # EOS embedding
self.transformer.set_input_embeddings(new_embedding)
class PromptLangSDClipModel(sd1_clip.SDClipModel):
def process_tokens(self, tokens, device): # type: ignore[override]
embedding_module = self.transformer.get_input_embeddings()
resolved: List[list] = []
pos_modifiers_per_batch: List[List[PostModifiers]] = []
for batch in tokens:
for seg_or_action in batch:
if isinstance(seg_or_action, Action):
segments = seg_or_action.get_all_segments()
row: list = []
modifiers: List[PostModifiers] = []
position = 0
for entry in batch:
if entry is ACTION_CONTINUATION:
# Slot already accounted for by the preceding Action's `position += length`.
continue
if isinstance(entry, Action):
length = entry.token_length()
result = entry.get_result(embedding_module)
if isinstance(result, tuple):
tensor, mods = result
modifiers.append(
dataclasses.replace(mods, start_idx=position, end_idx=position + length)
)
else:
tensor = result
row.append(tensor)
position += length
else:
segments = [seg_or_action]
row.append(entry)
position += 1
resolved.append(row)
pos_modifiers_per_batch.append(modifiers)
for segment in segments:
for tokenIdx in range(len(segment.tokens)):
if segment.tokens[tokenIdx] == -1:
segment.tokens[tokenIdx] = n
embeds, attention_mask, num_tokens, embeds_info = super().process_tokens(resolved, device)
# Support our set_up_textual_embeddings which modifies the input embeddings
def forward(self, tokens):
backup_embeds = self.transformer.get_input_embeddings()
device = backup_embeds.weight.device
self.set_up_textual_embeddings(tokens, backup_embeds)
# tokens = torch.LongTensor(tokens).to(device)
if any(pos_modifiers_per_batch):
embeds = _apply_pos_modifiers(
embeds, pos_modifiers_per_batch, self._get_position_embedding()
)
if backup_embeds.weight.dtype != torch.float32:
precision_scope = torch.autocast
else:
precision_scope = contextlib.nullcontext
return embeds, attention_mask, num_tokens, embeds_info
with precision_scope(model_management.get_autocast_device(device)):
outputs = self.transformer(input_ids=tokens, output_hidden_states=self.layer == "hidden")
self.transformer.set_input_embeddings(backup_embeds)
def _get_position_embedding(self) -> Embedding:
"""Isolated so a ComfyUI internal layout change only needs one fix."""
return self.transformer.text_model.embeddings.position_embedding
if self.layer == "last":
z = outputs.last_hidden_state
elif self.layer == "pooled":
z = outputs.pooler_output[:, None, :]
class PromptLangSDXLClipG(sdxl_clip.SDXLClipG, PromptLangSDClipModel):
"""SDXL's larger CLIP-G text encoder, with our DSL-aware process_tokens."""
def _apply_pos_modifiers(
embeds: Tensor,
pos_modifiers_per_batch: List[List[PostModifiers]],
position_embedding: Embedding,
) -> Tensor:
seq_len = embeds.shape[1]
pos_weights = position_embedding.weight[:seq_len].to(device=embeds.device, dtype=embeds.dtype)
out = embeds.clone()
for batch_idx, modifiers in enumerate(pos_modifiers_per_batch):
for mod in modifiers:
default_slice = pos_weights[mod.start_idx:mod.end_idx]
if mod.bypass_pos_embed:
modified_slice = torch.zeros_like(default_slice)
elif mod.position_embed_scale is not None:
modified_slice = default_slice * float(mod.position_embed_scale)
else:
z = outputs.hidden_states[self.layer_idx]
if self.layer_norm_hidden_state:
z = self.transformer.text_model.final_layer_norm(z)
continue
pooled_output = outputs.pooler_output
if self.text_projection is not None:
pooled_output = pooled_output.to(self.text_projection.device) @ self.text_projection
return z.float(), pooled_output.float()
# The transformer will add `default_slice` back; net effect is `+ modified_slice`.
out[batch_idx, mod.start_idx:mod.end_idx] += modified_slice - default_slice
def encode(self, tokens):
return self(tokens)
return out
def load_sd(self, sd):
return self.transformer.load_state_dict(sd, strict=False)
# Changed from comfy.sd1_clip.ClipTokenWeightEncoder
# Changed to use PromptSegments
def encode_token_weights(self, prompt_segments: list[list[SegOrAction]]):
to_encode = [[PromptSegment(text="_Empty Batch_", tokens=self.empty_tokens[0])]]
for batch in prompt_segments:
to_encode.append(batch)
class PromptLangSD1ClipModel(sd1_clip.SD1ClipModel):
def __init__(self, device="cpu", dtype=None, model_options=None, **kwargs):
super().__init__(
device=device,
dtype=dtype,
model_options=model_options or {},
clip_name="l",
clip_model=PromptLangSDClipModel,
**kwargs,
)
out, pooled = self.encode(to_encode)
z_empty = out[0:1]
if pooled.shape[0] > 1:
first_pooled = pooled[1:2]
else:
first_pooled = pooled[0:1]
output = []
for k in range(1, out.shape[0]):
z = out[k:k + 1]
# for i in range(len(z)):
# for j in range(len(z[i])):
# weight = token_dicts[k - 1][j][0].weight
# z[i][j] = (z[i][j] - z_empty[0][j]) * weight + z_empty[0][j]
output.append(z)
if (len(output) == 0):
return z_empty.cpu(), first_pooled.cpu()
return torch.cat(output, dim=-2).cpu(), first_pooled.cpu()
class PromptLangSDXLClipModel(sdxl_clip.SDXLClipModel):
def __init__(self, device="cpu", dtype=None, model_options=None) -> None:
torch.nn.Module.__init__(self)
opts = model_options or {}
self.clip_l = PromptLangSDClipModel(
layer="hidden",
layer_idx=-2,
device=device,
dtype=dtype,
layer_norm_hidden_state=False,
model_options=opts,
)
self.clip_g = PromptLangSDXLClipG(device=device, dtype=dtype, model_options=opts)
self.dtypes = {dtype} if dtype is not None else set()
-179
View File
@@ -1,179 +0,0 @@
from typing import Optional, Tuple, Union
import torch
from transformers import CLIPTextConfig
from transformers.modeling_outputs import BaseModelOutputWithPooling
from transformers.models.clip.modeling_clip import _expand_mask, CLIPTextEmbeddings, CLIPTextTransformer, \
CLIPTextModel
from custom_nodes.ClipStuff.lib.action.base import Action
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
def slerp(val, low, high):
low = low.unsqueeze(0)
high = high.unsqueeze(0)
low_norm = low/torch.norm(low, dim=1, keepdim=True)
high_norm = high/torch.norm(high, dim=1, keepdim=True)
omega = torch.acos((low_norm*high_norm).sum(1))
so = torch.sin(omega)
res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high
return res
class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
def __init__(self, config: CLIPTextConfig):
super().__init__(config)
def forward(
self,
input_dicts: Optional[list[list[SegOrAction]]] = None,
input_ids: Optional[torch.LongTensor] = None,
position_ids: Optional[torch.LongTensor] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
) -> torch.Tensor:
if input_dicts is None:
raise ValueError("You have to specify input_dicts")
batches = []
for batch_idx, batch in enumerate(input_dicts):
results = []
for seg_or_action in batch:
if isinstance(seg_or_action, Action):
results.append(seg_or_action.get_result(self.token_embedding))
else:
results.append(seg_or_action.get_embeddings(self.token_embedding))
batches.append(results)
seq_length = batches[0][0].shape[-2]
if position_ids is None:
position_ids = self.position_ids[:, :seq_length]
embeds = []
for batch in batches:
if len(batch) == 1:
embeds.append(batch[0])
else:
embeds.append(torch.cat(batch, dim=-2))
position_embeddings = self.position_embedding(position_ids)
embeddings = torch.cat(embeds, dim=0) + position_embeddings
return embeddings
class PrompLangCLIPTextTransformer(CLIPTextTransformer):
def __init__(self, config: CLIPTextConfig):
super().__init__(config)
self.embeddings = PromptLangCLIPTextEmbeddings(config)
def forward(
self,
input_ids: Optional[list[list[SegOrAction]]] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
) -> Union[Tuple, BaseModelOutputWithPooling]:
r"""
Returns:
"""
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
)
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
if input_ids is None:
raise ValueError("You have to specify input_ids")
# input_shape = input_ids.size()
# input_ids = input_ids.view(-1, input_shape[-1])
hidden_states = self.embeddings(input_dicts=input_ids)
bsz = len(input_ids)
# TODO: Properly gather this
seq_len = 77
# 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,
)
# This is necessary to pass the PromptLangCLIPTextTransformer
class PromptLangTextModel(CLIPTextModel):
def __init__(self, config: CLIPTextConfig):
super().__init__(config)
self.text_model = PrompLangCLIPTextTransformer(config)
def forward(
self,
input_ids: Optional[list[list[SegOrAction]]] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
) -> Union[Tuple, BaseModelOutputWithPooling]:
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
return self.text_model(
input_ids=input_ids,
attention_mask=attention_mask,
position_ids=position_ids,
output_attentions=output_attentions,
output_hidden_states=output_hidden_states,
return_dict=return_dict,
)
+68
View File
@@ -0,0 +1,68 @@
"""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)
+24 -14
View File
@@ -1,28 +1,38 @@
grammar = """
?start: item+
grammar = r"""
?start: stmt+
?stmt: assign
| item
assign: "$" NAME "=" arg ";"
item: embedding
| WORD
| function
| generic_function
| QUOTED_STRING
| weighted
| ref
function: sum_function
| neg_function
| norm_function
| diff_function
generic_function: FUNC_NAME "(" arg ("|" arg)* ")"
sum_function: "sum(" arg ("|" arg)* ")"
neg_function: "neg(" arg ")"
norm_function: "norm(" arg ")"
diff_function: "diff(" arg ("|" arg)* ")"
weighted: "(" arg ":" SIGNED_NUMBER ")"
ref: "$" NAME
arg: item+
embedding: "embedding:" WORD
WORD: /[A-Za-z0-9,_-]+/
QUOTED_STRING: /"([^"\\\]*(\\\.[^"\\\]*)*)"|'([^'\\\]*(\\\.[^'\\\]*)*)'/
// 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
"""
+15 -16
View File
@@ -1,4 +1,4 @@
from typing import Union
from typing import List, Union
import torch
from torch import Tensor
@@ -6,26 +6,25 @@ from torch.nn import Embedding
class PromptSegment:
def __init__(self, text: str, tokens: list[Union[int, Tensor]]):
"""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):
return f'"{self.text}"{self.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):
def token_length(self) -> int:
return len(self.tokens)
def get_embeddings(self, embedding_module: Embedding) -> Tensor:
tensors = torch.LongTensor(self.tokens).to(torch.device('cpu'))
unsqueezed_tensors = tensors.unsqueeze(0)
return embedding_module(unsqueezed_tensors)
"""Look up embeddings for plain int tokens.
def depth_repr(self, depth=1):
out = f'"{self.text}"('
cleaned_tokens = list(map(lambda x: str(x) if isinstance(x, int) else "EMBD", self.tokens))
out += ", ".join(cleaned_tokens)
out += ")"
return out
Inline TI tensors aren't handled here; the encoder splices them in at a higher level.
"""
ids = torch.LongTensor([t for t in self.tokens if isinstance(t, int)]).to(
embedding_module.weight.device
)
return embedding_module(ids.unsqueeze(0))
+18
View File
@@ -0,0 +1,18 @@
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]
+61 -41
View File
@@ -1,61 +1,81 @@
from lark import Transformer, Token
from typing import List
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.ClipStuff.lib.action.base import Action
from custom_nodes.ClipStuff.lib.actions.diff import DiffAction
from custom_nodes.ClipStuff.lib.parser.utils import build_prompt_segment
from custom_nodes.ClipStuff.lib.actions.neg import NegAction
from custom_nodes.ClipStuff.lib.actions.norm import NormAction
from custom_nodes.ClipStuff.lib.actions.sum import SumAction
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
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):
# def WORD(self, items):
# return items
"""Maps the Lark parse tree into a flat list of PromptSegments and Actions."""
def __init__(self, tokenizer: SD1Tokenizer):
def __init__(self, tokenizer: SDTokenizer):
super().__init__()
self.tokenizer = tokenizer
self.vars: dict = {}
def item(self, items: list[Token]):
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):
return item
if isinstance(item, PromptSegment):
if isinstance(item, (Action, PromptSegment, WeightedGroup)):
return item
if item.type == "WORD":
return build_prompt_segment(str(item), self.tokenizer)
elif item.type == "QUOTED_STRING":
# Remove the quotes
if item.type == "QUOTED_STRING":
# Strip surrounding quotes, unescape \" and \'.
unquoted = item[1:-1]
# Replace escaped quotes with quotes
unescaped = unquoted.replace("\\\"", "\"").replace("\\\'", "\'")
unescaped = unquoted.replace('\\"', '"').replace("\\'", "'")
return build_prompt_segment(unescaped, self.tokenizer)
elif item.type == "embedding":
return build_prompt_segment(item, self.tokenizer)
elif item.type == "function":
return item
else:
raise Exception("Unknown item type: " + str(item.type))
raise ValueError(f"Unknown item type: {item.type}")
def arg(self, items):
return items
def embedding(self, items):
return build_prompt_segment(f'{self.tokenizer.embedding_identifier}{items[0]}', self.tokenizer)
def weighted(self, items):
arg_items, weight_token = items
return WeightedGroup(arg_items, float(weight_token))
def function(self, items):
for item in items:
if item.data == 'sum_function':
return SumAction(item.children[0][:], item.children[1:][:])
elif item.data == 'neg_function':
return NegAction(item.children[0])
elif item.data == 'norm_function':
return NormAction(item.children[0])
elif item.data == 'diff_function':
return DiffAction(item.children[0][:], item.children[1:][:])
else:
raise Exception("Unknown function type: " + str(item.data))
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}")
+15 -12
View File
@@ -1,28 +1,30 @@
from lark import Token
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
from comfy.sd1_clip import SDTokenizer
from .prompt_segment import PromptSegment
def flatten_tree(tree):
if isinstance(tree, Token):
return [str(tree)]
else:
return [str(tree.data)] + sum([flatten_tree(child) for child in tree.children], [])
return [str(tree.data)] + sum([flatten_tree(child) for child in tree.children], [])
def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment:
split_text = text.split(" ")
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 split_text:
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')
get_embed_ret = tokenizer._try_get_embedding(embedding_name)
embedding = get_embed_ret[0]
leftover = get_embed_ret[1]
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)
@@ -33,6 +35,7 @@ def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment:
word = leftover
else:
continue
# Strip the SOT/EOT bracketing tokens added by the underlying CLIP tokenizer.
tokens.extend(tokenizer.tokenizer(word)["input_ids"][1:-1])
return PromptSegment(text, tokens)
+109 -53
View File
@@ -1,66 +1,122 @@
"""DSL-aware tokenizers.
Override `tokenize_with_weights` to parse our DSL and emit ComfyUI's native
`List[List[(token, weight)]]` format, where `token` is an int id, an inline
TI tensor, a lazily-evaluated `Action`, or `ACTION_CONTINUATION`.
Row alignment matters: comfy's stock `encode_token_weights` indexes weights by
post-transformer position, so each row must be exactly `max_length` entries.
A multi-slot Action is therefore emitted as one `(action, w)` entry followed by
`(ACTION_CONTINUATION, w)` placeholders; `process_tokens` drops the placeholders
and the action's tensor expands to fill those slots.
"""
from typing import Dict, Iterable, List, Tuple, Union
from lark import Tree
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
from comfy.sd1_clip import SD1Tokenizer, SDTokenizer
from custom_nodes.ClipStuff.lib.parser import PromptParser
from custom_nodes.ClipStuff.lib.parser.transformer import PromptTransformer
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
from .actions.base import ACTION_CONTINUATION, Action
from .actions.weighted import WeightedGroup
from .parser import PromptParser
from .parser.prompt_segment import PromptSegment
from .parser.transformer import PromptTransformer
class PromptLangTokenizer(SD1Tokenizer):
def __init__(self, tokenizer_path=None, max_length=77, pad_with_end=True, embedding_directory=None, embedding_size=768, embedding_key='clip_l', special_tokens=None):
super().__init__(tokenizer_path, max_length, pad_with_end, embedding_directory, embedding_size, embedding_key)
# Side-effect import: registers all built-in actions with the parser.
from . import actions # noqa: F401
"""
Doesn't actually tokenize...
Returns batches of segments and actions
:return: List of list(batches) of segments and actions
"""
def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs) -> list[list[SegOrAction]]:
if self.pad_with_end:
pad_token = self.end_token
else:
pad_token = 0
TokenEntry = Tuple[Union[int, "Action", object], float]
parsed_prompt = PromptParser.parse(text)
parsed_actions = PromptTransformer(self).transform(parsed_prompt)
# reshape token array to CLIP input size
batched_segments = []
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token])]
# batched_segments.append(batch)
batch_size = 1
if isinstance(parsed_actions, Tree):
segments_to_process = parsed_actions.children
else:
segments_to_process = [parsed_actions]
for segment in segments_to_process:
num_tokens = segment.token_length()
# determine if we're going to try and keep the tokens in a single batch
is_large = num_tokens >= self.max_word_length
def _flatten(item, weight: float) -> Iterable[TokenEntry]:
"""Walk the parsed item tree, yielding one (token, weight) entry per output slot."""
if isinstance(item, WeightedGroup):
for sub in item.items:
yield from _flatten(sub, weight * item.weight)
elif isinstance(item, Action):
yield (item, weight)
for _ in range(item.token_length() - 1):
yield (ACTION_CONTINUATION, weight)
elif isinstance(item, PromptSegment):
for tok in item.tokens:
yield (tok, weight)
else:
raise TypeError(f"Unexpected parse item {item!r} ({type(item).__name__})")
# If the segment is too large to fit in a single batch, pad the current batch and start a new one
if num_tokens + batch_size > self.max_length - 1:
remaining_length = self.max_length - batch_size - 1 # -1 for end token
# Pad batch
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length - 1))
batched_segments.append(batch)
# start new batch
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token]), segment]
batch_size = num_tokens + 1 # +1 for start token
continue
class PromptLangSDTokenizer(SDTokenizer):
def tokenize_with_weights( # type: ignore[override]
self, text: str, return_word_ids: bool = False, **kwargs
) -> List[List[TokenEntry]]:
# SDXL passes a pre-parsed tree to avoid re-running Lark per sub-tokenizer.
tree = kwargs.pop("_parsed_tree", None) or PromptParser.parse(text)
return self._batch_from_tree(tree)
# Since the segment fits in the current batch, add it
batch.append(segment)
batch_size += num_tokens
def _batch_from_tree(self, tree) -> List[List[TokenEntry]]:
pad_token = self.end_token if self.pad_with_end else 0
# 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)
parsed = PromptTransformer(self).transform(tree)
items = parsed.children if isinstance(parsed, Tree) else [parsed]
# assign stmts return None (they only populate the transformer's var table).
items = [i for i in items if i is not None]
# for batch in batched_segments:
# batch_size_info(batch)
batches: List[List[TokenEntry]] = []
current: List[TokenEntry] = [(self.start_token, 1.0)]
return batched_segments
def close(row: List[TokenEntry]) -> None:
row.append((self.end_token, 1.0))
row.extend([(pad_token, 1.0)] * (self.max_length - len(row)))
batches.append(row)
for item in items:
entries = list(_flatten(item, 1.0))
if len(current) + len(entries) > self.max_length - 1:
close(current)
current = [(self.start_token, 1.0)]
current.extend(entries)
close(current)
return batches
class PromptLangSD1Tokenizer(SD1Tokenizer):
def __init__(self, embedding_directory=None, tokenizer_data=None, clip_name="l", tokenizer=PromptLangSDTokenizer):
super().__init__(
embedding_directory=embedding_directory,
tokenizer_data=tokenizer_data or {},
clip_name=clip_name,
tokenizer=tokenizer,
)
class PromptLangSDXLClipGTokenizer(PromptLangSDTokenizer):
def __init__(self, tokenizer_path=None, embedding_directory=None, tokenizer_data=None):
super().__init__(
tokenizer_path=tokenizer_path,
pad_with_end=False,
embedding_directory=embedding_directory,
embedding_size=1280,
embedding_key="clip_g",
tokenizer_data=tokenizer_data or {},
)
class PromptLangSDXLTokenizer:
def __init__(self, embedding_directory=None, tokenizer_data=None) -> None:
td = tokenizer_data or {}
self.clip_l = PromptLangSDTokenizer(embedding_directory=embedding_directory, tokenizer_data=td)
self.clip_g = PromptLangSDXLClipGTokenizer(embedding_directory=embedding_directory, tokenizer_data=td)
def tokenize_with_weights(self, text: str, return_word_ids: bool = False, **kwargs) -> Dict[str, List[List[TokenEntry]]]:
tree = PromptParser.parse(text)
return {
"g": self.clip_g.tokenize_with_weights(text, return_word_ids, _parsed_tree=tree, **kwargs),
"l": self.clip_l.tokenize_with_weights(text, return_word_ids, _parsed_tree=tree, **kwargs),
}
def untokenize(self, token_weight_pair):
return self.clip_g.untokenize(token_weight_pair)
def state_dict(self):
return {}
+174 -104
View File
@@ -1,23 +1,24 @@
import random
import os
from dataclasses import dataclass
from typing import Any, List, Tuple
import numpy as np
from PIL import Image
import folder_paths
import comfy.sd
import comfy.ops
from custom_nodes.ClipStuff.lib.clip_model import PromptLangClipModel
import folder_paths
from comfy.supported_models_base import ClipTarget
from custom_nodes.ClipStuff.lib.tokenizer import PromptLangTokenizer
class EmptyClass:
pass
from .lib.clip_model import PromptLangSD1ClipModel, PromptLangSDXLClipModel
from .lib.inspect import inspect_prompt
from .lib.tokenizer import PromptLangSD1Tokenizer, PromptLangSDXLTokenizer
class SpecialClipLoader:
"""Wraps a loaded CLIP with our DSL-aware tokenizer + text encoder."""
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls): # type: ignore[no-untyped-def]
return {
"required": {
"source_clip": ("CLIP",),
@@ -26,131 +27,200 @@ class SpecialClipLoader:
RETURN_TYPES = ("CLIP",)
FUNCTION = "load_clip"
OUTPUT_IS_LIST = (False,)
CATEGORY = "conditioning"
@staticmethod
def load_clip(source_clip):
clip_target = EmptyClass()
clip_target.params = {}
clip_target.clip = PromptLangClipModel
clip_target.tokenizer = PromptLangTokenizer
def load_clip(source_clip: comfy.sd.CLIP) -> Tuple[comfy.sd.CLIP]:
is_sdxl = hasattr(source_clip.cond_stage_model, "clip_g") and hasattr(source_clip.cond_stage_model, "clip_l")
embedding_directory = source_clip.tokenizer.clip_l.embedding_directory
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.embedding_directory)
comfy.sd.load_clip_weights(
clip.cond_stage_model, source_clip.cond_stage_model.state_dict()
)
return (clip,)
if is_sdxl:
target = ClipTarget(PromptLangSDXLTokenizer, PromptLangSDXLClipModel)
else:
target = ClipTarget(PromptLangSD1Tokenizer, PromptLangSD1ClipModel)
new_clip = comfy.sd.CLIP(target=target, embedding_directory=embedding_directory)
new_clip.cond_stage_model.load_state_dict(source_clip.cond_stage_model.state_dict())
new_clip.layer_idx = source_clip.layer_idx
return (new_clip,)
def tensor2img(tensor_img):
i = 255.0 * tensor_img.cpu().numpy()
i_np_arr = np.clip(i, 0, 255, out=i).astype(np.uint8, copy=False)
return Image.fromarray(i_np_arr)
class PromptLangInspect:
"""Shows what a DSL prompt resolves to at the embedding layer: per-slot weight, L2 norm, nearest vocab."""
@classmethod
def INPUT_TYPES(cls): # type: ignore[no-untyped-def]
return {
"required": {
"clip": ("CLIP",),
"text": ("STRING", {"multiline": True}),
"top_k": ("INT", {"default": 3, "min": 1, "max": 10}),
}
}
RETURN_TYPES = ("STRING",)
FUNCTION = "inspect"
CATEGORY = "conditioning"
OUTPUT_NODE = True
def inspect(self, clip, text: str, top_k: int):
report = inspect_prompt(clip, text, top_k=top_k)
return {"ui": {"text": [report]}, "result": (report,)}
def tensor2img(tensor_img) -> Image.Image:
arr = (255.0 * tensor_img.cpu().numpy()).clip(0, 255).astype(np.uint8)
return Image.fromarray(arr)
class BuildGif:
def __init__(self):
pass
"""Builds an animated webp from a list of image batches.
Two output modes:
- "Big Grid": tiles batches across the X axis and chunks across the Y axis,
producing a single animated webp where each frame is the next image in a chunk.
- "One Per Split": one animation per (split, batch_index) combination.
"""
def __init__(self) -> None:
self.output_dir = folder_paths.get_output_directory()
@classmethod
def INPUT_TYPES(cls):
def INPUT_TYPES(cls): # type: ignore[no-untyped-def]
return {
"required": {
"images": ("IMAGE",),
"split_every": ("INT", {"default": -1}),
"output_mode": (
["One Per Split", "Big Grid"],
{"default": "Big Grid"},
),
"frame_duration": ("INT", {"default": 125}),
"output_mode": (["One Per Split", "Big Grid"], {"default": "Big Grid"}),
}
}
RELOAD_INST = True
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("Gifs",)
RETURN_TYPES = ()
INPUT_IS_LIST = True
FUNCTION = "build_gif"
OUTPUT_IS_LIST = (True,)
# OUTPUT_NODE = False
OUTPUT_NODE = True
CATEGORY = "List Stuff"
@staticmethod
def build_gif(images: list, split_every: list[int], output_mode: str):
print("Build GIF called!")
print(f"{type(images)}")
def build_gif(
self,
images: List[Any],
split_every: List[int],
frame_duration: List[int],
output_mode: List[str],
):
if len(split_every) > 1:
raise Exception("List input for split every is not supported.")
raise ValueError("List input for split_every is not supported.")
if len(output_mode) > 1:
raise ValueError("List input for output_mode is not supported.")
if len(frame_duration) > 1:
raise ValueError("List input for frame_duration is not supported.")
mode = output_mode[0]
duration = frame_duration[0]
split_requested = split_every[0]
full_output_folder, filename, counter, subfolder, _ = folder_paths.get_save_image_path(
filename_prefix="Gif", output_dir=self.output_dir, image_width=0, image_height=0
)
split_every = split_every[0]
batch_size = images[0].size()[0]
if split_every == -1:
# split_every=-1 means "don't split": one chunk containing everything.
if split_requested == -1:
split_chunks = 1
split_every = len(images)
chunk_len = len(images)
else:
split_chunks = int(len(images) / split_every)
out = []
num_wide = batch_size
num_tall = split_chunks
chunk_len = split_requested
split_chunks = len(images) // chunk_len
chunked_batches = [
images[split_every * chunk_idx : split_every * (chunk_idx + 1)]
for chunk_idx in range(split_chunks)
images[chunk_len * i : chunk_len * (i + 1)]
for i 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 = []
if output_mode == "Big Grid":
# For every image in gif
for idx_in_chunk in range(split_every):
img_shape = images[0][0].shape
img_frame = Image.new(
"RGB", size=(num_wide * img_shape[0], num_tall * img_shape[1])
)
# For every chunk of images
for split_idx in range(split_chunks):
img_chunk = chunked_batches[split_idx]
for batch_idx, img_tensor in enumerate(img_chunk[idx_in_chunk]):
img = tensor2img(img_tensor)
img_frame.paste(
img, (batch_idx * img_shape[0], split_idx * img_shape[1])
)
frames.append(img_frame)
save_path = (
f"{folder_paths.get_output_directory()}/{random.randint(1, 100)}"
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])
)
frames[0].save(
f"{save_path}.webp",
# quality=100,
# method=6,
lossless=True,
save_all=True,
append_images=frames[1:],
optimize=False,
duration=125,
loop=0,
)
elif output_mode == "One Per Split":
for split_idx in range(int(split_chunks)):
split_start = split_every * split_idx
split_end = split_every * (split_idx + 1)
for batch_idx in range(batch_size):
save_path = f"{folder_paths.get_output_directory()}/-{batch_idx}-{random.randint(1, 100)}"
print(save_path)
tensor2img(images[split_start][batch_idx]).save(
f"{save_path}.webp",
save_all=True,
append_images=[
tensor2img(nested_batch[batch_idx])
for nested_batch in images[split_start + 1 : split_end]
],
optimize=False,
duration=125,
loop=0,
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]),
)
return (out,)
frames.append(img_frame)
file = f"{ctx.filename}_{ctx.counter:05}_"
save_path = os.path.join(ctx.full_output_folder, file)
frames[0].save(
f"{save_path}.webp",
lossless=True,
save_all=True,
append_images=frames[1:],
optimize=False,
duration=ctx.duration,
loop=0,
)
return {"filename": f"{file}.webp", "subfolder": ctx.subfolder, "type": "output"}
def _save_one_per_split(self, ctx):
results = []
counter = ctx.counter
for split_idx in range(ctx.split_chunks):
split_start = ctx.chunk_len * split_idx
split_end = ctx.chunk_len * (split_idx + 1)
for batch_idx in range(ctx.batch_size):
file = f"{ctx.filename}_{counter:05}_"
save_path = os.path.join(ctx.full_output_folder, file)
counter += 1
tensor2img(ctx.images[split_start][batch_idx]).save(
f"{save_path}.webp",
save_all=True,
append_images=[
tensor2img(nested[batch_idx])
for nested in ctx.images[split_start + 1 : split_end]
],
optimize=False,
duration=ctx.duration,
loop=0,
)
results.append({"filename": f"{file}.webp", "subfolder": ctx.subfolder, "type": "output"})
return results
@dataclass
class _SaveContext:
images: Any
chunked_batches: Any
chunk_len: int
batch_size: int
split_chunks: int
full_output_folder: str
filename: str
counter: int
subfolder: str
duration: int
+26
View File
@@ -0,0 +1,26 @@
[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"]
+105
View File
@@ -0,0 +1,105 @@
"""Test setup that runs before any tests are collected.
Two things make this tricky:
1. The project's runtime imports use ComfyUI (`comfy.sd1_clip`), which we don't want to require for unit tests.
2. The package is normally installed under `custom_nodes/KepPromptLang/`, so we register `KepPromptLang` as a package alias.
"""
import os
import sys
import types
import pytest
REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, os.path.dirname(REPO_ROOT))
def _install_runtime_stubs():
"""Stub out runtime deps (numpy/PIL/comfy/folder_paths) so test imports of the package work.
Tests don't exercise the ComfyUI nodes; they only need the parser and action math layers.
"""
# Only stub modules that aren't actually installed; real numpy/torch must take precedence.
if "PIL" not in sys.modules:
try:
import PIL # noqa: F401
except ImportError:
pil = types.ModuleType("PIL")
pil.Image = types.ModuleType("PIL.Image")
sys.modules["PIL"] = pil
sys.modules["PIL.Image"] = pil.Image
if "numpy" not in sys.modules:
try:
import numpy # noqa: F401
except ImportError:
sys.modules["numpy"] = types.ModuleType("numpy")
if "folder_paths" not in sys.modules:
sys.modules["folder_paths"] = types.ModuleType("folder_paths")
if "comfy" in sys.modules:
return
comfy = types.ModuleType("comfy")
comfy_sd = types.ModuleType("comfy.sd")
comfy_sd.CLIP = type("CLIP", (), {})
comfy_supported = types.ModuleType("comfy.supported_models_base")
comfy_supported.ClipTarget = type("ClipTarget", (), {})
sdxl_clip = types.ModuleType("comfy.sdxl_clip")
sdxl_clip.SDXLClipModel = type("SDXLClipModel", (), {})
sdxl_clip.SDXLClipG = type("SDXLClipG", (), {})
sd1_clip = types.ModuleType("comfy.sd1_clip")
class SDTokenizer: # minimal stand-in
embedding_identifier = "embedding:"
def __init__(self, *args, **kwargs):
self.embedding_directory = None
self.embedding_size = 768
self.start_token = 49406
self.end_token = 49407
self.pad_with_end = True
self.max_length = 77
self.tokenizer = _FakeTokenizer()
def _try_get_embedding(self, name):
return None, ""
class SD1Tokenizer:
def __init__(self, *args, **kwargs):
pass
sd1_clip.SDTokenizer = SDTokenizer
sd1_clip.SD1Tokenizer = SD1Tokenizer
sd1_clip.SDClipModel = type("SDClipModel", (), {})
sd1_clip.SD1ClipModel = type("SD1ClipModel", (), {})
comfy.sd1_clip = sd1_clip
comfy.sd = comfy_sd
comfy.sdxl_clip = sdxl_clip
comfy.supported_models_base = comfy_supported
sys.modules["comfy"] = comfy
sys.modules["comfy.sd1_clip"] = sd1_clip
sys.modules["comfy.sdxl_clip"] = sdxl_clip
sys.modules["comfy.sd"] = comfy_sd
sys.modules["comfy.supported_models_base"] = comfy_supported
class _FakeTokenizer:
"""Tokenize each whitespace-separated word into a single deterministic int id."""
def __call__(self, word):
# Deterministic: sum of character codepoints, modulo a small range. SOT/EOT bracketing.
token = (sum(ord(c) for c in word) % 49000) + 100
return {"input_ids": [49406, token, 49407]}
_install_runtime_stubs()
@pytest.fixture
def tokenizer():
from comfy.sd1_clip import SDTokenizer
return SDTokenizer()
+160
View File
@@ -0,0 +1,160 @@
"""Action-level tests using a tiny in-memory torch.nn.Embedding.
These verify the math/shape contracts of each action without needing ComfyUI or a real CLIP.
"""
import pytest
torch = pytest.importorskip("torch")
from KepPromptLang.lib.actions.avg import AverageAction
from KepPromptLang.lib.actions.diff import DiffAction
from KepPromptLang.lib.actions.mult import MultiplyAction
from KepPromptLang.lib.actions.neg import NegAction
from KepPromptLang.lib.actions.norm import NormAction
from KepPromptLang.lib.actions.pos_scale import PosScaleAction
from KepPromptLang.lib.actions.post_pos import PostPosAction
from KepPromptLang.lib.actions.rand import RandAction
from KepPromptLang.lib.actions.scale_dims import ScaleDims
from KepPromptLang.lib.actions.set_dims import SetDims
from KepPromptLang.lib.actions.slerp import SlerpAction
from KepPromptLang.lib.actions.sum import SumAction
from KepPromptLang.lib.actions.utils import slerp
from KepPromptLang.lib.parser.prompt_segment import PromptSegment
EMBED_DIM = 4
VOCAB = 100
@pytest.fixture
def embedding():
torch.manual_seed(0)
emb = torch.nn.Embedding(VOCAB, EMBED_DIM)
return emb
def seg(*token_ids):
return PromptSegment(text="x", tokens=list(token_ids))
def test_sum_adds_embeddings(embedding):
a = seg(1, 2)
b = seg(3, 4)
action = SumAction([[a], [b]])
expected = embedding(torch.LongTensor([[1, 2]])) + embedding(torch.LongTensor([[3, 4]]))
assert torch.allclose(action.get_result(embedding), expected)
def test_diff_subtracts_embeddings(embedding):
a = seg(1, 2)
b = seg(3, 4)
action = DiffAction([[a], [b]])
expected = embedding(torch.LongTensor([[1, 2]])) - embedding(torch.LongTensor([[3, 4]]))
assert torch.allclose(action.get_result(embedding), expected)
def test_neg_negates(embedding):
action = NegAction([seg(1, 2)])
expected = -embedding(torch.LongTensor([[1, 2]]))
assert torch.allclose(action.get_result(embedding), expected)
def test_mult_scales(embedding):
# The multiplier is read from PromptSegment.text (mimicking parser output).
target_seg = PromptSegment(text="3.5", tokens=[1])
action = MultiplyAction([[seg(1, 2)], [target_seg]])
expected = embedding(torch.LongTensor([[1, 2]])) * 3.5
assert torch.allclose(action.get_result(embedding), expected)
def test_norm_unit_length(embedding):
action = NormAction([seg(1, 2)])
result = action.get_result(embedding)
norms = torch.norm(result, dim=-1)
assert torch.allclose(norms, torch.ones_like(norms), atol=1e-5)
def test_avg_weighted_mix(embedding):
weight_seg = PromptSegment(text="0.25", tokens=[1])
action = AverageAction([[seg(1, 2)], [seg(3, 4)], [weight_seg]])
expected = embedding(torch.LongTensor([[1, 2]])) * 0.75 + embedding(torch.LongTensor([[3, 4]])) * 0.25
assert torch.allclose(action.get_result(embedding), expected)
def test_avg_mismatched_lengths_errors():
weight_seg = PromptSegment(text="0.5", tokens=[1])
with pytest.raises(ValueError, match="same length"):
AverageAction([[seg(1, 2)], [seg(3)], [weight_seg]])
def test_slerp_endpoints(embedding):
weight0 = PromptSegment(text="0.0", tokens=[1])
weight1 = PromptSegment(text="1.0", tokens=[1])
a, b = seg(1, 2), seg(3, 4)
a_emb = embedding(torch.LongTensor([[1, 2]]))
b_emb = embedding(torch.LongTensor([[3, 4]]))
assert torch.allclose(SlerpAction([[a], [b], [weight0]]).get_result(embedding), a_emb, atol=1e-5)
assert torch.allclose(SlerpAction([[a], [b], [weight1]]).get_result(embedding), b_emb, atol=1e-5)
def test_slerp_helper_endpoints_and_midpoint():
low = torch.tensor([1.0, 0.0])
high = torch.tensor([0.0, 1.0]) # 90 degrees apart, both unit length
assert torch.allclose(slerp(0.0, low, high), low, atol=1e-5)
assert torch.allclose(slerp(1.0, low, high), high, atol=1e-5)
midpoint = slerp(0.5, low, high)
# Midpoint of orthogonal unit vectors on the unit sphere is (sqrt(2)/2, sqrt(2)/2).
expected = torch.tensor([2 ** 0.5 / 2, 2 ** 0.5 / 2])
assert torch.allclose(midpoint, expected, atol=1e-5)
def test_rand_token_length_and_bounds():
length_seg = PromptSegment(text="3", tokens=[1])
min_seg = PromptSegment(text="-2", tokens=[1])
max_seg = PromptSegment(text="2", tokens=[1])
action = RandAction([[length_seg], [min_seg], [max_seg]])
emb = torch.nn.Embedding(VOCAB, EMBED_DIM)
result = action.get_result(emb)
assert result.shape == (1, 3, EMBED_DIM)
assert (result >= -2).all() and (result <= 2).all()
def test_scale_dims_modifies_only_target_dim(embedding):
pair_seg = PromptSegment(text="0,3.0", tokens=[1])
action = ScaleDims([[seg(1, 2)], [pair_seg]])
base = embedding(torch.LongTensor([[1, 2]])).clone()
result = action.get_result(embedding)
assert torch.allclose(result[0, :, 0], base[0, :, 0] * 3.0)
assert torch.allclose(result[0, :, 1:], base[0, :, 1:])
def test_set_dims_overwrites_value(embedding):
pair_seg = PromptSegment(text="2,-9.5", tokens=[1])
action = SetDims([[seg(1, 2)], [pair_seg]])
result = action.get_result(embedding)
assert torch.allclose(result[0, :, 2], torch.tensor([-9.5, -9.5]))
def test_pos_scale_returns_modifier(embedding):
multiplier = PromptSegment(text="1.5", tokens=[1])
action = PosScaleAction([[seg(1, 2)], [multiplier]])
tensor, modifiers = action.get_result(embedding)
assert tensor.shape == (1, 2, EMBED_DIM)
assert modifiers.position_embed_scale == 1.5
def test_post_pos_returns_bypass(embedding):
action = PostPosAction([seg(1, 2)])
tensor, modifiers = action.get_result(embedding)
assert tensor.shape == (1, 2, EMBED_DIM)
assert modifiers.bypass_pos_embed is True
def test_action_token_lengths():
a, b = seg(1, 2, 3), seg(4, 5, 6)
assert SumAction([[a], [b]]).token_length() == 3
assert DiffAction([[a], [b]]).token_length() == 3
assert NegAction([a]).token_length() == 3
assert NormAction([a]).token_length() == 3
+96
View File
@@ -0,0 +1,96 @@
import pytest
torch = pytest.importorskip("torch")
from KepPromptLang.lib.actions.nearest import NearestAction
from KepPromptLang.lib.actions.noise import NoiseAction
from KepPromptLang.lib.actions.project import ProjectAction, RejectAction
from KepPromptLang.lib.actions.renorm import RenormAction
from KepPromptLang.lib.parser.prompt_segment import PromptSegment
from KepPromptLang.lib.parser.registration import get_action_by_name
EMBED_DIM = 4
VOCAB = 50
@pytest.fixture
def embedding():
torch.manual_seed(0)
return torch.nn.Embedding(VOCAB, EMBED_DIM)
def seg(*token_ids):
return PromptSegment(text="x", tokens=list(token_ids))
def test_proj_plus_reject_reconstructs_input(embedding):
a, b = seg(1, 2), seg(3)
proj = ProjectAction([[a], [b]]).get_result(embedding)
rej = RejectAction([[a], [b]]).get_result(embedding)
a_emb = embedding(torch.LongTensor([[1, 2]]))
assert torch.allclose(proj + rej, a_emb, atol=1e-5)
def test_reject_is_orthogonal_to_b(embedding):
a, b = seg(1, 2), seg(3)
rej = RejectAction([[a], [b]]).get_result(embedding)
b_dir = torch.nn.functional.normalize(
embedding(torch.LongTensor([[3]])).mean(dim=1, keepdim=True), dim=-1
)
dots = (rej * b_dir).sum(dim=-1)
assert torch.allclose(dots, torch.zeros_like(dots), atol=1e-5)
def test_renorm_matches_ref_norm(embedding):
a, ref = seg(1, 2), seg(3)
out = RenormAction([[a], [ref]]).get_result(embedding)
ref_norm = torch.norm(embedding(torch.LongTensor([[3]])), dim=-1).mean()
out_norms = torch.norm(out, dim=-1)
assert torch.allclose(out_norms, ref_norm.expand_as(out_norms), atol=1e-5)
def test_noise_shape_and_mean(embedding):
std = PromptSegment(text="0.01", tokens=[1])
out = NoiseAction([[seg(1, 2)], [std]]).get_result(embedding)
base = embedding(torch.LongTensor([[1, 2]]))
assert out.shape == base.shape
# Perturbation magnitude bounded (5σ with margin); std=0.01, EMBED_DIM=4.
assert (out - base).abs().max() < 0.2
def test_noise_zero_std_is_identity(embedding):
std = PromptSegment(text="0.0", tokens=[1])
out = NoiseAction([[seg(1, 2)], [std]]).get_result(embedding)
base = embedding(torch.LongTensor([[1, 2]]))
assert torch.allclose(out, base)
def test_nearest_returns_exact_token_for_that_token(embedding):
out = NearestAction([[seg(7)]]).get_result(embedding)
assert out.shape == (1, 1, EMBED_DIM)
assert torch.allclose(out[0, 0], embedding.weight[7])
def test_nearest_k_tokens(embedding):
k = PromptSegment(text="3", tokens=[1])
action = NearestAction([[seg(7)], [k]])
assert action.token_length() == 3
out = action.get_result(embedding)
assert out.shape == (1, 3, EMBED_DIM)
# First match should be the token itself.
assert torch.allclose(out[0, 0], embedding.weight[7])
def test_lerp_is_registered_as_avg_alias():
lerp_cls = get_action_by_name("lerp")
avg_cls = get_action_by_name("avg")
assert issubclass(lerp_cls, avg_cls)
def test_token_lengths():
a, b = seg(1, 2, 3), seg(4)
assert ProjectAction([[a], [b]]).token_length() == 3
assert RejectAction([[a], [b]]).token_length() == 3
assert RenormAction([[a], [b]]).token_length() == 3
std = PromptSegment(text="0.1", tokens=[1])
assert NoiseAction([[a], [std]]).token_length() == 3
+58
View File
@@ -0,0 +1,58 @@
from KepPromptLang.lib.actions.diff import DiffAction
from KepPromptLang.lib.actions.norm import NormAction
from KepPromptLang.lib.actions.sum import SumAction
from KepPromptLang.lib.parser import PromptParser
from KepPromptLang.lib.parser.prompt_segment import PromptSegment
from KepPromptLang.lib.parser.transformer import PromptTransformer
def parse(text, tokenizer):
tree = PromptParser.parse(text)
return PromptTransformer(tokenizer).transform(tree)
def test_plain_words_become_segments(tokenizer):
result = parse("hello world", tokenizer)
items = result.children
assert len(items) == 2
assert all(isinstance(i, PromptSegment) for i in items)
assert items[0].text == "hello"
assert items[1].text == "world"
def test_sum_action_parses(tokenizer):
action = parse("sum(king|man|woman)", tokenizer)
items = action.children if hasattr(action, "children") else [action]
assert len(items) == 1
assert isinstance(items[0], SumAction)
assert len(items[0].all_args) == 3
def test_nested_actions(tokenizer):
action = parse("sum(diff(king|man)|woman)", tokenizer)
items = action.children if hasattr(action, "children") else [action]
outer = items[0]
assert isinstance(outer, SumAction)
inner = outer.all_args[0][0]
assert isinstance(inner, DiffAction)
def test_norm_single_arg(tokenizer):
action = parse("norm(cat)", tokenizer)
items = action.children if hasattr(action, "children") else [action]
assert isinstance(items[0], NormAction)
def test_quoted_string(tokenizer):
result = parse('"hello world"', tokenizer)
items = result.children if hasattr(result, "children") else [result]
assert isinstance(items[0], PromptSegment)
assert items[0].text == "hello world"
def test_unknown_action_errors(tokenizer):
import pytest
from lark.exceptions import VisitError
with pytest.raises((ValueError, VisitError), match="not found in registry"):
parse("nonexistentAction(cat)", tokenizer)
+93
View File
@@ -0,0 +1,93 @@
"""Verify the tokenizer emits ComfyUI's native (token, weight) format with lazy Actions
and per-position weights.
Uses the comfy stub from conftest, so no real ComfyUI needed.
"""
import pytest
torch = pytest.importorskip("torch")
from KepPromptLang.lib.actions.base import ACTION_CONTINUATION, Action
from KepPromptLang.lib.actions.sum import SumAction
from KepPromptLang.lib.actions.weighted import WeightedGroup
from KepPromptLang.lib.tokenizer import PromptLangSDTokenizer
@pytest.fixture
def tok():
return PromptLangSDTokenizer()
def test_plain_text_is_int_tuples_at_max_length(tok):
[row] = tok.tokenize_with_weights("hello world")
assert all(isinstance(t, int) and w == 1.0 for t, w in row)
assert row[0] == (tok.start_token, 1.0)
assert len(row) == tok.max_length
def test_action_emits_one_entry_plus_continuations(tok):
[row] = tok.tokenize_with_weights("a sum(king|man|woman) here")
assert len(row) == tok.max_length
actions = [t for t, _ in row if isinstance(t, Action)]
continuations = [t for t, _ in row if t is ACTION_CONTINUATION]
assert len(actions) == 1
assert isinstance(actions[0], SumAction)
assert len(continuations) == actions[0].token_length() - 1
def test_paren_weight_syntax(tok):
[row] = tok.tokenize_with_weights("a (cat:1.3) here")
weighted = [(t, w) for t, w in row if w != 1.0]
# "cat" is one token under the fake tokenizer.
assert len(weighted) == 1
assert weighted[0][1] == pytest.approx(1.3)
assert isinstance(weighted[0][0], int)
def test_paren_weight_on_action_propagates_to_continuations(tok):
[row] = tok.tokenize_with_weights("(sum(king|man|woman):0.7)")
action_entry = next((t, w) for t, w in row if isinstance(t, Action))
cont_weights = [w for t, w in row if t is ACTION_CONTINUATION]
assert action_entry[1] == pytest.approx(0.7)
assert all(w == pytest.approx(0.7) for w in cont_weights)
def test_nested_paren_weights_multiply(tok):
[row] = tok.tokenize_with_weights("((cat:1.2):0.5)")
weighted = [(t, w) for t, w in row if w != 1.0]
assert len(weighted) == 1
assert weighted[0][1] == pytest.approx(0.6)
def test_emph_is_alias_for_paren_weight(tok):
[row] = tok.tokenize_with_weights("emph(cat|1.3)")
weighted = [(t, w) for t, w in row if w != 1.0]
assert len(weighted) == 1
assert weighted[0][1] == pytest.approx(1.3)
def test_nested_actions_stay_nested(tok):
[row] = tok.tokenize_with_weights("sum(diff(king|man)|woman)")
actions = [t for t, _ in row if isinstance(t, Action)]
assert len(actions) == 1
assert isinstance(actions[0], SumAction)
from KepPromptLang.lib.actions.diff import DiffAction
assert isinstance(actions[0].all_args[0][0], DiffAction)
def test_overflow_splits_into_multiple_batches(tok):
text = " ".join(f"w{i}" for i in range(80))
batches = tok.tokenize_with_weights(text)
assert len(batches) >= 2
for row in batches:
assert len(row) == tok.max_length
assert row[0] == (tok.start_token, 1.0)
def test_weighted_group_token_length():
from KepPromptLang.lib.parser.prompt_segment import PromptSegment
grp = WeightedGroup([PromptSegment("a", [1, 2]), PromptSegment("b", [3])], 1.5)
assert grp.token_length() == 3
+77
View File
@@ -0,0 +1,77 @@
import pytest
torch = pytest.importorskip("torch")
from KepPromptLang.lib.actions.base import Action
from KepPromptLang.lib.actions.sum import SumAction
from KepPromptLang.lib.tokenizer import PromptLangSDTokenizer
@pytest.fixture
def tok():
return PromptLangSDTokenizer()
def content_tokens(row, tok):
"""Non-SOT/EOT/pad int tokens from a row, in order."""
return [
t for t, _ in row
if isinstance(t, int) and t not in (tok.start_token, tok.end_token, 0)
]
def test_var_substitutes_at_top_level(tok):
[direct] = tok.tokenize_with_weights("a cat dog")
[via_var] = tok.tokenize_with_weights("$x = cat dog; a $x")
assert content_tokens(via_var, tok) == content_tokens(direct, tok)
def test_var_holding_action(tok):
[row] = tok.tokenize_with_weights("$axis = sum(king|man); $axis")
actions = [t for t, _ in row if isinstance(t, Action)]
assert len(actions) == 1
assert isinstance(actions[0], SumAction)
def test_var_inside_function_arg(tok):
[row] = tok.tokenize_with_weights("$a = king; sum($a|woman)")
actions = [t for t, _ in row if isinstance(t, Action)]
assert len(actions) == 1
# token_length should be 1 (single-token base arg via the fake tokenizer)
assert actions[0].token_length() == 1
def test_var_under_weight(tok):
[row] = tok.tokenize_with_weights("$x = cat; ($x:1.5)")
weighted = [w for t, w in row if isinstance(t, int) and w != 1.0]
assert weighted == [pytest.approx(1.5)]
def test_var_ref_before_assign_errors(tok):
from lark.exceptions import VisitError
with pytest.raises((ValueError, VisitError), match="referenced before assignment"):
tok.tokenize_with_weights("$x and then $x = cat;")
def test_var_reassignment_errors(tok):
from lark.exceptions import VisitError
with pytest.raises((ValueError, VisitError), match="already defined"):
tok.tokenize_with_weights("$x = cat; $x = dog; $x")
def test_var_chains(tok):
[direct] = tok.tokenize_with_weights("cat")
[chained] = tok.tokenize_with_weights("$a = cat; $b = $a; $b")
assert content_tokens(chained, tok) == content_tokens(direct, tok)
def test_comments_ignored(tok):
[a] = tok.tokenize_with_weights("cat dog")
[b] = tok.tokenize_with_weights("cat # this is ignored\ndog")
assert content_tokens(a, tok) == content_tokens(b, tok)
def test_assign_only_produces_empty_prompt(tok):
[row] = tok.tokenize_with_weights("$x = cat;")
# SOT + EOT + padding only
assert content_tokens(row, tok) == []
+72
View File
@@ -0,0 +1,72 @@
"""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()))