73 Commits
Author SHA1 Message Date
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
37 changed files with 1799 additions and 252 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
+23 -25
View File
@@ -29,33 +29,10 @@ See example workflow in examples folder.
- Example: `"Hello World"`, `'It\'s a sunny day'`
- Represents string literals.
## Functions
Here are the available functions and their usage:
1. **Sum Function**:
- Syntax: `sum(arg1 | arg2 | ... | argN)`
- Adds together multiple embeddings.
- Example: `sum(embedding:face1 | dog)`
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.
- An argument (`arg`) can be an embedding, multiple words, another function, or a quoted string.
- For functions that accept multiple arguments, they are separated by the `|` symbol.
## Examples
@@ -78,4 +55,25 @@ Here are the available functions and their usage:
```
sum(king|neg(man)|woman)
```
```
## Functions
Here are the available functions and their usage:
| Display Name | Action Name | Description | Usage Examples |
| --- | --- | --- | --- |
| 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> |
| 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> |
| Negate | neg | Negates the provided segments or actions. | <ul><li>neg(cat)</li><li>sum(king\|neg(man)\|women)</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> |
| 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> |
| 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> |
| 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> |
| Pooled Average(Experimental) | _exp-pooledAvg | Processes the provided segments or actions fully through CLIP and creates a pooled average of the last hidden state by averaging the last hidden state of each token. | <ul><li>A cat on a _exp-pooledAvg(beautiful sunny day)</li><li>A _exp-pooledAvg(broken glass) bottle</li></ul> |
| Sum | sum | Adds the embeddings of the provided segments or actions. | <ul><li>A happy sum(cat\|dog\|shark)</li></ul> |
| Pooler Output(Experimental) | _exp-pooler | Processes the provided segments or actions fully through CLIP and returns the pooler_output from the transformer | <ul><li>A cat on a _exp-pooler(beautiful sunny day)</li><li>A _exp-pooler(broken glass) bottle</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> |
| 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> |
| 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> |
+31
View File
@@ -0,0 +1,31 @@
from custom_nodes.KepPromptLang.lib.actions.avg import AverageAction
from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction
from custom_nodes.KepPromptLang.lib.actions.mult import MultiplyAction
from custom_nodes.KepPromptLang.lib.actions.neg import NegAction
from custom_nodes.KepPromptLang.lib.actions.norm import NormAction
from custom_nodes.KepPromptLang.lib.actions.pooled_avg import PooledAvgAction
from custom_nodes.KepPromptLang.lib.actions.pooler import PoolerAction
from custom_nodes.KepPromptLang.lib.actions.pos_scale import PosScaleAction
from custom_nodes.KepPromptLang.lib.actions.post_pos import PostPosAction
from custom_nodes.KepPromptLang.lib.actions.rand import RandAction
from custom_nodes.KepPromptLang.lib.actions.scale_dims import ScaleDims
from custom_nodes.KepPromptLang.lib.actions.set_dims import SetDims
from custom_nodes.KepPromptLang.lib.actions.slerp import SlerpAction
from custom_nodes.KepPromptLang.lib.actions.sum import SumAction
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
register_action(DiffAction)
register_action(MultiplyAction)
register_action(NegAction)
register_action(NormAction)
register_action(RandAction)
register_action(SumAction)
register_action(SlerpAction)
register_action(AverageAction)
register_action(ScaleDims)
register_action(SetDims)
register_action(PosScaleAction)
register_action(PoolerAction)
register_action(PostPosAction)
register_action(PooledAvgAction)
+79 -21
View File
@@ -1,23 +1,61 @@
from abc import ABC, abstractmethod
from typing import Union
from enum import Enum
from typing import Union, List, TypedDict, Tuple
from torch import Tensor
from torch.nn import Embedding
from transformers.models.clip.modeling_clip import CLIPTextTransformer
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class ActionArity(Enum):
NONE = 0
SINGLE = 1
MULTI = 2
class PostModifiers(TypedDict):
"""
A dictionary of post modifiers for an action result.
"""
position_embed_scale: Union[float, None]
bypass_pos_embed: Union[bool, None]
class Action(ABC):
@property
@abstractmethod
def chars(self) -> list[str] | None:
def chars(self) -> Union[List[str], None]:
pass
@property
@abstractmethod
def name(self) -> str:
def arity(self) -> ActionArity:
"""
Determines the arity of the action. This is used to determine how many arguments the action supports.
:return:
"""
pass
@property
@abstractmethod
def display_name(self) -> str:
pass
@property
@abstractmethod
def action_name(self) -> str:
pass
@property
@abstractmethod
def description(self) -> str:
pass
@property
def usage_examples(self) -> List[str]:
return []
@property
@abstractmethod
def grammar(self) -> str:
@@ -36,7 +74,7 @@ class Action(ABC):
pass
@abstractmethod
def get_all_segments(self) -> list[PromptSegment]:
def get_all_segments(self) -> List[PromptSegment]:
"""
Get all segments, including nested segments.
:return:
@@ -44,7 +82,15 @@ class Action(ABC):
pass
@abstractmethod
def get_result(self, embedding_module: Embedding) -> Tensor:
def __init__(self, *args, **kwargs) -> None:
"""
Initialize the action. This is called when the action is parsed from the prompt.
:param args: The arguments for the action.
"""
pass
@abstractmethod
def get_result(self, embedding_module: Embedding) -> Union[Tensor, Tuple[Tensor, PostModifiers]]:
"""
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.
@@ -52,12 +98,20 @@ class Action(ABC):
"""
pass
def process_with_transformer(self, transformer: CLIPTextTransformer, embedding_module: Embedding) -> None:
"""
For actions that need access to the TextTransformer, this method is called. Results are expected to be returned via get_result still.
:param transformer: An instance of CLIPTextTransformer
"""
pass
def depth_repr(self, depth: int = 1) -> str:
raise NotImplementedError()
class SingleArgAction(Action, ABC):
def get_all_segments(self) -> list[PromptSegment]:
arity = ActionArity.SINGLE
def get_all_segments(self) -> List[PromptSegment]:
segments = []
for seg_or_action in self.arg:
if isinstance(seg_or_action, Action):
@@ -66,23 +120,23 @@ class SingleArgAction(Action, ABC):
segments.append(seg_or_action)
return segments
def __init__(self, arg: list[PromptSegment | Action]):
def process_with_transformer(self, transformer: CLIPTextTransformer, embedding_module: Embedding) -> None:
for seg_or_action in self.arg:
if isinstance(seg_or_action, Action):
seg_or_action.process_with_transformer(transformer, embedding_module)
def __init__(self, arg: List[Union[PromptSegment, Action]]):
# TODO: Target is a list now... what does this mean for us..
self.arg = arg
def __repr__(self) -> str:
return f"{self.name}({self.arg})"
return f"{self.display_name}({self.arg})"
class MultiArgAction(Action, ABC):
def get_all_segments(self) -> list[PromptSegment]:
arity = ActionArity.MULTI
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 arg in self.all_args:
for seg_or_action in arg:
if isinstance(seg_or_action, Action):
segments.extend(seg_or_action.get_all_segments())
@@ -91,10 +145,14 @@ class MultiArgAction(Action, ABC):
return segments
def process_with_transformer(self, transformer: CLIPTextTransformer, embedding_module: Embedding) -> None:
for arg in self.all_args:
for seg_or_action in arg:
if isinstance(seg_or_action, Action):
seg_or_action.process_with_transformer(transformer, embedding_module)
def __init__(
self,
base_segment: list[PromptSegment | Action],
args: list[list[Union[PromptSegment, Action]]],
args: List[List[Union[PromptSegment, Action]]],
):
self.base_segment = base_segment
self.args = args
self.all_args = args
+11 -2
View File
@@ -1,11 +1,20 @@
from typing import List
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 custom_nodes.KepPromptLang.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
def get_embedding(seg_or_action: SegOrAction, embedding_module: Embedding) -> Tensor:
if isinstance(seg_or_action, Action):
return seg_or_action.get_result(embedding_module)
return seg_or_action.get_embeddings(embedding_module)
def get_total_length(args: List[SegOrAction]) -> int:
total_length = 0
for seg_or_action in args:
total_length += seg_or_action.token_length()
return total_length
+108
View File
@@ -0,0 +1,108 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import MultiArgAction, Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.actions.utils import slerp
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class AverageAction(MultiArgAction):
grammar = 'avg(" arg "|" arg "|" arg ")"'
chars = ["+", "+"]
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[Union[PromptSegment, Action]]]) -> None:
super().__init__(args)
if len(args) != 3:
raise ValueError("Average action should have exactly three arguments(2 vectors and a weight)")
self.first_arg = args[0]
self.second_arg = args[1]
self._parse_weight(args[2])
self._validate_args()
def _parse_weight(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Average weight should have exactly one segment")
weight_seg_or_action = arg[0]
if isinstance(weight_seg_or_action, Action):
raise ValueError("Average weight should not have an action as an argument")
try:
self.parsed_weight = float(weight_seg_or_action.text)
except ValueError:
raise ValueError("Average should have an integer/float as the weight")
def _validate_args(self) -> None:
first_arg_token_length = sum(seg_or_action.token_length() for seg_or_action in self.first_arg)
second_arg_token_length = sum(seg_or_action.token_length() for seg_or_action in self.second_arg)
if first_arg_token_length != second_arg_token_length:
raise ValueError(f"Average start and end arguments should have the same length. Got {start_arg_token_length} and {end_arg_token_length}")
if self.parsed_weight < 0 or self.parsed_weight > 1:
print(f"WARNING: Average weight should be between 0 and 1. Got {self.parsed_weight}")
def token_length(self) -> int:
# Average interpolates between the embeddings of the start and end segments, so the length is the length of the start segment
return sum(seg_or_action.token_length() for seg_or_action in self.first_arg)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the start segment
all_start_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.first_arg
]
start_embedding = torch.cat(all_start_embeddings, dim=1)
# Calculate the embeddings for the end segment
all_end_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.second_arg
]
end_embedding = torch.cat(all_end_embeddings, dim=1)
# Perform the weighted average
result = start_embedding * (1 - self.parsed_weight) + end_embedding * self.parsed_weight
return result
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.additional_args))})"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_arg, Action):
base_segment_repr = self.base_arg.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.additional_args, Action):
target_repr = self.additional_args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
+33 -14
View File
@@ -1,29 +1,47 @@
from typing import Union, List
import torch
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 custom_nodes.KepPromptLang.lib.action.base import Action, MultiArgAction
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class DiffAction(MultiArgAction):
grammar = 'diff(" arg ("|" arg)* ")"'
name = "diff"
chars = ["-", "-"]
display_name = "Difference"
action_name = "diff"
description = "Subtracts the segments in the order they are given. The first segment is subtracted from the second, then the third from the result, and so on."
usage_examples = [
"diff(The cat is|The dog is)",
"diff(Cat|Dog)",
"sum(diff(king|man)|woman)",
]
def __init__(self, args: List[List[Union[PromptSegment, Action]]]):
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)
return sum(seg_or_action.token_length() for seg_or_action in self.base_arg)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the base segment
all_base_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.base_segment
for seg_or_action in self.base_arg
]
result = torch.cat(all_base_embeddings, dim=1)
for arg in self.args:
for arg in self.additional_args:
all_arg_embeddings = [
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
]
@@ -44,23 +62,24 @@ class DiffAction(MultiArgAction):
return result
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
# return f"sum(\n\tbase_segment={self.base_segment},\n\tadditional_args={self.additional_args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.args))})"
return f"sum({', '.join(map(str, self.additional_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)
if isinstance(self.base_arg, Action):
base_segment_repr = self.base_arg.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f"base_segment={self.base_segment.depth_repr()},\n"
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.args, Action):
target_repr = self.args.depth_repr(depth + 1)
if isinstance(self.additional_args, Action):
target_repr = self.additional_args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.args.depth_repr()},\n"
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
+69
View File
@@ -0,0 +1,69 @@
from typing import List
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import (
Action,
MultiArgAction,
)
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class MultiplyAction(MultiArgAction):
grammar = 'mult(" arg+ ")"'
chars = ["[", "]"]
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 should have exactly two arguments")
self.target_arg = args[0]
self._parse_multiplier(args[1])
def _parse_multiplier(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Multiply actions multiplier should have exactly one segment")
multiplier_seg_or_action = arg[0]
if isinstance(multiplier_seg_or_action, Action):
raise ValueError("Multiply actions multiplier must be a number")
try:
self.parsed_multiplier = float(multiplier_seg_or_action.text)
except ValueError:
raise ValueError("Multiply action should have an integer/float as the multiplier")
def token_length(self) -> int:
"""
Mult multiplies the embeddings of the base segment, so the length is the length of the base segment
:return:
"""
total_length = 0
for seg_or_action in self.target_arg:
total_length += seg_or_action.token_length()
return total_length
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
all_embeddings = []
for seg_or_action in self.target_arg:
if isinstance(seg_or_action, Action):
all_embeddings.append(seg_or_action.get_result(embedding_module))
else:
all_embeddings.append(seg_or_action.get_embeddings(embedding_module))
target_embeddings = torch.cat(all_embeddings, dim=1)
return target_embeddings * self.parsed_multiplier
+11 -2
View File
@@ -1,14 +1,22 @@
import torch
from torch.nn import Embedding
from custom_nodes.ClipStuff.lib.action.base import Action, SingleArgAction
from custom_nodes.KepPromptLang.lib.action.base import Action, SingleArgAction
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class NegAction(SingleArgAction):
grammar = 'neg(" arg+ ")"'
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
@@ -30,3 +38,4 @@ class NegAction(SingleArgAction):
target_embeddings = torch.cat(all_embeddings, dim=1)
return target_embeddings * -1
+10 -2
View File
@@ -1,17 +1,25 @@
import torch
from torch.nn import Embedding
from custom_nodes.ClipStuff.lib.action.base import (
from custom_nodes.KepPromptLang.lib.action.base import (
Action,
SingleArgAction,
)
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class NormAction(SingleArgAction):
grammar = 'norm(" arg+ ")"'
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
+68
View File
@@ -0,0 +1,68 @@
import torch
from torch.nn import Embedding
from transformers.modeling_outputs import BaseModelOutputWithPooling
from transformers.models.clip.modeling_clip import CLIPTextTransformer
from custom_nodes.KepPromptLang.lib.action.base import Action, SingleArgAction
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_total_length
from custom_nodes.KepPromptLang.lib.fun_clip_stuff import (
PromptLangCLIPTextEmbeddings,
PrompLangCLIPTextTransformer,
)
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class PooledAvgAction(SingleArgAction):
grammar = 'pooledAvg(" arg+ ")"'
chars = ["[", "]"]
display_name = "Pooled Average(Experimental)"
action_name = "_exp-pooledAvg"
description = "Processes the provided segments or actions fully through CLIP and creates a pooled average of the last hidden state by averaging the last hidden state of each token."
usage_examples = [
"A cat on a _exp-pooledAvg(beautiful sunny day)",
"A _exp-pooledAvg(broken glass) bottle",
]
def __init__(self, args):
super().__init__(args)
self.result = None
def token_length(self) -> int:
"""
PooledAvg returns the average of the last hidden state, so the length is 1
"""
return 1
def process_with_transformer(
self, transformer: CLIPTextTransformer, embedding_module: Embedding
) -> None:
""" """
# SOT + tokens + EOT
eot_token = embedding_module.num_embeddings - 1
print("Using EOT token", eot_token)
#TODO: Play with impact of padding on pooled output
arg_length = get_total_length(self.arg)
# SOT + arg length + EOT
empty_tokens = [[49406] + [eot_token] * (arg_length + 1)]
transformer_results: BaseModelOutputWithPooling = transformer(
[
[PromptSegment(text="_Empty Batch_", tokens=empty_tokens[0])],
[PromptSegment(text="[SOT]", tokens=[49406])]
+ self.arg
+ [PromptSegment(text="[EOT]", tokens=[eot_token])],
]
)
self.result = (
transformer_results.last_hidden_state[1, 1:-1, :].mean(dim=0).unsqueeze(0).unsqueeze(0)
)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
if self.result is not None:
return self.result
raise Exception(
"PooledAvg action result is not set. Did you forget to call process_with_transformer?"
)
+67
View File
@@ -0,0 +1,67 @@
import torch
from torch.nn import Embedding
from transformers.modeling_outputs import BaseModelOutputWithPooling
from transformers.models.clip.modeling_clip import CLIPTextTransformer
from custom_nodes.KepPromptLang.lib.action.base import Action, SingleArgAction
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_total_length
from custom_nodes.KepPromptLang.lib.fun_clip_stuff import (
PromptLangCLIPTextEmbeddings,
PrompLangCLIPTextTransformer,
)
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class PoolerAction(SingleArgAction):
grammar = 'pooler(" arg+ ")"'
chars = ["[", "]"]
display_name = "Pooler Output(Experimental)"
action_name = "_exp-pooler"
description = "Processes the provided segments or actions fully through CLIP and returns the pooler_output from the transformer"
usage_examples = [
"A cat on a _exp-pooler(beautiful sunny day)",
"A _exp-pooler(broken glass) bottle",
]
def __init__(self, args):
super().__init__(args)
self.result = None
def token_length(self) -> int:
"""
Pooler returns the embedding of the EOT token, so the length is 1
"""
return 1
def process_with_transformer(
self, transformer: CLIPTextTransformer, embedding_module: Embedding
) -> None:
""" """
# SOT + tokens + EOT
eot_token = embedding_module.num_embeddings - 1
print("Using EOT token", eot_token)
arg_length = get_total_length(self.arg)
# SOT + arg length + EOT
empty_tokens = [[49406] + [eot_token] * (arg_length + 1)]
transformer_results: BaseModelOutputWithPooling = transformer(
[
[PromptSegment(text="_Empty Batch_", tokens=empty_tokens[0])],
[PromptSegment(text="[SOT]", tokens=[49406])]
+ self.arg
+ [PromptSegment(text="[EOT]", tokens=[eot_token])],
]
)
self.result = transformer_results.pooler_output[1].unsqueeze(0).unsqueeze(0)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
if self.result is not None:
return self.result
raise Exception(
"Pooled action result is not set. Did you forget to call process_with_transformer?"
)
+71
View File
@@ -0,0 +1,71 @@
from typing import Tuple, List
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import (
Action,
PostModifiers,
MultiArgAction,
)
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
class PosScaleAction(MultiArgAction):
grammar = 'posScale(" arg+ ")"'
chars = ["[", "]"]
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 should have exactly two arguments")
self.target_arg = args[0]
self._parse_multiplier(args[1])
def _parse_multiplier(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError(
"PosScale actions multiplier should have exactly one segment"
)
multiplier_seg_or_action = arg[0]
if isinstance(multiplier_seg_or_action, Action):
raise ValueError("PosScale actions multiplier must be a number")
try:
self.parsed_multiplier = float(multiplier_seg_or_action.text)
except ValueError:
raise ValueError(
"PosScale action should have an integer/float as the multiplier"
)
def token_length(self) -> int:
"""
PosScale modifies the posional embeddings of the base segment, so the length is the length of the base segment
:return:
"""
total_length = 0
for seg_or_action in self.target_arg:
total_length += seg_or_action.token_length()
return total_length
def get_result(self, embedding_module: Embedding) -> Tuple[torch.Tensor, PostModifiers]:
all_embeddings = []
for seg_or_action in self.target_arg:
if isinstance(seg_or_action, Action):
all_embeddings.append(seg_or_action.get_result(embedding_module))
else:
all_embeddings.append(seg_or_action.get_embeddings(embedding_module))
target_embeddings = torch.cat(all_embeddings, dim=1)
return target_embeddings, {"position_embed_scale": self.parsed_multiplier}
+47
View File
@@ -0,0 +1,47 @@
from typing import Tuple
import torch
from torch import Tensor
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import (
SingleArgAction,
Action,
PostModifiers,
)
class PostPosAction(SingleArgAction):
grammar = 'postPos(" arg+ ")"'
chars = ["[", "]"]
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 __init__(self, args):
super().__init__(args)
self.result = None
def token_length(self) -> int:
"""
PostPos returns the results of the wrapped action, so the length is the length of the wrapped action
"""
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) -> Tuple[Tensor, PostModifiers]:
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))
return torch.cat(all_embeddings, dim=1), PostModifiers(position_embed_scale=None, bypass_pos_embed=True)
+95
View File
@@ -0,0 +1,95 @@
from typing import List
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import (
Action,
MultiArgAction,
)
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class RandAction(MultiArgAction):
grammar = 'rand(" arg ")"'
name = "rand"
chars = None
display_name = "Random Embedding"
action_name = "rand"
description = "Returns a random embedding of the specified token length, with the values optionally bounded by the second and third arguments."
usage_examples = [
"A rand(1) cat",
"A rand(1|-1|1) cat",
]
parsed_token_length = 0
range_min = 0
range_max = 1
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 1 and len(args) != 3:
raise ValueError("Random action should have exactly one argument or three arguments")
self._parse_token_length(args[0])
if len(args) == 3:
self._parse_range(args[1], args[2])
def _parse_token_length(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Random action first argument should have exactly one segment")
token_length_seg_or_action = arg[0]
if isinstance(token_length_seg_or_action, Action):
raise ValueError("Random action should not have an action as an argument")
try:
self.parsed_token_length = int(token_length_seg_or_action.text)
except ValueError:
raise ValueError("Random action should have an integer as the first argument")
def _parse_range(self, min_arg: List[SegOrAction], max_arg: List[SegOrAction]) -> None:
if len(min_arg) != 1:
raise ValueError("Random action second argument should have exactly one segment")
if len(max_arg) != 1:
raise ValueError("Random action third argument should have exactly one segment")
min_seg_or_action = min_arg[0]
max_seg_or_action = max_arg[0]
if isinstance(min_seg_or_action, Action):
raise ValueError("Random action should not have an action as an argument")
if isinstance(max_seg_or_action, Action):
raise ValueError("Random action should not have an action as an argument")
try:
self.range_min = int(min_seg_or_action.text)
except ValueError:
raise ValueError("Random action should have an integer as the second argument")
try:
self.range_max = int(max_seg_or_action.text)
except ValueError:
raise ValueError("Random action should have an integer as the third argument")
if self.range_min > self.range_max:
raise ValueError("Random action should have the second argument be less than the third argument")
def token_length(self) -> int:
"""
Random returns a random embedding whose length is the number in the argument
:return:
"""
return self.parsed_token_length
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Create random tensor of size
result = torch.empty(1, self.parsed_token_length, embedding_module.embedding_dim).uniform_(self.range_min, self.range_max)
return result
+76
View File
@@ -0,0 +1,76 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import MultiArgAction, Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class ScaleDims(MultiArgAction):
grammar = 'scaleDims(" arg ("|" arg)* ")"'
chars = ["-", "-"]
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[Union[PromptSegment, Action]]]):
super().__init__(args)
self.base_arg = args[0]
self._parse_scale_args(args[1:])
def _parse_scale_args(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
# scaleDims scale args should have format of "<dim>,<scale>" where dim is the dimension to scale and scale is the amount to scale it by
# scaleDims(some words|4,1.5|76,1.2)
self.scale_args = []
for arg in args:
if isinstance(arg, Action):
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got an action")
if len(arg) != 1:
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got multiple segments")
extracted_arg = arg[0]
assert isinstance(extracted_arg, PromptSegment)
if "," not in extracted_arg.text:
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got a segment with no comma: " + extracted_arg.text)
# Split prompt segment into text and scale args
dim, scale = extracted_arg.text.split(",")
try:
# TODO: Check that dim is within the bounds of the embedding
parsed_dim = int(dim)
except ValueError:
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got a segment with a non-integer dim: " + str(dim))
try:
parsed_scale = float(scale)
except ValueError:
raise ValueError("ScaleDims scale args must be in the format of <dim>,<scale>(e.g. 4,1.5) but got a segment with a non-float scale: " + str(scale))
self.scale_args.append((parsed_dim, parsed_scale))
def token_length(self) -> int:
# scaleDims modifies the embeddings of the base segment, so the length is the length of the base segment
return sum(seg_or_action.token_length() for seg_or_action in self.base_arg)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the base segment
all_base_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.base_arg
]
base_embeddings = torch.cat(all_base_embeddings, dim=1)
for dim, scale in self.scale_args:
base_embeddings[0, :, dim] *= scale
return base_embeddings
+76
View File
@@ -0,0 +1,76 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import MultiArgAction, Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class SetDims(MultiArgAction):
grammar = 'setDims(" arg ("|" arg)* ")"'
chars = ["-", "-"]
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[Union[PromptSegment, Action]]]):
super().__init__(args)
self.base_arg = args[0]
self._parse_value_args(args[1:])
def _parse_value_args(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
# setDims args should have format of "<dim>,<value>" where dim is the dimension to set and value is the value to set it to
# setDims(some words|4,-0.01254|76,1.2)
self.value_args = []
for arg in args:
if isinstance(arg, Action):
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got an action")
if len(arg) != 1:
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got multiple segments")
extracted_arg = arg[0]
assert isinstance(extracted_arg, PromptSegment)
if "," not in extracted_arg.text:
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got a segment with no comma: " + extracted_arg.text)
# Split prompt segment into text and value args
dim, value = extracted_arg.text.split(",")
try:
# TODO: Check that dim is within the bounds of the embedding
parsed_dim = int(dim)
except ValueError:
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got a segment with a non-integer dim: " + str(dim))
try:
parsed_value = float(value)
except ValueError:
raise ValueError("SetDims value args must be in the format of <dim>,<value>(e.g. 4,1.5) but got a segment with a non-float scale: " + str(value))
self.value_args.append((parsed_dim, parsed_value))
def token_length(self) -> int:
# setDims modifies the embeddings of the base segment, so the length is the length of the base segment
return sum(seg_or_action.token_length() for seg_or_action in self.base_arg)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the base segment
all_base_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.base_arg
]
base_embeddings = torch.cat(all_base_embeddings, dim=1)
for dim, value in self.value_args:
base_embeddings[0, :, dim] = value
return base_embeddings
+106
View File
@@ -0,0 +1,106 @@
from typing import List, Union
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import MultiArgAction, Action
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.actions.utils import slerp
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class SlerpAction(MultiArgAction):
grammar = 'slerp(" arg "|" arg "|" arg ")"'
chars = ["+", "+"]
display_name = "Slerp"
action_name = "slerp"
description = "Performs a slerp(Interpolation) between two segments or actions, with the given weight. The recommended weight is 0 - 1"
usage_examples = [
"The slerp(cat|dog|0.5) is happy",
]
def __init__(self, args: List[List[Union[PromptSegment, Action]]]) -> None:
super().__init__(args)
if len(args) != 3:
raise ValueError("Slerp action should have exactly three arguments(2 vectors and a weight)")
self.start_argument = args[0]
self.end_argument = args[1]
self._parse_weight(args[2])
self._validate_args()
def _parse_weight(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError("Slerp weight should have exactly one segment")
weight_seg_or_action = arg[0]
if isinstance(weight_seg_or_action, Action):
raise ValueError("Slerp weight should not have an action as an argument")
try:
self.parsed_weight = float(weight_seg_or_action.text)
except ValueError:
raise ValueError("Slerp should have an integer/float as the weight")
def _validate_args(self) -> None:
start_arg_token_length = sum(seg_or_action.token_length() for seg_or_action in self.start_argument)
end_arg_token_length = sum(seg_or_action.token_length() for seg_or_action in self.end_argument)
if start_arg_token_length != end_arg_token_length:
raise ValueError(f"Slerp start and end arguments should have the same length. Got {start_arg_token_length} and {end_arg_token_length}")
if self.parsed_weight < 0 or self.parsed_weight > 1:
print(f"WARNING: Slerp weight should be between 0 and 1. Got {self.parsed_weight}")
def token_length(self) -> int:
# Slerp interpolates between the embeddings of the start and end segments, so the length is the length of the start segment
return sum(seg_or_action.token_length() for seg_or_action in self.start_argument)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the start segment
all_start_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.start_argument
]
start_embedding = torch.cat(all_start_embeddings, dim=1)
# Calculate the embeddings for the end segment
all_end_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.end_argument
]
end_embedding = torch.cat(all_end_embeddings, dim=1)
# Perform the slerp
result = slerp(self.parsed_weight, start_embedding, end_embedding)
return result
# def __repr__(self):
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
def __repr__(self) -> str:
return f"sum({', '.join(map(str, self.additional_args))})"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_arg, Action):
base_segment_repr = self.base_arg.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.additional_args, Action):
target_repr = self.additional_args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
+30 -41
View File
@@ -1,57 +1,45 @@
from typing import Union
from typing import Union, List
import torch
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 custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.action.base import Action, MultiArgAction
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
class SumAction(Action):
grammar = 'sum(" arg ("|" arg)* ")"'
name = "sum"
class SumAction(MultiArgAction):
grammar = 'sum(" arg ("|" arg)+ ")"'
chars = ["+", "+"]
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[Union[PromptSegment, Action]]]) -> 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
return sum(seg_or_action.token_length() for seg_or_action in self.base_arg)
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
# Calculate the embeddings for the base segment
all_base_embeddings = [
get_embedding(seg_or_action, embedding_module)
for seg_or_action in self.base_segment
for seg_or_action in self.base_arg
]
result = torch.cat(all_base_embeddings, dim=1)
for arg in self.args:
for arg in self.additional_args:
all_arg_embeddings = [
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
]
@@ -76,21 +64,22 @@ class SumAction(Action):
# 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))})"
return f"sum({', '.join(map(str, self.additional_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)
if isinstance(self.base_arg, Action):
base_segment_repr = self.base_arg.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f"base_segment={self.base_segment.depth_repr()},\n"
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
if isinstance(self.args, Action):
target_repr = self.args.depth_repr(depth + 1)
if isinstance(self.additional_args, Action):
target_repr = self.additional_args.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.args.depth_repr()},\n"
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
+2 -2
View File
@@ -1,6 +1,6 @@
from typing import Union
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
from custom_nodes.ClipStuff.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
from custom_nodes.KepPromptLang.lib.action.base import Action
SegOrAction = Union[PromptSegment, Action]
+34 -2
View File
@@ -1,6 +1,38 @@
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
from typing import List
def batch_size_info(batch: list[SegOrAction]):
import torch
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
def batch_size_info(batch: List[SegOrAction]):
for segment in batch:
print("Token Len: " + str(segment.token_length()))
print(segment.depth_repr())
def slerp(val: float, low: torch.Tensor, high: torch.Tensor, epsilon=1e-5):
# Convert val to tensor and clamp between 0 and 1
val = torch.tensor(val, dtype=torch.float32).clamp(0, 1)
# Normalize the vectors
low_norm = low / torch.norm(low, dim=-1, keepdim=True)
high_norm = high / torch.norm(high, dim=-1, keepdim=True)
# Calculate the cosine of the angle between the vectors
dot = (low_norm * high_norm).sum(-1, keepdim=True)
# Clamp to prevent numerical errors
dot = torch.clamp(dot, -1, 1)
omega = torch.acos(dot)
# Slerp formula
sin_omega = torch.sin(omega)
scale_0 = torch.sin((1.0 - val) * omega) / (sin_omega + epsilon)
scale_1 = torch.sin(val * omega) / (sin_omega + epsilon)
# Handle the case where omega is small (the vectors are close)
close_condition = sin_omega < epsilon
scale_0 = torch.where(close_condition, 1.0 - val, scale_0)
scale_1 = torch.where(close_condition, val, scale_1)
return scale_0 * low + scale_1 * high
+23
View File
@@ -0,0 +1,23 @@
{
"architectures": [
"CLIPTextModel"
],
"attention_dropout": 0.0,
"bos_token_id": 0,
"dropout": 0.0,
"eos_token_id": 2,
"hidden_act": "gelu",
"hidden_size": 1280,
"initializer_factor": 1.0,
"initializer_range": 0.02,
"intermediate_size": 5120,
"layer_norm_eps": 1e-05,
"max_position_embeddings": 77,
"model_type": "clip_text_model",
"num_attention_heads": 20,
"num_hidden_layers": 32,
"pad_token_id": 1,
"projection_dim": 1280,
"torch_dtype": "float32",
"vocab_size": 49408
}
+55 -15
View File
@@ -1,19 +1,22 @@
import contextlib
import os
from typing import List
import torch
from transformers import CLIPTextConfig, modeling_utils
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.sd1_clip import SD1ClipModel
from comfy.sdxl_clip import SDXLClipModel
from custom_nodes.KepPromptLang.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.fun_clip_stuff import PromptLangTextModel
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
# Methods with no comment can be assumed to be the same as comfy.sd1_clip.SD1ClipModel
class PromptLangClipModel(torch.nn.Module):
class PromptLangSDClipModel(torch.nn.Module):
"""Uses the CLIP transformer encoder for text (from huggingface)"""
LAYERS = [
"last",
@@ -23,7 +26,7 @@ class PromptLangClipModel(torch.nn.Module):
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
textmodel_path=None, dtype=None): # clip-vit-base-patch32
super().__init__()
assert layer in self.LAYERS
self.num_layers = 12
@@ -38,18 +41,22 @@ class PromptLangClipModel(torch.nn.Module):
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 comfy.ops.use_comfy_ops(device, dtype):
with modeling_utils.no_init_weights():
# Our transformer
self.transformer = PromptLangTextModel(config)
if dtype is not None:
self.transformer.to(dtype)
self.max_length = max_length
if freeze:
self.freeze()
self.layer = layer
self.layer_idx = None
self.empty_tokens = [[49406] + [49407] * 76]
self.text_projection = None
self.text_projection = torch.nn.Parameter(torch.eye(self.transformer.get_input_embeddings().weight.shape[1]))
self.logit_scale = torch.nn.Parameter(torch.tensor(4.6055))
self.layer_norm_hidden_state = True
if layer == "hidden":
assert layer_idx is not None
@@ -75,7 +82,7 @@ class PromptLangClipModel(torch.nn.Module):
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):
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 = []
@@ -101,12 +108,11 @@ class PromptLangClipModel(torch.nn.Module):
tokens_temp += [next_new_token]
next_new_token += 1
else:
print("WARNING: shape mismatch when trying to apply embedding, embedding will be ignored",
raise Exception("WARNING: shape mismatch when trying to apply embedding. Should have been caught during tokenization.",
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))])
# This should never happen...
raise Exception("Segment size mismatch. Please submit an issue on Github.")
segment.tokens = tokens_temp
n = token_dict_size
@@ -165,18 +171,22 @@ class PromptLangClipModel(torch.nn.Module):
pooled_output = outputs.pooler_output
if self.text_projection is not None:
pooled_output = pooled_output.to(self.text_projection.device) @ self.text_projection
pooled_output = pooled_output.float().to(self.text_projection.device) @ self.text_projection.float()
return z.float(), pooled_output.float()
def encode(self, tokens):
return self(tokens)
def load_sd(self, sd):
if "text_projection" in sd:
self.text_projection[:] = sd.pop("text_projection")
if "text_projection.weight" in sd:
self.text_projection[:] = sd.pop("text_projection.weight").transpose(0, 1)
return self.transformer.load_state_dict(sd, strict=False)
# Changed from comfy.sd1_clip.ClipTokenWeightEncoder
# Changed to use PromptSegments
def encode_token_weights(self, prompt_segments: list[list[SegOrAction]]):
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)
@@ -200,3 +210,33 @@ class PromptLangClipModel(torch.nn.Module):
if (len(output) == 0):
return z_empty.cpu(), first_pooled.cpu()
return torch.cat(output, dim=-2).cpu(), first_pooled.cpu()
class PromptLangSD1ClipModel(SD1ClipModel):
def __init__(self, device="cpu", dtype=None, clip_name="l", clip_model=PromptLangSDClipModel):
super().__init__()
self.clip_name = clip_name
self.clip = "clip_{}".format(self.clip_name)
setattr(self, self.clip, clip_model(device=device, dtype=dtype))
class PromptLangSDXLClipModel(SDXLClipModel):
def __init__(self, device="cpu", dtype=None) -> None:
# Skip SDXLClipModel's init
super(SDXLClipModel, self).__init__()
self.clip_l = PromptLangSDClipModel(layer="hidden", layer_idx=11, device=device, dtype=dtype)
self.clip_l.layer_norm_hidden_state = False
self.clip_g = PromptLangSDXLClipG(device, dtype)
class PromptLangSDXLClipG(PromptLangSDClipModel):
def __init__(self, device="cpu", max_length=77, freeze=True, layer="penultimate", layer_idx=None, textmodel_path=None, dtype=None):
if layer == "penultimate":
layer="hidden"
layer_idx=-2
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "clip_config_bigg.json")
super().__init__(device=device, freeze=freeze, layer=layer, layer_idx=layer_idx, textmodel_json_config=textmodel_json_config, textmodel_path=textmodel_path, dtype=dtype)
self.empty_tokens = [[49406] + [49407] + [0] * 75]
self.layer_norm_hidden_state = False
def load_sd(self, sd):
return super().load_sd(sd)
+126 -27
View File
@@ -1,13 +1,23 @@
from typing import Optional, Tuple, Union
from typing import Optional, Tuple, Union, List, TypedDict, TYPE_CHECKING
from importlib.metadata import version as import_version
from packaging import version
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 transformers.models.clip.modeling_clip import (
CLIPTextEmbeddings,
CLIPTextTransformer,
CLIPTextModel,
)
from custom_nodes.KepPromptLang.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
if TYPE_CHECKING:
from torch import Tensor
from custom_nodes.KepPromptLang.lib.action.base import PostModifiers
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)
@@ -19,30 +29,64 @@ def slerp(val, low, high):
res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high
return res
class PosModifier(TypedDict):
"""
A dictionary of post modifiers for an action result.
"""
position_embed_scale: Union[float]
start_idx: Union[int]
end_idx: Union[int]
class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
def __init__(self, config: CLIPTextConfig):
super().__init__(config)
def forward(
self,
input_dicts: Optional[list[list[SegOrAction]]] = None,
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 = []
batches: List[List[Tensor | Tuple[Tensor, PostModifiers] | Action]] = []
pos_modifiers: List[List[PosModifier]] = []
for batch_idx, batch in enumerate(input_dicts):
results = []
batch_pos_modifiers = []
token_idx = 0
for seg_or_action in batch:
if isinstance(seg_or_action, Action):
results.append(seg_or_action.get_result(self.token_embedding))
action_result: Union[
Tensor, Tuple[Tensor, PostModifiers]
] = seg_or_action.get_result(self.token_embedding)
if isinstance(action_result, tuple):
result, post_modifiers = action_result
if post_modifiers.get("position_embed_scale", None) is not None:
post_modifiers["start_idx"] = token_idx
post_modifiers["end_idx"] = (
token_idx + seg_or_action.token_length()
)
if post_modifiers.get("bypass_pos_embed", False):
post_modifiers["start_idx"] = token_idx
post_modifiers["end_idx"] = (
token_idx + seg_or_action.token_length()
)
batch_pos_modifiers.append(post_modifiers)
else:
result = action_result
else:
results.append(seg_or_action.get_embeddings(self.token_embedding))
result = seg_or_action.get_embeddings(self.token_embedding)
results.append(result)
token_idx += seg_or_action.token_length()
batches.append(results)
pos_modifiers.append(batch_pos_modifiers)
seq_length = batches[0][0].shape[-2]
@@ -56,8 +100,28 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
else:
embeds.append(torch.cat(batch, dim=-2))
position_embeddings = self.position_embedding(position_ids)
embeddings = torch.cat(embeds, dim=0) + position_embeddings
# Iterate over the batches and apply the pos modifiers to the position embeddings then add them to the embeddings
for idx, batch_pos_modifiers in enumerate(pos_modifiers):
position_embeddings = self.position_embedding(position_ids)
if len(batch_pos_modifiers) > 0:
print(f"Found {len(batch_pos_modifiers)} pos modifiers for batch {idx}")
# Apply each pos modifier to the position embeddings at the specified indices
for post_modifier in batch_pos_modifiers:
if post_modifier.get("bypass_pos_embed", False):
position_embeddings[
0, post_modifier["start_idx"] : post_modifier["end_idx"]
] = 0
elif post_modifier["position_embed_scale"] is not None:
position_embeddings[
0, post_modifier["start_idx"] : post_modifier["end_idx"]
] *= post_modifier["position_embed_scale"]
else:
raise ValueError(
"Pos modifier must have a scale or bypass_pos_embed"
)
# Add the possibly modified position embeddings to the embeddings
embeds[idx] = embeds[idx] + position_embeddings
embeddings = torch.cat(embeds, dim=0)
return embeddings
@@ -66,10 +130,44 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
def __init__(self, config: CLIPTextConfig):
super().__init__(config)
self.embeddings = PromptLangCLIPTextEmbeddings(config)
self.transformers_version = version.parse(import_version('transformers'))
def process_attention_mask(self, hidden_states, attention_mask, bsz, seq_len):
# Parse the transformer version
input_shape = torch.Size([bsz, seq_len])
v4_30 = version.parse('4.30.0')
v4_35 = version.parse('4.35')
if self.transformers_version < v4_30:
print("Using transformers < 4.30.0")
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(
hidden_states.device)
elif v4_30 <= self.transformers_version < v4_35:
print("Using transformers >= 4.30.0 and <= 4.34.*")
from transformers.models.clip.modeling_clip import _make_causal_mask
causal_attention_mask = _make_causal_mask(input_shape, hidden_states.dtype, device=hidden_states.device)
else:
print("Using transformers >= 4.35")
from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask
causal_attention_mask = _create_4d_causal_attention_mask(
input_shape, hidden_states.dtype, device=hidden_states.device
)
# Expand attention_mask if it exists
if attention_mask is not None:
# Import _expand_mask or _prepare_4d_attention_mask based on version
if self.transformers_version < v4_35:
from transformers.models.clip.modeling_clip import _expand_mask
attention_mask = _expand_mask(attention_mask, hidden_states.dtype)
else:
from transformers.modeling_attn_mask_utils import _prepare_4d_attention_mask
attention_mask = _prepare_4d_attention_mask(attention_mask, hidden_states.dtype)
return causal_attention_mask, attention_mask
def forward(
self,
input_ids: Optional[list[list[SegOrAction]]] = None,
input_ids: Optional[List[List[SegOrAction]]] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
@@ -92,21 +190,19 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
# input_shape = input_ids.size()
# input_ids = input_ids.view(-1, input_shape[-1])
for batch_idx, batch in enumerate(input_ids):
for seg_or_action in batch:
if isinstance(seg_or_action, Action):
seg_or_action.process_with_transformer(
self, self.embeddings.token_embedding
)
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)
seq_len = hidden_states.shape[1]
causal_attention_mask, attention_mask = self.process_attention_mask(hidden_states, attention_mask, bsz, seq_len)
encoder_outputs = self.encoder(
inputs_embeds=hidden_states,
@@ -129,8 +225,11 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
if isinstance(seg_or_action, Action):
idx += seg_or_action.token_length()
else:
if seg_or_action.text == '__PAD__':
if seg_or_action.text == "__PAD__" or seg_or_action.text == "[EOT]":
break
# Is a segment, and isn't the pad segment
idx += seg_or_action.token_length()
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)
@@ -160,7 +259,7 @@ class PromptLangTextModel(CLIPTextModel):
def forward(
self,
input_ids: Optional[list[list[SegOrAction]]] = None,
input_ids: Optional[List[List[SegOrAction]]] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
+4 -12
View File
@@ -3,24 +3,16 @@ grammar = """
item: embedding
| WORD
| function
| generic_function
| QUOTED_STRING
function: sum_function
| neg_function
| norm_function
| diff_function
sum_function: "sum(" arg ("|" arg)* ")"
neg_function: "neg(" arg ")"
norm_function: "norm(" arg ")"
diff_function: "diff(" arg ("|" arg)* ")"
generic_function: FUNC_NAME "(" arg ("|" arg)* ")"
arg: item+
embedding: "embedding:" WORD
WORD: /[A-Za-z0-9,_-]+/
FUNC_NAME: /[A-Za-z_-]+/
WORD: /[A-Za-z0-9,_\.-]+/
QUOTED_STRING: /"([^"\\\]*(\\\.[^"\\\]*)*)"|'([^'\\\]*(\\\.[^'\\\]*)*)'/
%import common.WS
+3 -3
View File
@@ -1,4 +1,4 @@
from typing import Union
from typing import Union, List
import torch
from torch import Tensor
@@ -6,7 +6,7 @@ from torch.nn import Embedding
class PromptSegment:
def __init__(self, text: str, tokens: list[Union[int, Tensor]]):
def __init__(self, text: str, tokens: List[Union[int, Tensor]]):
self.text = text
self.tokens = tokens
@@ -17,7 +17,7 @@ class PromptSegment:
return len(self.tokens)
def get_embeddings(self, embedding_module: Embedding) -> Tensor:
tensors = torch.LongTensor(self.tokens).to(torch.device('cpu'))
tensors = torch.LongTensor(self.tokens).to(embedding_module.weight.device)
unsqueezed_tensors = tensors.unsqueeze(0)
return embedding_module(unsqueezed_tensors)
+20
View File
@@ -0,0 +1,20 @@
from typing import Type, Dict
from custom_nodes.KepPromptLang.lib.action.base import Action
action_registry: Dict[str, Type[Action]] = {}
def register_action(action: Type[Action]) -> None:
"""
:rtype: object
"""
if action.action_name in action_registry:
raise ValueError(f"Action {action.action_name} already registered")
action_registry[str(action.action_name)] = action
def get_action_by_name(name: str) -> Type[Action]:
if name not in action_registry:
raise ValueError(f"Action {name} not found in registry")
return action_registry[name]
+25 -27
View File
@@ -1,29 +1,29 @@
from typing import List
from lark import Transformer, Token
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 comfy.sd1_clip import SDTokenizer
from custom_nodes.KepPromptLang.lib.action.base import Action, ActionArity
from custom_nodes.KepPromptLang.lib.parser.registration import get_action_by_name
from custom_nodes.KepPromptLang.lib.parser.utils import build_prompt_segment
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class PromptTransformer(Transformer):
"""
Transforms the parsed prompt into a list of segments and actions
Types from the grammar are mapped to the methods in this class
"""
# def WORD(self, items):
# return items
def __init__(self, tokenizer: SD1Tokenizer):
def __init__(self, tokenizer: SDTokenizer):
super().__init__()
self.tokenizer = tokenizer
def item(self, items: list[Token]):
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)):
return item
if item.type == "WORD":
@@ -37,7 +37,7 @@ class PromptTransformer(Transformer):
elif item.type == "embedding":
return build_prompt_segment(item, self.tokenizer)
elif item.type == "function":
return item
raise Exception("Unexpected type in prompt transformer: function. Please report this issue on GitHub.")
else:
raise Exception("Unknown item type: " + str(item.type))
@@ -47,15 +47,13 @@ class PromptTransformer(Transformer):
def embedding(self, items):
return build_prompt_segment(f'{self.tokenizer.embedding_identifier}{items[0]}', self.tokenizer)
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 generic_function(self, items):
action = get_action_by_name(items[0])
if action.arity == ActionArity.SINGLE:
if len(items) != 2:
raise ValueError(f"Action {action.name} should have exactly one argument")
return action(items[1])
elif action.arity == ActionArity.MULTI:
return action(items[1:][:])
else:
raise ValueError(f"Unknown action arity: {action.arity}")
+9 -6
View File
@@ -1,7 +1,7 @@
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 custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
def flatten_tree(tree):
@@ -11,7 +11,7 @@ def flatten_tree(tree):
return [str(tree.data)] + sum([flatten_tree(child) for child in tree.children], [])
def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment:
def build_prompt_segment(text: str, tokenizer: SDTokenizer) -> PromptSegment:
split_text = text.split(" ")
tokens = []
for word in split_text:
@@ -24,10 +24,13 @@ def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment:
if embedding is None:
print(f"warning, embedding:{embedding_name} does not exist, ignoring")
else:
if len(embedding.shape) == 1:
tokens.append(embedding)
if embedding.shape[1] != tokenizer.embedding_size:
print(f"warning, embedding:{embedding_name} has size {embedding.shape[1]}, expected {tokenizer.embedding_size}, ignoring")
else:
tokens.extend(embedding)
if len(embedding.shape) == 1:
tokens.append(embedding)
else:
tokens.extend(embedding)
if leftover != "":
word = leftover
+33 -11
View File
@@ -1,22 +1,24 @@
from typing import List, Dict
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.KepPromptLang.lib.actions.types import SegOrAction
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 custom_nodes.KepPromptLang.lib.parser import PromptParser
from custom_nodes.KepPromptLang.lib.parser.transformer import PromptTransformer
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
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):
class PromptLangSDTokenizer(SDTokenizer):
def __init__(self, tokenizer_path=None, max_length=77, pad_with_end=True, embedding_directory=None, embedding_size=768, embedding_key='clip_l'):
super().__init__(tokenizer_path, max_length, pad_with_end, embedding_directory, embedding_size, embedding_key)
"""
Doesn't actually tokenize...
Returns batches of segments and actions
:return: List of list(batches) of segments and actions
"""
def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs) -> list[list[SegOrAction]]:
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:
@@ -41,9 +43,9 @@ class PromptLangTokenizer(SD1Tokenizer):
# 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
remaining_length = self.max_length - batch_size
# Pad batch
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length - 1))
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * (remaining_length - 1))) # -1 for end token
batched_segments.append(batch)
# start new batch
@@ -64,3 +66,23 @@ class PromptLangTokenizer(SD1Tokenizer):
# batch_size_info(batch)
return batched_segments
class PromptLangSD1Tokenizer(SD1Tokenizer):
def __init__(self, embedding_directory=None, clip_name='l', tokenizer=PromptLangSDTokenizer) -> None:
super().__init__(embedding_directory, clip_name, tokenizer)
class PromptLangSDXLClipGTokenizer(PromptLangSDTokenizer):
def __init__(self, tokenizer_path=None, embedding_directory=None):
super().__init__(tokenizer_path, pad_with_end=False, embedding_directory=embedding_directory, embedding_size=1280, embedding_key='clip_g')
class PromptLangSDXLTokenizer(SD1Tokenizer):
def __init__(self, embedding_directory=None) -> None:
self.clip_l = PromptLangSDTokenizer(embedding_directory=embedding_directory)
self.clip_g = PromptLangSDXLClipGTokenizer(embedding_directory=embedding_directory)
def tokenize_with_weights(self, text:str, return_word_ids=False) -> Dict[str, List[List[SegOrAction]]]:
out = {}
out["g"] = self.clip_g.tokenize_with_weights(text, return_word_ids)
out["l"] = self.clip_l.tokenize_with_weights(text, return_word_ids)
return out
+80 -40
View File
@@ -1,4 +1,5 @@
import random
import os
from typing import List, Tuple, Any
import numpy as np
from PIL import Image
@@ -6,9 +7,18 @@ from PIL import Image
import folder_paths
import comfy.sd
import comfy.ops
from custom_nodes.ClipStuff.lib.clip_model import PromptLangClipModel
from comfy.sd2_clip import SD2ClipModel
from comfy.sdxl_clip import SDXLClipModel
from comfy.supported_models_base import ClipTarget
from custom_nodes.KepPromptLang.lib.clip_model import (
PromptLangSDXLClipModel,
PromptLangSD1ClipModel,
)
from custom_nodes.ClipStuff.lib.tokenizer import PromptLangTokenizer
from custom_nodes.KepPromptLang.lib.tokenizer import (
PromptLangSDXLTokenizer,
PromptLangSD1Tokenizer,
)
class EmptyClass:
@@ -17,7 +27,7 @@ class EmptyClass:
class SpecialClipLoader:
@classmethod
def INPUT_TYPES(s):
def INPUT_TYPES(cls): # type: ignore
return {
"required": {
"source_clip": ("CLIP",),
@@ -30,35 +40,43 @@ class SpecialClipLoader:
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]:
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()
)
if isinstance(source_clip.cond_stage_model, SDXLClipModel):
clip_target = ClipTarget(PromptLangSDXLTokenizer, PromptLangSDXLClipModel)
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.clip_g.embedding_directory)
comfy.sd.load_clip_weights(clip.cond_stage_model.clip_g,source_clip.cond_stage_model.clip_g.state_dict())
comfy.sd.load_clip_weights(
clip.cond_stage_model.clip_l, source_clip.cond_stage_model.clip_l.state_dict()
)
elif isinstance(source_clip, SD2ClipModel):
raise ValueError("SD2 Clip model is not supported.")
else:
clip_target = ClipTarget(PromptLangSD1Tokenizer, PromptLangSD1ClipModel)
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.clip_l.embedding_directory)
comfy.sd.load_clip_weights(
clip.cond_stage_model, source_clip.cond_stage_model.state_dict()
)
return (clip,)
def tensor2img(tensor_img):
def tensor2img(tensor_img) -> Image.Image:
i = 255.0 * tensor_img.cpu().numpy()
i_np_arr = np.clip(i, 0, 255, out=i).astype(np.uint8, copy=False)
return Image.fromarray(i_np_arr)
class BuildGif:
def __init__(self):
def __init__(self) -> None:
self.output_dir = folder_paths.get_output_directory()
pass
@classmethod
def INPUT_TYPES(cls):
def INPUT_TYPES(cls): # type: ignore
return {
"required": {
"images": ("IMAGE",),
"split_every": ("INT", {"default": -1}),
"frame_duration": ("INT", {"default": 125}),
"output_mode": (
["One Per Split", "Big Grid"],
{"default": "Big Grid"},
@@ -67,46 +85,53 @@ class BuildGif:
}
RELOAD_INST = True
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("Gifs",)
RETURN_TYPES = ()
# RETURN_NAMES = ("Gifs",)
INPUT_IS_LIST = True
FUNCTION = "build_gif"
OUTPUT_IS_LIST = (True,)
# OUTPUT_NODE = False
# OUTPUT_IS_LIST = (True,)
OUTPUT_NODE = True
CATEGORY = "List Stuff"
@staticmethod
def build_gif(images: list, split_every: list[int], output_mode: str):
def build_gif(self, images: List[Any], split_every: List[int], frame_duration: List[int], output_mode: List[str]):
print("Build GIF called!")
print(f"{type(images)}")
if len(split_every) > 1:
raise Exception("List input for split every is not supported.")
split_every = split_every[0]
batch_size = images[0].size()[0]
if split_every == -1:
split_chunks = 1
split_every = len(images)
else:
split_chunks = int(len(images) / split_every)
if len(output_mode) > 1:
raise Exception("List input for output_mode is not supported.")
output_mode = output_mode[0]
out = []
if len(frame_duration) > 1:
raise Exception("List input for frame_duration is not supported.")
frame_duration = frame_duration[0]
full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path(filename_prefix="Gif", output_dir=self.output_dir, image_width=0, image_height=0)
split_every_val = split_every[0]
batch_size = images[0].size()[0]
if split_every_val == -1:
split_chunks = 1
split_every_val = len(images)
else:
split_chunks = int(len(images) / split_every_val)
num_wide = batch_size
num_tall = split_chunks
chunked_batches = [
images[split_every * chunk_idx : split_every * (chunk_idx + 1)]
images[split_every_val * chunk_idx : split_every_val * (chunk_idx + 1)]
for chunk_idx in range(split_chunks)
]
frames = []
results = list()
if output_mode == "Big Grid":
# For every image in gif
for idx_in_chunk in range(split_every):
for idx_in_chunk in range(split_every_val):
img_shape = images[0][0].shape
img_frame = Image.new(
"RGB", size=(num_wide * img_shape[0], num_tall * img_shape[1])
@@ -121,8 +146,9 @@ class BuildGif:
)
frames.append(img_frame)
file = f"{filename}_{counter:05}_"
save_path = (
f"{folder_paths.get_output_directory()}/{random.randint(1, 100)}"
f"{os.path.join(full_output_folder, file)}"
)
frames[0].save(
f"{save_path}.webp",
@@ -132,15 +158,24 @@ class BuildGif:
save_all=True,
append_images=frames[1:],
optimize=False,
duration=125,
duration=frame_duration,
loop=0,
)
results.append({
"filename": f"{file}.webp",
"subfolder": subfolder,
"type": "output"
})
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)
split_start = split_every_val * split_idx
split_end = split_every_val * (split_idx + 1)
for batch_idx in range(batch_size):
save_path = f"{folder_paths.get_output_directory()}/-{batch_idx}-{random.randint(1, 100)}"
file = f"{filename}_{counter:05}_"
save_path = (
f"{os.path.join(full_output_folder, file)}"
)
counter += 1
print(save_path)
tensor2img(images[split_start][batch_idx]).save(
f"{save_path}.webp",
@@ -150,7 +185,12 @@ class BuildGif:
for nested_batch in images[split_start + 1 : split_end]
],
optimize=False,
duration=125,
duration=frame_duration,
loop=0,
)
return (out,)
results.append({
"filename": f"{file}.webp",
"subfolder": subfolder,
"type": "output"
})
return { "ui": { "images": results } }
+1
View File
@@ -1 +1,2 @@
lark
packaging
+12
View File
@@ -0,0 +1,12 @@
# If the models/checkpoints folder does not have test.txt, then download the model.
import os
from huggingface_hub import hf_hub_download
FILE = "v1-5-pruned-emaonly.safetensors"
REPO_ID = "runwayml/stable-diffusion-v1-5"
if not os.path.exists(f"models/checkpoints/{FILE}"):
print("Downloading model...")
hf_hub_download(repo_id=REPO_ID, filename=FILE, local_dir="models/checkpoints", local_dir_use_symlinks=False)
else:
print("Model already downloaded.")
+109
View File
@@ -0,0 +1,109 @@
#This is an example that uses the websockets api to know when a prompt execution is done
#Once the prompt execution is done it downloads the images using the /history endpoint
import os
import websocket #NOTE: websocket-client (https://github.com/websocket-client/websocket-client)
import uuid
import json
import urllib.request
import urllib.parse
server_address = "127.0.0.1:8188"
client_id = str(uuid.uuid4())
def queue_prompt(prompt):
p = {"prompt": prompt, "client_id": client_id}
data = json.dumps(p).encode('utf-8')
req = urllib.request.Request("http://{}/prompt".format(server_address), data=data)
try:
response = urllib.request.urlopen(req)
return json.loads(response.read())
except urllib.error.HTTPError as e:
print(f"HTTP Error {e.code}: {e.reason}")
error_body = e.read()
# Attempt to read and print the JSON error body
try:
json_error_body = json.loads(error_body)
print(json_error_body)
raise e
except json.JSONDecodeError:
print("Failed to decode error response as JSON.")
print(error_body)
raise e
except Exception as e:
print(f"An unexpected error occurred: {e}")
raise e
def get_image(filename, subfolder, folder_type):
data = {"filename": filename, "subfolder": subfolder, "type": folder_type}
url_values = urllib.parse.urlencode(data)
with urllib.request.urlopen("http://{}/view?{}".format(server_address, url_values)) as response:
return response.read()
def get_history(prompt_id):
with urllib.request.urlopen("http://{}/history/{}".format(server_address, prompt_id)) as response:
return json.loads(response.read())
def get_images(ws, prompt):
prompt_id = queue_prompt(prompt)['prompt_id']
output_images = {}
while True:
out = ws.recv()
if isinstance(out, str):
message = json.loads(out)
if message['type'] == 'executing':
data = message['data']
if data['node'] is None and data['prompt_id'] == prompt_id:
break #Execution is done
else:
continue #previews are binary data
history = get_history(prompt_id)[prompt_id]
for o in history['outputs']:
for node_id in history['outputs']:
node_output = history['outputs'][node_id]
if 'images' in node_output:
images_output = []
for image in node_output['images']:
image_data = get_image(image['filename'], image['subfolder'], image['type'])
images_output.append(image_data)
output_images[node_id] = images_output
return output_images
# Load json from file relative to this script
prompt = json.load(open(os.path.join(os.path.dirname(os.path.realpath(__file__)), "workflow_api.json")))
#set the text prompt for our positive CLIPTextEncode
# prompt["6"]["inputs"]["text"] = "masterpiece best quality man"
#set the seed for our KSampler node
print(queue_prompt(prompt))
ws = websocket.WebSocket()
ws.connect("ws://{}/ws?clientId={}".format(server_address, client_id))
while True:
out = ws.recv()
if isinstance(out, str):
message = json.loads(out)
# print(message)
if message["type"] == "executing" and message["data"]["node"] is None:
print("Execution is done")
break
if message["type"] == "execution_error":
print("Execution error")
print(json.dumps(message["data"], indent=4))
raise Exception("Execution error")
# images = get_images(ws, prompt)
#Commented out code to display the output images:
# for node_id in images:
# for image_data in images[node_id]:
# from PIL import Image
# import io
# image = Image.open(io.BytesIO(image_data))
# image.show()
+84
View File
@@ -0,0 +1,84 @@
{
"1": {
"inputs": {
"text": "A sum(cat|norm(sum(neg(parrot)|rabbit))) outside",
"clip": [
"2",
0
]
},
"class_type": "CLIPTextEncode"
},
"2": {
"inputs": {
"source_clip": [
"4",
1
]
},
"class_type": "Special CLIP Loader"
},
"3": {
"inputs": {
"seed": 556492279461741,
"steps": 1,
"cfg": 8,
"sampler_name": "euler",
"scheduler": "normal",
"denoise": 1,
"model": [
"4",
0
],
"positive": [
"1",
0
],
"negative": [
"1",
0
],
"latent_image": [
"7",
0
]
},
"class_type": "KSampler"
},
"4": {
"inputs": {
"ckpt_name": "v1-5-pruned-emaonly.safetensors"
},
"class_type": "CheckpointLoaderSimple"
},
"5": {
"inputs": {
"samples": [
"3",
0
],
"vae": [
"4",
2
]
},
"class_type": "VAEDecode"
},
"6": {
"inputs": {
"images": [
"5",
0
]
},
"class_type": "PreviewImage"
},
"7": {
"inputs": {
"width": 512,
"height": 512,
"batch_size": 1
},
"class_type": "EmptyLatentImage"
}
}
+83
View File
@@ -0,0 +1,83 @@
import importlib
import inspect
import os
from typing import List, Type
from custom_nodes.KepPromptLang.lib.action.base import Action
EXCLUDED_MODULES = ["utils.py", "action_utils.py", "types.py"]
def import_module_from_path(path: str):
module_name = path.replace("/", ".")[:-3]
return importlib.import_module(module_name, package="custom_nodes.KepPromptLang.lib.actions")
# Function to find and import action classes
def find_action_classes(directory: str) -> List[Type[Action]]:
action_classes = []
for filename in os.listdir(directory):
if filename.endswith(".py") and not filename.startswith("__") and filename not in EXCLUDED_MODULES:
module_path = os.path.join(directory, filename)
module = import_module_from_path(module_path)
for name, obj in inspect.getmembers(module, inspect.isclass):
if issubclass(obj, Action) and obj is not Action and obj.__name__ != "MultiArgAction" and obj.__name__ != "SingleArgAction":
action_classes.append(obj)
return action_classes
# Function to extract info from an action class
def extract_class_info(cls: Type[Action]) -> dict:
class_info = {
'class_name': cls.__name__,
'properties': {
'display_name': getattr(cls, 'display_name', None),
'action_name': getattr(cls, 'action_name', None),
'description': getattr(cls, 'description', None),
'usage_examples': getattr(cls, 'usage_examples', None)
}
}
return class_info
def escape_pipes(text: str) -> str:
return text.replace('|', '\\|')
def generate_markdown_documentation(classes_info: List[dict]) -> str:
documentation = "# Actions Documentation\n\n"
# Define table columns
# columns = ["Class", "Display Name", "Action Name", "Description", "Usage Examples"]
columns = ["Display Name", "Action Name", "Description", "Usage Examples"]
documentation += "| " + " | ".join(columns) + " |\n"
documentation += "| --- " * len(columns) + "|\n"
for cls_info in classes_info:
# row = [cls_info['class_name']]
row = []
# Iterate over properties in a predefined order
for prop in ["display_name", "action_name", "description", "usage_examples"]:
prop_doc = cls_info['properties'].get(prop, 'N/A')
# Format and escape usage examples
if isinstance(prop_doc, list):
escaped_examples = [escape_pipes(example) for example in prop_doc]
prop_doc = "<ul>" + "".join([f"<li>{example}</li>" for example in escaped_examples]) + "</ul>"
else:
prop_doc = escape_pipes(prop_doc)
row.append(prop_doc)
documentation += "| " + " | ".join(row) + " |\n"
return documentation
# Main execution
if __name__ == "__main__":
actions_directory = (
"../lib/actions" # Update this path as per your project structure
)
action_classes = find_action_classes(actions_directory)
class_infos = [extract_class_info(cls) for cls in action_classes]
docs = generate_markdown_documentation(class_infos)
print(docs) # Or write to a file