Compare commits
85
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f2256932c8 | ||
|
|
94c46fe1ba | ||
|
|
ca11ea8e82 | ||
|
|
fb363aed7a | ||
|
|
d67d8600ab | ||
|
|
6444e2890a | ||
|
|
28141f2cbe | ||
|
|
70704f5e68 | ||
|
|
e59aaa499d | ||
|
|
c98df8d289 | ||
|
|
a7c4bbe332 | ||
|
|
ce795c52bc | ||
|
|
ef9693ec73 | ||
|
|
5362ac75fa | ||
|
|
4fcdee837b | ||
|
|
6b2bf7485e | ||
|
|
d330663624 | ||
|
|
5de8d6821c | ||
|
|
0f36544b0e | ||
|
|
05920e0392 | ||
|
|
9ccf5d2583 | ||
|
|
f185b39f06 | ||
|
|
af761a5620 | ||
|
|
fee9c56abd | ||
|
|
6879267391 | ||
|
|
a40ed34eac | ||
|
|
94b2347d01 | ||
|
|
bcc083f2fa | ||
|
|
f91d3aa233 | ||
|
|
b497a81b43 | ||
|
|
4076d2beb1 | ||
|
|
e644f219dd | ||
|
|
9e7e79a1a1 | ||
|
|
f18ef4b292 | ||
|
|
78c14ce95a | ||
|
|
60bc4e5a4d | ||
|
|
c11a6c5178 | ||
|
|
c0adb16e30 | ||
|
|
df98b07ffc | ||
|
|
79f17aac60 | ||
|
|
7804426fa2 | ||
|
|
5be4787a12 | ||
|
|
fd2316c8fa | ||
|
|
137a3f0f24 | ||
|
|
c422e6000c | ||
|
|
5054bfdb82 | ||
|
|
7a854c7fea | ||
|
|
9f12d2f16d | ||
|
|
bd369930bd | ||
|
|
9bc7f54bcf | ||
|
|
337dad1cc9 | ||
|
|
ade09bf806 | ||
|
|
88c3804446 | ||
|
|
acbaf7cefe | ||
|
|
1f1e74cd30 | ||
|
|
919f2dbebf | ||
|
|
bd0093f64e | ||
|
|
bdb4d00910 | ||
|
|
c1cbdbe9df | ||
|
|
c28099021b | ||
|
|
2d2d29bd5a | ||
|
|
d0719cb9bc | ||
|
|
3fc1a7d7b9 | ||
|
|
f66eeb8117 | ||
|
|
d5a2cba894 | ||
|
|
0eaa2c1c41 | ||
|
|
553f90691e | ||
|
|
c4769a797b | ||
|
|
66eef6d7d6 | ||
|
|
0fdc4244f1 | ||
|
|
4057ff8347 | ||
|
|
5f84a4b530 | ||
|
|
fc3560be65 | ||
|
|
833a027486 | ||
|
|
6b8a1ea243 | ||
|
|
07cec3f68f | ||
|
|
6a516d093c | ||
|
|
c9f7895f62 | ||
|
|
b125677e6d | ||
|
|
7ced8b7533 | ||
|
|
d883444fdd | ||
|
|
c6dc33a9ae | ||
|
|
110da543ed | ||
|
|
21eecda007 | ||
|
|
235a681f70 |
@@ -0,0 +1,10 @@
|
||||
name: Test
|
||||
on:
|
||||
workflow_dispatch:
|
||||
jobs:
|
||||
build:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v2
|
||||
- name: Setup upterm session
|
||||
uses: lhotari/action-upterm@v1
|
||||
@@ -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
|
||||
|
||||
@@ -1 +1,81 @@
|
||||
# ClipStuff
|
||||
## Basic Instructions.
|
||||
Clone repo into custom_nodes folder.
|
||||
Install the requirements.txt file via pip.
|
||||
|
||||
Pass CLIP output from Load Checkpoint into SpecialClipLoader node, then use the outputted clip with standard Clip Text Encode.
|
||||
|
||||
See example workflow in examples folder.
|
||||
|
||||
### Example Photo
|
||||

|
||||
|
||||
## Functions
|
||||
|
||||
## Syntax Elements
|
||||
|
||||
1. **Embedding**:
|
||||
- Syntax: `embedding:WORD`
|
||||
- Example: `embedding:face_vector`
|
||||
- Represents a named vector embedding(Textual Inversion).
|
||||
|
||||
2. **Word**:
|
||||
- Syntax: Any alphanumeric word including characters such as `,`, `_`, and `-`.
|
||||
- Example: `cat, dog_face, id_123`
|
||||
- Represents simple words or identifiers.
|
||||
|
||||
3. **Quoted String**:
|
||||
- Syntax: A string enclosed within double or single quotes. You can escape quotes inside the string using a backslash (`\`).
|
||||
- Example: `"Hello World"`, `'It\'s a sunny day'`
|
||||
- Represents string literals.
|
||||
|
||||
## 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.
|
||||
- For functions that accept multiple arguments, they are separated by the `|` symbol.
|
||||
|
||||
## Examples
|
||||
|
||||
1. Add two embeddings and normalize the result:
|
||||
```
|
||||
norm(sum(cat | dog | horse | parrot))
|
||||
```
|
||||
|
||||
2. Negate an embedding:
|
||||
```
|
||||
neg(embedding:body_vector)
|
||||
```
|
||||
|
||||
3. King - Man + Woman = Queen:
|
||||
```
|
||||
sum(diff(king|man)|woman)
|
||||
```
|
||||
or
|
||||
```
|
||||
sum(king|neg(man)|woman)
|
||||
```
|
||||
```
|
||||
|
||||
+4
-2
@@ -1,11 +1,13 @@
|
||||
from .nodes import (
|
||||
KepAdvTextEncode,
|
||||
BuildGif,
|
||||
SpecialClipLoader,
|
||||
MonacoPrompt,
|
||||
)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"Kep Adv Text Encode": KepAdvTextEncode,
|
||||
"Build Gif": BuildGif,
|
||||
"Special CLIP Loader": SpecialClipLoader,
|
||||
"Monaco Prompt": MonacoPrompt,
|
||||
}
|
||||
|
||||
WEB_DIRECTORY = ("./web/dist", ["app.bundle.js"])
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 1.8 MiB |
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,22 @@
|
||||
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.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)
|
||||
|
||||
@@ -0,0 +1,117 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from enum import Enum
|
||||
from typing import Union, List
|
||||
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
|
||||
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
|
||||
class ActionArity(Enum):
|
||||
NONE = 0
|
||||
SINGLE = 1
|
||||
MULTI = 2
|
||||
|
||||
class Action(ABC):
|
||||
@property
|
||||
@abstractmethod
|
||||
def chars(self) -> Union[List[str], None]:
|
||||
pass
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def arity(self) -> ActionArity:
|
||||
"""
|
||||
Determines the arity of the action. This is used to determine how many arguments the action supports.
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def name(self) -> str:
|
||||
pass
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def grammar(self) -> str:
|
||||
"""
|
||||
The grammar for this action. This is used to parse the action from the prompt.
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def token_length(self) -> int:
|
||||
"""
|
||||
The length of the tokens that this action will add to the prompt.
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_all_segments(self) -> List[PromptSegment]:
|
||||
"""
|
||||
Get all segments, including nested segments.
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
"""
|
||||
Initialize the action. This is called when the action is parsed from the prompt.
|
||||
:param args: The arguments for the action.
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
||||
"""
|
||||
Get the result of this action. This is called when the embeddings are being calculated.
|
||||
:param embedding_module: The embedding module to use to get the base embeddings for tokens.
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
def depth_repr(self, depth: int = 1) -> str:
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
class SingleArgAction(Action, ABC):
|
||||
arity = ActionArity.SINGLE
|
||||
def get_all_segments(self) -> List[PromptSegment]:
|
||||
segments = []
|
||||
for seg_or_action in self.arg:
|
||||
if isinstance(seg_or_action, Action):
|
||||
segments.extend(seg_or_action.get_all_segments())
|
||||
else:
|
||||
segments.append(seg_or_action)
|
||||
return segments
|
||||
|
||||
def __init__(self, arg: List[Union[PromptSegment, Action]]):
|
||||
# TODO: Target is a list now... what does this mean for us..
|
||||
self.arg = arg
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.name}({self.arg})"
|
||||
|
||||
class MultiArgAction(Action, ABC):
|
||||
arity = ActionArity.MULTI
|
||||
def get_all_segments(self) -> List[PromptSegment]:
|
||||
segments = []
|
||||
for arg in self.all_args:
|
||||
for seg_or_action in arg:
|
||||
if isinstance(seg_or_action, Action):
|
||||
segments.extend(seg_or_action.get_all_segments())
|
||||
else:
|
||||
segments.append(seg_or_action)
|
||||
|
||||
return segments
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
args: List[List[Union[PromptSegment, Action]]],
|
||||
):
|
||||
self.all_args = args
|
||||
@@ -1,6 +0,0 @@
|
||||
from custom_nodes.ClipStuff.lib.actions.arith import ArithAction
|
||||
from custom_nodes.ClipStuff.lib.actions.nudge import NudgeAction
|
||||
|
||||
ALL_ACTIONS = [NudgeAction, ArithAction]
|
||||
ALL_START_CHARS = [action.START_CHAR for action in ALL_ACTIONS]
|
||||
ALL_END_CHARS = [action.END_CHAR for action in ALL_ACTIONS]
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
|
||||
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)
|
||||
@@ -1,131 +0,0 @@
|
||||
from typing import Callable, Union
|
||||
|
||||
import torch
|
||||
from torch.nn import Embedding
|
||||
|
||||
from comfy.sd1_clip import SD1Tokenizer
|
||||
from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment
|
||||
|
||||
|
||||
class ArithAction(Action):
|
||||
START_CHAR = "<"
|
||||
END_CHAR = ">"
|
||||
|
||||
def __init__(self, base_segment: PromptSegment | Action, ops: dict[str, list[PromptSegment | Action]]):
|
||||
self.base_segment = base_segment
|
||||
self.ops = ops
|
||||
|
||||
def __repr__(self):
|
||||
return f"ArithAction(\n\tbase_segment={self.base_segment},\n\tops={self.ops}\n)"
|
||||
|
||||
def depth_repr(self, depth=1):
|
||||
out = "ArithAction(\n"
|
||||
if isinstance(self.base_segment, Action):
|
||||
base_segment_repr = self.base_segment.depth_repr(depth + 1)
|
||||
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
|
||||
elif isinstance(self.base_segment, PromptSegment):
|
||||
out += "\t" * depth + f'base_segment={self.base_segment.depth_repr(depth)}'
|
||||
else:
|
||||
out += "\t" * depth + f'base_segment="{self.base_segment}",'
|
||||
|
||||
for op_key, ops in self.ops.items():
|
||||
for op in ops:
|
||||
out += "\n" + "\t" * depth + f'"{op_key}":[\n'
|
||||
if isinstance(op, Action):
|
||||
op_repr = op.depth_repr(depth + 2)
|
||||
out += "\t" * (depth + 1) + f"{op_repr}\n"
|
||||
else:
|
||||
out += "\t" * (depth + 1) + f'{op.depth_repr()},\n'
|
||||
out += "\t" * depth + "],"
|
||||
out += "\n" + "\t" * (depth - 1) + ")"
|
||||
return out
|
||||
|
||||
def token_length(self):
|
||||
# ArithAction modifies the embeddings of the base segment, so the length is the length of the base segment
|
||||
if isinstance(self.base_segment, Action):
|
||||
return self.base_segment.token_length()
|
||||
|
||||
return len(self.base_segment.tokens)
|
||||
|
||||
|
||||
def get_all_segments(self):
|
||||
segments = []
|
||||
if isinstance(self.base_segment, Action):
|
||||
segments += self.base_segment.get_all_segments()
|
||||
else:
|
||||
segments.append(self.base_segment)
|
||||
|
||||
for op_key, ops in self.ops.items():
|
||||
for op in ops:
|
||||
if isinstance(op, Action):
|
||||
segments += op.get_all_segments()
|
||||
else:
|
||||
segments.append(op)
|
||||
|
||||
return segments
|
||||
|
||||
def get_result(self, embedding_module: Embedding):
|
||||
if isinstance(self.base_segment, Action):
|
||||
base_segment_result = self.base_segment.get_result(embedding_module)
|
||||
else:
|
||||
base_segment_result = self.base_segment.get_embeddings(embedding_module)
|
||||
|
||||
for op_key, ops in self.ops.items():
|
||||
for op in ops:
|
||||
if isinstance(op, Action):
|
||||
op_result = op.get_result(embedding_module)
|
||||
else:
|
||||
op_result = op.get_embeddings(embedding_module)
|
||||
|
||||
|
||||
if op_result.shape[1] > base_segment_result.shape[1]:
|
||||
print('[WARN] ArithAction: op_result.shape[1] > base_segment_result.shape[1] - averaging op_result')
|
||||
op_result = torch.mean(op_result, dim=1, keepdim=True)
|
||||
|
||||
if op_key == "+":
|
||||
base_segment_result.add(op_result)
|
||||
elif op_key == "-":
|
||||
base_segment_result.subtract(op_result)
|
||||
|
||||
return base_segment_result
|
||||
|
||||
@classmethod
|
||||
def parse_segment(
|
||||
cls,
|
||||
tokens: list[str],
|
||||
start_chars: list[str],
|
||||
end_chars: list[str],
|
||||
parent_parser: Callable[[list[str], SD1Tokenizer], Union[PromptSegment, 'Action']],
|
||||
tokenizer: SD1Tokenizer,
|
||||
) -> Action:
|
||||
"""
|
||||
Parse an arithmetic action from a list of tokens
|
||||
Supported formats:
|
||||
<base_segment:+op1-op2-op3>
|
||||
|
||||
:param tokens: List of tokens, will be modified
|
||||
:param start_chars: List of start chars for all actions
|
||||
:param end_chars: List of end chars for all actions
|
||||
:param parent_parser: Function to parse segments to allow for nested actions
|
||||
:return:
|
||||
"""
|
||||
token = tokens.pop(0)
|
||||
assert token == cls.START_CHAR, "ArithAction must start with " + cls.START_CHAR + " but got " + token
|
||||
|
||||
# Parse base segment
|
||||
base_segment = parent_parser(tokens, tokenizer)
|
||||
|
||||
token = tokens.pop(0)
|
||||
assert token == ":", "ArithAction must have a ':' after the base segment" + " but got " + token
|
||||
|
||||
# Parse ops string
|
||||
ops = {'+': [], '-': []}
|
||||
while tokens[0] != cls.END_CHAR:
|
||||
op_char = tokens.pop(0)
|
||||
assert op_char in ["+", "-"], "ArithAction must have a '+' or '-' as an op char but got " + op_char
|
||||
ops[op_char].append(parent_parser(tokens, tokenizer))
|
||||
|
||||
token = tokens.pop(0)
|
||||
assert token == cls.END_CHAR, "ArithAction must end with " + cls.END_CHAR + " but got " + token
|
||||
|
||||
return cls(base_segment, ops)
|
||||
@@ -0,0 +1,100 @@
|
||||
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 ")"'
|
||||
name = "avg"
|
||||
chars = ["+", "+"]
|
||||
|
||||
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
|
||||
@@ -1,95 +0,0 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Callable, Union
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
|
||||
from comfy.sd1_clip import SD1Tokenizer
|
||||
|
||||
|
||||
class Action(ABC):
|
||||
@property
|
||||
@abstractmethod
|
||||
def START_CHAR(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def END_CHAR(self):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def token_length(self):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_all_segments(self):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_result(self, embedding_module: Embedding):
|
||||
pass
|
||||
|
||||
@classmethod
|
||||
@abstractmethod
|
||||
def parse_segment(
|
||||
cls,
|
||||
tokens: list[str],
|
||||
start_chars: list[str],
|
||||
end_chars: list[str],
|
||||
parent_parser: Callable[[list[str], SD1Tokenizer], Union[str, 'Action']],
|
||||
tokenizer: SD1Tokenizer,
|
||||
) -> 'Action':
|
||||
pass
|
||||
|
||||
def depth_repr(self, depth=1):
|
||||
raise NotImplementedError()
|
||||
|
||||
class PromptSegment:
|
||||
def __init__(self, text: str, tokens: list[Union[int, Tensor]]):
|
||||
self.text = text
|
||||
self.tokens = tokens
|
||||
|
||||
def token_length(self):
|
||||
return len(self.tokens)
|
||||
|
||||
def get_embeddings(self, embedding_module: Embedding):
|
||||
tensors = torch.LongTensor(self.tokens).to(torch.device('cpu'))
|
||||
unsqueezed_tensors = tensors.unsqueeze(0)
|
||||
return embedding_module(unsqueezed_tensors)
|
||||
|
||||
def depth_repr(self, depth=1):
|
||||
out = f'"{self.text}"('
|
||||
|
||||
cleaned_tokens = list(map(lambda x: str(x) if isinstance(x, int) else "EMBD", self.tokens))
|
||||
out += ", ".join(cleaned_tokens)
|
||||
|
||||
out += ")"
|
||||
return out
|
||||
|
||||
def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment:
|
||||
split_text = text.split(" ")
|
||||
tokens = []
|
||||
for word in split_text:
|
||||
if word.startswith(tokenizer.embedding_identifier) and tokenizer.embedding_directory is not None:
|
||||
embedding_name = word[len(tokenizer.embedding_identifier):].strip('\n')
|
||||
|
||||
get_embed_ret = tokenizer._try_get_embedding(embedding_name)
|
||||
embedding = get_embed_ret[0]
|
||||
leftover = get_embed_ret[1]
|
||||
if embedding is None:
|
||||
print(f"warning, embedding:{embedding_name} does not exist, ignoring")
|
||||
else:
|
||||
if len(embedding.shape) == 1:
|
||||
tokens.append(embedding)
|
||||
else:
|
||||
tokens.extend(embedding)
|
||||
|
||||
if leftover != "":
|
||||
word = leftover
|
||||
else:
|
||||
continue
|
||||
tokens.extend(tokenizer.tokenizer(word)["input_ids"][1:-1])
|
||||
|
||||
return PromptSegment(text, tokens)
|
||||
@@ -0,0 +1,77 @@
|
||||
from typing import Union, List
|
||||
|
||||
import torch
|
||||
from torch.nn import 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 = ["-", "-"]
|
||||
|
||||
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_arg)
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
|
||||
# Calculate the embeddings for the base segment
|
||||
all_base_embeddings = [
|
||||
get_embedding(seg_or_action, embedding_module)
|
||||
for seg_or_action in self.base_arg
|
||||
]
|
||||
|
||||
result = torch.cat(all_base_embeddings, dim=1)
|
||||
|
||||
for arg in self.additional_args:
|
||||
all_arg_embeddings = [
|
||||
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
|
||||
]
|
||||
|
||||
arg_embedding = torch.cat(all_arg_embeddings, dim=1)
|
||||
|
||||
if (
|
||||
arg_embedding.shape[-2] == 1
|
||||
or result.shape[-2] == arg_embedding.shape[-2]
|
||||
):
|
||||
result = result.sub(arg_embedding)
|
||||
else:
|
||||
print(
|
||||
"WARNING: shape mismatch when trying to apply sum, arg will be averaged"
|
||||
)
|
||||
result = result.sub(torch.mean(arg_embedding, dim=1, keepdim=True))
|
||||
|
||||
return result
|
||||
|
||||
# def __repr__(self):
|
||||
# return f"sum(\n\tbase_segment={self.base_segment},\n\tadditional_args={self.additional_args}\n)"
|
||||
def __repr__(self) -> str:
|
||||
return f"sum({', '.join(map(str, self.additional_args))})"
|
||||
|
||||
def depth_repr(self, depth=1):
|
||||
out = "NudgeAction(\n"
|
||||
if isinstance(self.base_arg, Action):
|
||||
base_segment_repr = self.base_arg.depth_repr(depth + 1)
|
||||
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
|
||||
else:
|
||||
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
|
||||
|
||||
if isinstance(self.additional_args, Action):
|
||||
target_repr = self.additional_args.depth_repr(depth + 1)
|
||||
out += "\t" * depth + f"target={target_repr},\n"
|
||||
else:
|
||||
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
|
||||
out += "\t" * depth + f"weight={self.weight},\n"
|
||||
out += "\t" * (depth - 1) + ")"
|
||||
return out
|
||||
|
||||
@@ -1,17 +0,0 @@
|
||||
from custom_nodes.ClipStuff.lib.actions import ALL_START_CHARS, ALL_END_CHARS
|
||||
from custom_nodes.ClipStuff.lib.actions.base import Action
|
||||
|
||||
|
||||
def is_action_segment(action_class: Action.__class__, segment: str):
|
||||
if not issubclass(action_class, Action):
|
||||
raise Exception(
|
||||
f"action_class must be a subclass of Action, got {action_class}"
|
||||
)
|
||||
|
||||
return (
|
||||
segment[0] == action_class.START_CHAR and segment[-1] == action_class.END_CHAR
|
||||
)
|
||||
|
||||
|
||||
def is_any_action_segment(segment: str):
|
||||
return segment[0] in ALL_START_CHARS and segment[-1] in ALL_END_CHARS
|
||||
@@ -0,0 +1,62 @@
|
||||
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+ ")"'
|
||||
name = "mult"
|
||||
chars = ["[", "]"]
|
||||
|
||||
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
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
import torch
|
||||
from torch.nn import Embedding
|
||||
|
||||
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 = ["[", "]"]
|
||||
|
||||
def token_length(self) -> int:
|
||||
"""
|
||||
Neg negates the embeddings of the base segment, so the length is the length of the base segment
|
||||
:return:
|
||||
"""
|
||||
total_length = 0
|
||||
for seg_or_action in self.arg:
|
||||
total_length += seg_or_action.token_length()
|
||||
|
||||
return total_length
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
|
||||
all_embeddings = []
|
||||
for seg_or_action in self.arg:
|
||||
if isinstance(seg_or_action, Action):
|
||||
all_embeddings.append(seg_or_action.get_result(embedding_module))
|
||||
else:
|
||||
all_embeddings.append(seg_or_action.get_embeddings(embedding_module))
|
||||
|
||||
target_embeddings = torch.cat(all_embeddings, dim=1)
|
||||
return target_embeddings * -1
|
||||
|
||||
@@ -0,0 +1,38 @@
|
||||
import torch
|
||||
from torch.nn import Embedding
|
||||
|
||||
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
|
||||
|
||||
def token_length(self) -> int:
|
||||
"""
|
||||
Norm normalizes the embeddings of the base segment, so the length is the length of the base segment
|
||||
:return:
|
||||
"""
|
||||
total_length = 0
|
||||
for seg_or_action in self.arg:
|
||||
total_length += seg_or_action.token_length()
|
||||
|
||||
return total_length
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
|
||||
all_embeddings = []
|
||||
for seg_or_action in self.arg:
|
||||
if isinstance(seg_or_action, Action):
|
||||
target_embeddings = seg_or_action.get_result(embedding_module)
|
||||
else:
|
||||
target_embeddings = seg_or_action.get_embeddings(embedding_module)
|
||||
all_embeddings.append(target_embeddings)
|
||||
|
||||
target_embeddings = torch.cat(all_embeddings, dim=1)
|
||||
return torch.div(target_embeddings, torch.norm(target_embeddings, dim=-1, keepdim=True))
|
||||
|
||||
@@ -1,128 +0,0 @@
|
||||
from typing import Optional, Union, Callable
|
||||
|
||||
import torch
|
||||
from torch.nn import Embedding
|
||||
|
||||
from comfy.sd1_clip import SD1Tokenizer
|
||||
from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment
|
||||
|
||||
|
||||
class NudgeAction(Action):
|
||||
def token_length(self):
|
||||
# Nudge nudges the embeddings of the base segment, so the length is the length of the base segment
|
||||
if isinstance(self.base_segment, Action):
|
||||
return self.base_segment.token_length()
|
||||
|
||||
return len(self.base_segment.tokens)
|
||||
|
||||
def get_all_segments(self):
|
||||
segments = []
|
||||
if isinstance(self.base_segment, Action):
|
||||
segments += self.base_segment.get_all_segments()
|
||||
else:
|
||||
segments.append(self.base_segment)
|
||||
|
||||
if isinstance(self.target, Action):
|
||||
segments += self.target.get_all_segments()
|
||||
else:
|
||||
segments.append(self.target)
|
||||
|
||||
return segments
|
||||
|
||||
def get_result(self, embedding_module: Embedding):
|
||||
if isinstance(self.base_segment, Action):
|
||||
base_segment_result = self.base_segment.get_result(embedding_module)
|
||||
else:
|
||||
base_segment_result = self.base_segment.get_embeddings(embedding_module)
|
||||
|
||||
if isinstance(self.target, Action):
|
||||
target_segment_result = self.target.get_result(embedding_module)
|
||||
else:
|
||||
target_segment_result = self.target.get_embeddings(embedding_module)
|
||||
|
||||
base_mean = torch.mean(base_segment_result, dim=1, keepdim=True)
|
||||
if target_segment_result.shape[1] == 1:
|
||||
translation_vector = target_segment_result - base_mean
|
||||
else:
|
||||
translation_vector = torch.mean(target_segment_result, dim=1, keepdim=True) - base_mean
|
||||
|
||||
return base_segment_result.add(translation_vector, alpha=self.weight)
|
||||
|
||||
START_CHAR = "["
|
||||
END_CHAR = "]"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_segment: PromptSegment | Action,
|
||||
target: Union[PromptSegment, Action],
|
||||
weight: Optional[float] = None,
|
||||
):
|
||||
self.base_segment = base_segment
|
||||
self.weight = weight
|
||||
self.target = target
|
||||
|
||||
def __repr__(self):
|
||||
return f"NudgeAction(\n\tbase_segment={self.base_segment},\n\ttarget={self.target},\n\tweight={self.weight}\n)"
|
||||
|
||||
def depth_repr(self, depth=1):
|
||||
out = "NudgeAction(\n"
|
||||
if isinstance(self.base_segment, Action):
|
||||
base_segment_repr = self.base_segment.depth_repr(depth + 1)
|
||||
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
|
||||
else:
|
||||
out += "\t" * depth + f'base_segment={self.base_segment.depth_repr()},\n'
|
||||
|
||||
if isinstance(self.target, Action):
|
||||
target_repr = self.target.depth_repr(depth + 1)
|
||||
out += "\t" * depth + f"target={target_repr},\n"
|
||||
else:
|
||||
out += "\t" * depth + f"target={self.target.depth_repr()},\n"
|
||||
out += "\t" * depth + f"weight={self.weight},\n"
|
||||
out += "\t" * (depth - 1) + ")"
|
||||
return out
|
||||
|
||||
@classmethod
|
||||
def parse_segment(
|
||||
cls,
|
||||
tokens: list[str],
|
||||
start_chars: list[str],
|
||||
end_chars: list[str],
|
||||
parent_parser: Callable[[list[str], SD1Tokenizer], PromptSegment | Action],
|
||||
tokenizer: SD1Tokenizer,
|
||||
) -> Action:
|
||||
"""
|
||||
Parse a nudge action from a list of tokens
|
||||
Supported formats:
|
||||
[base_segment:target_segment]
|
||||
[base_segment:target_segment:weight]
|
||||
|
||||
Weight is optional, if not provided it will be None
|
||||
:param tokens: List of tokens, will be modified
|
||||
:param start_chars: List of start chars for all actions
|
||||
:param end_chars: List of end chars for all actions
|
||||
:param parent_parser: Function to parse segments to allow for nested actions
|
||||
:return:
|
||||
"""
|
||||
token = tokens.pop(0)
|
||||
assert token == cls.START_CHAR, "NudgeAction must start with " + cls.START_CHAR + " got " + token
|
||||
|
||||
# Parse base segment
|
||||
base_segment = parent_parser(tokens, tokenizer)
|
||||
|
||||
token = tokens.pop(0)
|
||||
assert token == ":", "NudgeAction must have a ':' after the base segment" + " but got " + token
|
||||
|
||||
# Parse target segment
|
||||
target_segment = parent_parser(tokens, tokenizer)
|
||||
|
||||
# Parse weight if it exists
|
||||
weight = None
|
||||
if tokens[0] == ":":
|
||||
# Parse weight
|
||||
tokens.pop(0)
|
||||
weight = float(tokens.pop(0))
|
||||
|
||||
token = tokens.pop(0)
|
||||
assert token == cls.END_CHAR, "NudgeAction must end with " + cls.END_CHAR + " got " + token
|
||||
|
||||
return cls(base_segment, target_segment, weight)
|
||||
@@ -0,0 +1,87 @@
|
||||
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
|
||||
|
||||
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
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
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)* ")"'
|
||||
name = "scaleDims"
|
||||
description = "Scales the specified dimensions of the input embeddings by the specified amount"
|
||||
example = "'The scaleDims(cat|4,1.5|76,1.2) is happy' scales the 4th dimension by 1.5 and the 76th dimension by 1.2 for the word 'cat'"
|
||||
chars = ["-", "-"]
|
||||
|
||||
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
|
||||
@@ -0,0 +1,72 @@
|
||||
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)* ")"'
|
||||
name = "setDims"
|
||||
description = "Sets the specified dimensions of the input embeddings to the specified value"
|
||||
example = "'The scaleDims(cat|4,1.5|76,1.2) is happy' scales the 4th dimension by 1.5 and the 76th dimension by 1.2 for the word 'cat'"
|
||||
chars = ["-", "-"]
|
||||
|
||||
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
|
||||
@@ -0,0 +1,100 @@
|
||||
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 ")"'
|
||||
name = "slerp"
|
||||
chars = ["+", "+"]
|
||||
|
||||
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
|
||||
@@ -0,0 +1,79 @@
|
||||
from typing import Union, List
|
||||
|
||||
import torch
|
||||
from torch.nn import Embedding
|
||||
|
||||
from custom_nodes.KepPromptLang.lib.actions.action_utils import get_embedding
|
||||
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||
from custom_nodes.KepPromptLang.lib.action.base import Action, MultiArgAction
|
||||
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
|
||||
|
||||
|
||||
class SumAction(MultiArgAction):
|
||||
grammar = 'sum(" arg ("|" arg)+ ")"'
|
||||
name = "sum"
|
||||
chars = ["+", "+"]
|
||||
|
||||
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_arg)
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
|
||||
# Calculate the embeddings for the base segment
|
||||
all_base_embeddings = [
|
||||
get_embedding(seg_or_action, embedding_module)
|
||||
for seg_or_action in self.base_arg
|
||||
]
|
||||
|
||||
result = torch.cat(all_base_embeddings, dim=1)
|
||||
|
||||
for arg in self.additional_args:
|
||||
all_arg_embeddings = [
|
||||
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
|
||||
]
|
||||
|
||||
arg_embedding = torch.cat(all_arg_embeddings, dim=1)
|
||||
|
||||
if (
|
||||
arg_embedding.shape[-2] == 1
|
||||
or result.shape[-2] == arg_embedding.shape[-2]
|
||||
):
|
||||
result = result.add(arg_embedding)
|
||||
else:
|
||||
print(
|
||||
"WARNING: shape mismatch when trying to apply sum, arg will be averaged"
|
||||
)
|
||||
result = result.add(
|
||||
torch.mean(arg_embedding, dim=1, keepdim=True)
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
# def __repr__(self):
|
||||
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
|
||||
def __repr__(self) -> str:
|
||||
return f"sum({', '.join(map(str, self.additional_args))})"
|
||||
|
||||
def depth_repr(self, depth=1):
|
||||
out = "NudgeAction(\n"
|
||||
if isinstance(self.base_arg, Action):
|
||||
base_segment_repr = self.base_arg.depth_repr(depth + 1)
|
||||
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
|
||||
else:
|
||||
out += "\t" * depth + f"base_segment={self.base_arg.depth_repr()},\n"
|
||||
|
||||
if isinstance(self.additional_args, Action):
|
||||
target_repr = self.additional_args.depth_repr(depth + 1)
|
||||
out += "\t" * depth + f"target={target_repr},\n"
|
||||
else:
|
||||
out += "\t" * depth + f"target={self.additional_args.depth_repr()},\n"
|
||||
out += "\t" * depth + f"weight={self.weight},\n"
|
||||
out += "\t" * (depth - 1) + ")"
|
||||
return out
|
||||
|
||||
@@ -0,0 +1,6 @@
|
||||
from typing import Union
|
||||
|
||||
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||
from custom_nodes.KepPromptLang.lib.action.base import Action
|
||||
|
||||
SegOrAction = Union[PromptSegment, Action]
|
||||
@@ -0,0 +1,38 @@
|
||||
from typing import List
|
||||
|
||||
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
|
||||
+38
-27
@@ -1,18 +1,20 @@
|
||||
import contextlib
|
||||
import os
|
||||
from typing import Union
|
||||
from typing import List
|
||||
|
||||
import torch
|
||||
from transformers import CLIPTextConfig, modeling_utils
|
||||
|
||||
from comfy import model_management
|
||||
import comfy.ops
|
||||
from comfy.sd import CLIP
|
||||
from custom_nodes.ClipStuff.lib.actions.base import PromptSegment, Action
|
||||
from custom_nodes.ClipStuff.lib.fun_clip_stuff import MyCLIPTextModel
|
||||
from custom_nodes.ClipStuff.lib.tokenizer import TokenDict
|
||||
from custom_nodes.KepPromptLang.lib.action.base import Action
|
||||
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
||||
from custom_nodes.KepPromptLang.lib.fun_clip_stuff import PromptLangTextModel
|
||||
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
class SD1FunClipModel(torch.nn.Module):
|
||||
|
||||
# Methods with no comment can be assumed to be the same as comfy.sd1_clip.SD1ClipModel
|
||||
class PromptLangClipModel(torch.nn.Module):
|
||||
"""Uses the CLIP transformer encoder for text (from huggingface)"""
|
||||
LAYERS = [
|
||||
"last",
|
||||
@@ -22,28 +24,37 @@ class SD1FunClipModel(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
|
||||
if textmodel_path is not None:
|
||||
self.transformer = MyCLIPTextModel.from_pretrained(textmodel_path)
|
||||
# Our transformer
|
||||
self.transformer = PromptLangTextModel.from_pretrained(textmodel_path)
|
||||
else:
|
||||
if textmodel_json_config is None:
|
||||
# TODO: Maybe re-use clip config?
|
||||
# Config could come from cond_stage_model.transformer.config
|
||||
# Copied clip_config
|
||||
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "clip_config.json")
|
||||
config = CLIPTextConfig.from_json_file(textmodel_json_config)
|
||||
self.num_layers = config.num_hidden_layers
|
||||
with comfy.ops.use_comfy_ops():
|
||||
with comfy.ops.use_comfy_ops(device, dtype):
|
||||
with modeling_utils.no_init_weights():
|
||||
self.transformer = MyCLIPTextModel(config)
|
||||
# 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
|
||||
@@ -68,7 +79,8 @@ class SD1FunClipModel(torch.nn.Module):
|
||||
self.layer = self.layer_default[0]
|
||||
self.layer_idx = self.layer_default[1]
|
||||
|
||||
def set_up_textual_embeddings(self, tokens: list[list[PromptSegment | Action]], current_embeds):
|
||||
# Completely changed to support Segments and actions
|
||||
def set_up_textual_embeddings(self, tokens: List[List[SegOrAction]], current_embeds):
|
||||
next_new_token = token_dict_size = current_embeds.weight.shape[0] - 1
|
||||
embedding_weights = []
|
||||
|
||||
@@ -131,7 +143,8 @@ class SD1FunClipModel(torch.nn.Module):
|
||||
if segment.tokens[tokenIdx] == -1:
|
||||
segment.tokens[tokenIdx] = n
|
||||
|
||||
def forward(self, tokens, **kwargs):
|
||||
# Support our set_up_textual_embeddings which modifies the input embeddings
|
||||
def forward(self, tokens):
|
||||
backup_embeds = self.transformer.get_input_embeddings()
|
||||
device = backup_embeds.weight.device
|
||||
self.set_up_textual_embeddings(tokens, backup_embeds)
|
||||
@@ -142,16 +155,8 @@ class SD1FunClipModel(torch.nn.Module):
|
||||
else:
|
||||
precision_scope = contextlib.nullcontext
|
||||
|
||||
|
||||
if (kwargs.get("position_ids", None) is not None):
|
||||
position_ids = torch.LongTensor(kwargs["position_ids"]).to(device)
|
||||
else:
|
||||
position_ids = None
|
||||
|
||||
|
||||
with precision_scope(model_management.get_autocast_device(device)):
|
||||
outputs = self.transformer(input_ids=tokens, output_hidden_states=self.layer == "hidden",
|
||||
position_ids=position_ids)
|
||||
outputs = self.transformer(input_ids=tokens, output_hidden_states=self.layer == "hidden")
|
||||
self.transformer.set_input_embeddings(backup_embeds)
|
||||
|
||||
if self.layer == "last":
|
||||
@@ -165,21 +170,27 @@ class SD1FunClipModel(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, **kwargs):
|
||||
return self(tokens, **kwargs)
|
||||
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)
|
||||
|
||||
def encode_token_weights(self, prompt_segments: list[list[Union[PromptSegment | Action]]], **kwargs):
|
||||
# Changed from comfy.sd1_clip.ClipTokenWeightEncoder
|
||||
# Changed to use PromptSegments
|
||||
def encode_token_weights(self, prompt_segments: List[List[SegOrAction]]):
|
||||
to_encode = [[PromptSegment(text="_Empty Batch_", tokens=self.empty_tokens[0])]]
|
||||
for batch in prompt_segments:
|
||||
to_encode.append(batch)
|
||||
|
||||
out, pooled = self.encode(to_encode, **kwargs)
|
||||
out, pooled = self.encode(to_encode)
|
||||
z_empty = out[0:1]
|
||||
if pooled.shape[0] > 1:
|
||||
first_pooled = pooled[1:2]
|
||||
|
||||
+32
-58
@@ -1,14 +1,17 @@
|
||||
from typing import Optional, Tuple, Union
|
||||
from typing import Optional, Tuple, Union, List
|
||||
|
||||
import torch
|
||||
from torch import device
|
||||
from transformers import CLIPTextConfig
|
||||
from transformers.modeling_outputs import BaseModelOutputWithPooling
|
||||
from transformers.models.clip.modeling_clip import _expand_mask, CLIPTextEmbeddings, CLIPTextTransformer, \
|
||||
CLIPTextModel
|
||||
from transformers.models.clip.modeling_clip import (
|
||||
_expand_mask,
|
||||
CLIPTextEmbeddings,
|
||||
CLIPTextTransformer,
|
||||
CLIPTextModel,
|
||||
)
|
||||
|
||||
from custom_nodes.ClipStuff.lib.actions.base import PromptSegment, Action
|
||||
from custom_nodes.ClipStuff.lib.tokenizer import TokenDict
|
||||
from custom_nodes.KepPromptLang.lib.action.base import Action
|
||||
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
||||
|
||||
def slerp(val, low, high):
|
||||
low = low.unsqueeze(0)
|
||||
@@ -20,18 +23,20 @@ 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 MyCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||
class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||
def __init__(self, config: CLIPTextConfig):
|
||||
super().__init__(config)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_dicts: Optional[list[list[tuple[TokenDict]]]] = 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 = []
|
||||
for batch_idx, batch in enumerate(input_dicts):
|
||||
@@ -48,29 +53,6 @@ class MyCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||
if position_ids is None:
|
||||
position_ids = self.position_ids[:, :seq_length]
|
||||
|
||||
# if inputs_embeds is None:
|
||||
# inputs_embeds = self.token_embedding(input_ids)
|
||||
|
||||
# for batch_idx, batch in enumerate(input_dicts):
|
||||
# for token_idx, token in enumerate(batch):
|
||||
# if token[0].nudge_id is not None:
|
||||
# nudged_embed = inputs_embeds[batch_idx, token_idx][:] + self.token_embedding(torch.LongTensor([token[0].nudge_id]).to(torch.device('cpu')))[0]
|
||||
# if token[0].nudge_index_start is not None and token[0].nudge_index_stop is not None:
|
||||
# nudge_start = token[0].nudge_index_start
|
||||
# nudge_end = token[0].nudge_index_stop
|
||||
# else:
|
||||
# nudge_start = 0
|
||||
# nudge_end = 768
|
||||
# inputs_embeds[batch_idx, token_idx][nudge_start:nudge_end] = (slerp(token[0].nudge_weight, inputs_embeds[batch_idx, token_idx][:], nudged_embed)[0][nudge_start:nudge_end])
|
||||
# elif token[0].arith_ops is not None:
|
||||
# for op, id_list in token[0].arith_ops.items():
|
||||
# if op == '+':
|
||||
# for this_id in id_list:
|
||||
# inputs_embeds[batch_idx, token_idx] += self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0]
|
||||
# elif op == '-':
|
||||
# for this_id in id_list:
|
||||
# inputs_embeds[batch_idx, token_idx] -= self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0]
|
||||
|
||||
embeds = []
|
||||
for batch in batches:
|
||||
if len(batch) == 1:
|
||||
@@ -84,14 +66,14 @@ class MyCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||
return embeddings
|
||||
|
||||
|
||||
class MyCLIPTextTransformer(CLIPTextTransformer):
|
||||
class PrompLangCLIPTextTransformer(CLIPTextTransformer):
|
||||
def __init__(self, config: CLIPTextConfig):
|
||||
super().__init__(config)
|
||||
self.embeddings = MyCLIPTextEmbeddings(config)
|
||||
self.embeddings = PromptLangCLIPTextEmbeddings(config)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[list[list[PromptSegment | Action]]] = 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,
|
||||
@@ -119,12 +101,20 @@ class MyCLIPTextTransformer(CLIPTextTransformer):
|
||||
bsz = len(input_ids)
|
||||
# TODO: Properly gather this
|
||||
seq_len = 77
|
||||
# bsz, seq_len = input_shape
|
||||
input_shape = torch.Size([bsz, seq_len])
|
||||
# CLIP's text model uses causal mask, prepare it here.
|
||||
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
|
||||
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(
|
||||
hidden_states.device
|
||||
)
|
||||
## VERSION DIFF ##
|
||||
# transformers < 4.30.0
|
||||
if hasattr(self, "_build_causal_attention_mask"):
|
||||
print("Using transformers < 4.30.0")
|
||||
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(hidden_states.device)
|
||||
else:
|
||||
# transformers >= 4.30.0
|
||||
print("Using transformers >= 4.30.0")
|
||||
from transformers.models.clip.modeling_clip import _make_causal_mask
|
||||
causal_attention_mask = _make_causal_mask(input_shape, hidden_states.dtype, device=hidden_states.device)
|
||||
|
||||
# expand attention_mask
|
||||
if attention_mask is not None:
|
||||
# [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
|
||||
@@ -174,37 +164,21 @@ class MyCLIPTextTransformer(CLIPTextTransformer):
|
||||
)
|
||||
|
||||
|
||||
class MyCLIPTextModel(CLIPTextModel):
|
||||
# This is necessary to pass the PromptLangCLIPTextTransformer
|
||||
class PromptLangTextModel(CLIPTextModel):
|
||||
def __init__(self, config: CLIPTextConfig):
|
||||
super().__init__(config)
|
||||
self.text_model = MyCLIPTextTransformer(config)
|
||||
self.text_model = PrompLangCLIPTextTransformer(config)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[list[list[tuple[TokenDict]]]] = 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,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
return_dict: Optional[bool] = None,
|
||||
) -> Union[Tuple, BaseModelOutputWithPooling]:
|
||||
r"""
|
||||
Returns:
|
||||
|
||||
Examples:
|
||||
|
||||
```python
|
||||
>>> from transformers import AutoTokenizer, CLIPTextModel
|
||||
|
||||
>>> model = CLIPTextModel.from_pretrained("openai/clip-vit-base-patch32")
|
||||
>>> tokenizer = AutoTokenizer.from_pretrained("openai/clip-vit-base-patch32")
|
||||
|
||||
>>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="pt")
|
||||
|
||||
>>> outputs = model(**inputs)
|
||||
>>> last_hidden_state = outputs.last_hidden_state
|
||||
>>> pooled_output = outputs.pooler_output # pooled (EOS token) states
|
||||
```"""
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
|
||||
return self.text_model(
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
from lark import Lark
|
||||
|
||||
from .grammar import grammar
|
||||
|
||||
PromptParser = Lark(grammar, start="start", parser="earley")
|
||||
@@ -0,0 +1,20 @@
|
||||
grammar = """
|
||||
?start: item+
|
||||
|
||||
item: embedding
|
||||
| WORD
|
||||
| generic_function
|
||||
| QUOTED_STRING
|
||||
|
||||
generic_function: FUNC_NAME "(" arg ("|" arg)* ")"
|
||||
|
||||
arg: item+
|
||||
|
||||
embedding: "embedding:" WORD
|
||||
FUNC_NAME: /[A-Za-z_-]+/
|
||||
WORD: /[A-Za-z0-9,_\.-]+/
|
||||
QUOTED_STRING: /"([^"\\\]*(\\\.[^"\\\]*)*)"|'([^'\\\]*(\\\.[^'\\\]*)*)'/
|
||||
|
||||
%import common.WS
|
||||
%ignore WS
|
||||
"""
|
||||
@@ -0,0 +1,31 @@
|
||||
from typing import Union, List
|
||||
|
||||
import torch
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
|
||||
|
||||
class PromptSegment:
|
||||
def __init__(self, text: str, tokens: List[Union[int, Tensor]]):
|
||||
self.text = text
|
||||
self.tokens = tokens
|
||||
|
||||
def __repr__(self):
|
||||
return f'"{self.text}"{self.tokens}'
|
||||
|
||||
def token_length(self):
|
||||
return len(self.tokens)
|
||||
|
||||
def get_embeddings(self, embedding_module: Embedding) -> Tensor:
|
||||
tensors = torch.LongTensor(self.tokens).to(embedding_module.weight.device)
|
||||
unsqueezed_tensors = tensors.unsqueeze(0)
|
||||
return embedding_module(unsqueezed_tensors)
|
||||
|
||||
def depth_repr(self, depth=1):
|
||||
out = f'"{self.text}"('
|
||||
|
||||
cleaned_tokens = list(map(lambda x: str(x) if isinstance(x, int) else "EMBD", self.tokens))
|
||||
out += ", ".join(cleaned_tokens)
|
||||
|
||||
out += ")"
|
||||
return out
|
||||
@@ -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.name in action_registry:
|
||||
raise ValueError(f"Action {action.name} already registered")
|
||||
action_registry[str(action.name)] = action
|
||||
|
||||
def get_action_by_name(name: str) -> Type[Action]:
|
||||
if name not in action_registry:
|
||||
raise ValueError(f"Action {name} not found in registry")
|
||||
|
||||
return action_registry[name]
|
||||
@@ -0,0 +1,63 @@
|
||||
from typing import List
|
||||
|
||||
from lark import Transformer, Token
|
||||
|
||||
from comfy.sd1_clip import SD1Tokenizer
|
||||
from custom_nodes.KepPromptLang.lib.action.base import Action, ActionArity
|
||||
from custom_nodes.KepPromptLang.lib.actions.diff import DiffAction
|
||||
from custom_nodes.KepPromptLang.lib.actions.rand import RandAction
|
||||
from custom_nodes.KepPromptLang.lib.parser.registration import get_action_by_name
|
||||
from custom_nodes.KepPromptLang.lib.parser.utils import build_prompt_segment
|
||||
from custom_nodes.KepPromptLang.lib.actions.neg import NegAction
|
||||
from custom_nodes.KepPromptLang.lib.actions.norm import NormAction
|
||||
from custom_nodes.KepPromptLang.lib.actions.sum import SumAction
|
||||
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
|
||||
class PromptTransformer(Transformer):
|
||||
# def WORD(self, items):
|
||||
# return items
|
||||
|
||||
def __init__(self, tokenizer: SD1Tokenizer):
|
||||
super().__init__()
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def item(self, items: List[Token]):
|
||||
for item in items:
|
||||
if isinstance(item, Action):
|
||||
return item
|
||||
|
||||
if isinstance(item, PromptSegment):
|
||||
return item
|
||||
|
||||
if item.type == "WORD":
|
||||
return build_prompt_segment(str(item), self.tokenizer)
|
||||
elif item.type == "QUOTED_STRING":
|
||||
# Remove the quotes
|
||||
unquoted = item[1:-1]
|
||||
# Replace escaped quotes with quotes
|
||||
unescaped = unquoted.replace("\\\"", "\"").replace("\\\'", "\'")
|
||||
return build_prompt_segment(unescaped, self.tokenizer)
|
||||
elif item.type == "embedding":
|
||||
return build_prompt_segment(item, self.tokenizer)
|
||||
elif item.type == "function":
|
||||
return item
|
||||
else:
|
||||
raise Exception("Unknown item type: " + str(item.type))
|
||||
|
||||
def arg(self, items):
|
||||
return items
|
||||
|
||||
def embedding(self, items):
|
||||
return build_prompt_segment(f'{self.tokenizer.embedding_identifier}{items[0]}', self.tokenizer)
|
||||
|
||||
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}")
|
||||
@@ -0,0 +1,38 @@
|
||||
from lark import Token
|
||||
|
||||
from comfy.sd1_clip import SD1Tokenizer
|
||||
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
|
||||
def flatten_tree(tree):
|
||||
if isinstance(tree, Token):
|
||||
return [str(tree)]
|
||||
else:
|
||||
return [str(tree.data)] + sum([flatten_tree(child) for child in tree.children], [])
|
||||
|
||||
|
||||
def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment:
|
||||
split_text = text.split(" ")
|
||||
tokens = []
|
||||
for word in split_text:
|
||||
if word.startswith(tokenizer.embedding_identifier) and tokenizer.embedding_directory is not None:
|
||||
embedding_name = word[len(tokenizer.embedding_identifier):].strip('\n')
|
||||
|
||||
get_embed_ret = tokenizer._try_get_embedding(embedding_name)
|
||||
embedding = get_embed_ret[0]
|
||||
leftover = get_embed_ret[1]
|
||||
if embedding is None:
|
||||
print(f"warning, embedding:{embedding_name} does not exist, ignoring")
|
||||
else:
|
||||
if len(embedding.shape) == 1:
|
||||
tokens.append(embedding)
|
||||
else:
|
||||
tokens.extend(embedding)
|
||||
|
||||
if leftover != "":
|
||||
word = leftover
|
||||
else:
|
||||
continue
|
||||
tokens.extend(tokenizer.tokenizer(word)["input_ids"][1:-1])
|
||||
|
||||
return PromptSegment(text, tokens)
|
||||
+19
-129
@@ -1,116 +1,15 @@
|
||||
import re
|
||||
from typing import Union
|
||||
from typing import List
|
||||
|
||||
from lark import Tree
|
||||
|
||||
from comfy.sd1_clip import SD1Tokenizer
|
||||
from custom_nodes.ClipStuff.lib.actions import (
|
||||
NudgeAction,
|
||||
ArithAction,
|
||||
ALL_START_CHARS,
|
||||
ALL_END_CHARS,
|
||||
ALL_ACTIONS,
|
||||
)
|
||||
from custom_nodes.ClipStuff.lib.actions.base import (
|
||||
Action,
|
||||
PromptSegment,
|
||||
build_prompt_segment,
|
||||
)
|
||||
from custom_nodes.ClipStuff.lib.actions.lib import (
|
||||
is_any_action_segment,
|
||||
is_action_segment,
|
||||
)
|
||||
from custom_nodes.ClipStuff.lib.actions.utils import batch_size_info
|
||||
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
||||
|
||||
arith_action = r'(<[a-zA-Z0-9\-_]+:[a-zA-Z0-9\-_]+>)'
|
||||
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
|
||||
|
||||
# TODO: Get embedding identifier from tokenizer
|
||||
tokenizer_regex = re.compile(
|
||||
fr"""
|
||||
\d+\.\d+ # Capture decimals
|
||||
|
|
||||
(?:(?!embedding:)[\w\s]|embedding:[a-zA-Z0-9_]+)+ # Capture sequences of characters, including "embedding:"
|
||||
|
|
||||
\d+ # Capture whole numbers
|
||||
|
|
||||
[:+-{re.escape("".join(ALL_START_CHARS))}{re.escape("".join(ALL_END_CHARS))}] # Capture special characters including start and end characters
|
||||
""",
|
||||
re.VERBOSE
|
||||
)
|
||||
def tokenize(text: str) -> list[str]:
|
||||
# Captures:
|
||||
# 1. Words
|
||||
# 2. Numbers(1.0, 1)
|
||||
# 3. Special characters(ALL_START_CHARS, ALL_END_CHARS, :, +, -)
|
||||
tokens = re.findall(tokenizer_regex, text)
|
||||
print(tokens)
|
||||
return [token.strip() for token in tokens]
|
||||
|
||||
|
||||
|
||||
def parse_segment(tokens: list[str], tokenizer: SD1Tokenizer) -> PromptSegment | Action:
|
||||
print("Parse segment: Checking token: " + tokens[0])
|
||||
for action in ALL_ACTIONS:
|
||||
if tokens[0] == action.START_CHAR:
|
||||
return action.parse_segment(tokens, ALL_START_CHARS, ALL_END_CHARS, parse_segment, tokenizer)
|
||||
# If we get here, it's a text segment
|
||||
return build_prompt_segment(tokens.pop(0), tokenizer)
|
||||
|
||||
def parse(tokens: list[str], tokenizer: SD1Tokenizer) -> list[PromptSegment | Action]:
|
||||
parsed = []
|
||||
while tokens:
|
||||
if tokens[0] == '':
|
||||
tokens.pop(0)
|
||||
continue
|
||||
print("Parse: Checking token: " + tokens[0])
|
||||
if tokens[0] in ALL_START_CHARS:
|
||||
parsed.append(parse_segment(tokens, tokenizer))
|
||||
else:
|
||||
parsed.append(build_prompt_segment(tokens.pop(0), tokenizer))
|
||||
return parsed
|
||||
|
||||
|
||||
def parse_special_tokens(string) -> list[str]:
|
||||
out = []
|
||||
current = ""
|
||||
|
||||
for char in string:
|
||||
if char in ALL_START_CHARS:
|
||||
out += [current]
|
||||
current = char
|
||||
elif char in ALL_END_CHARS:
|
||||
out += [current + char]
|
||||
current = ""
|
||||
else:
|
||||
current += char
|
||||
out += [current]
|
||||
return out
|
||||
|
||||
|
||||
def parse_segment_actions(string, tokenizer: SD1Tokenizer) -> list[PromptSegment | NudgeAction | ArithAction]:
|
||||
tokens = tokenize(string)
|
||||
parsed = parse(tokens, tokenizer)
|
||||
return parsed
|
||||
|
||||
class TokenDict:
|
||||
def __init__(self,
|
||||
token_id: int,
|
||||
weight: float = None,
|
||||
nudge_id=None, nudge_weight=None, nudge_start: int = None, nudge_end: int = None,
|
||||
arith_ops: dict[str, list[str]] = None):
|
||||
if weight is None:
|
||||
self.weight = 1.0
|
||||
else:
|
||||
self.weight = weight
|
||||
|
||||
self.token_id = token_id
|
||||
self.nudge_id = nudge_id
|
||||
self.nudge_weight = nudge_weight
|
||||
self.nudge_index_start = nudge_start
|
||||
self.nudge_index_stop = nudge_end
|
||||
|
||||
self.arith_ops = arith_ops
|
||||
|
||||
|
||||
class MyTokenizer(SD1Tokenizer):
|
||||
class PromptLangTokenizer(SD1Tokenizer):
|
||||
def __init__(self, tokenizer_path=None, max_length=77, pad_with_end=True, embedding_directory=None, embedding_size=768, embedding_key='clip_l', special_tokens=None):
|
||||
super().__init__(tokenizer_path, max_length, pad_with_end, embedding_directory, embedding_size, embedding_key)
|
||||
|
||||
@@ -119,34 +18,25 @@ class MyTokenizer(SD1Tokenizer):
|
||||
Returns batches of segments and actions
|
||||
:return: List of list(batches) of segments and actions
|
||||
"""
|
||||
def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs) -> list[list[PromptSegment | Action]]:
|
||||
def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs) -> List[List[SegOrAction]]:
|
||||
if self.pad_with_end:
|
||||
pad_token = self.end_token
|
||||
else:
|
||||
pad_token = 0
|
||||
|
||||
parsed_actions = parse_segment_actions(text, self)
|
||||
|
||||
# nudge_start = kwargs.get("nudge_start")
|
||||
# nudge_end = kwargs.get("nudge_end")
|
||||
#
|
||||
# if nudge_start is not None and nudge_end is not None:
|
||||
# nudge_start = int(nudge_start)
|
||||
# nudge_end = int(nudge_end)
|
||||
#
|
||||
# # tokenize words
|
||||
for segment in parsed_actions:
|
||||
if isinstance(segment, Action):
|
||||
print(segment.depth_repr())
|
||||
else:
|
||||
print(segment.depth_repr())
|
||||
parsed_prompt = PromptParser.parse(text)
|
||||
parsed_actions = PromptTransformer(self).transform(parsed_prompt)
|
||||
|
||||
# reshape token array to CLIP input size
|
||||
batched_segments = []
|
||||
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token])]
|
||||
# batched_segments.append(batch)
|
||||
batch_size = 1
|
||||
for segment in parsed_actions:
|
||||
if isinstance(parsed_actions, Tree):
|
||||
segments_to_process = parsed_actions.children
|
||||
else:
|
||||
segments_to_process = [parsed_actions]
|
||||
for segment in segments_to_process:
|
||||
num_tokens = segment.token_length()
|
||||
# determine if we're going to try and keep the tokens in a single batch
|
||||
is_large = num_tokens >= self.max_word_length
|
||||
@@ -163,7 +53,7 @@ class MyTokenizer(SD1Tokenizer):
|
||||
batch_size = num_tokens + 1 # +1 for start token
|
||||
continue
|
||||
|
||||
# If the segment is small enough to fit in the current batch, add it
|
||||
# Since the segment fits in the current batch, add it
|
||||
batch.append(segment)
|
||||
batch_size += num_tokens
|
||||
|
||||
@@ -172,7 +62,7 @@ class MyTokenizer(SD1Tokenizer):
|
||||
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length))
|
||||
batched_segments.append(batch)
|
||||
|
||||
for batch in batched_segments:
|
||||
batch_size_info(batch)
|
||||
# for batch in batched_segments:
|
||||
# batch_size_info(batch)
|
||||
|
||||
return batched_segments
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
import random
|
||||
import os
|
||||
from typing import List, Tuple, Any
|
||||
|
||||
import numpy as np
|
||||
from PIL import Image
|
||||
@@ -6,18 +8,37 @@ from PIL import Image
|
||||
import folder_paths
|
||||
import comfy.sd
|
||||
import comfy.ops
|
||||
from custom_nodes.ClipStuff.lib.clip_model import SD1FunClipModel
|
||||
from custom_nodes.KepPromptLang.lib.clip_model import PromptLangClipModel
|
||||
|
||||
from custom_nodes.ClipStuff.lib.tokenizer import MyTokenizer
|
||||
from custom_nodes.KepPromptLang.lib.tokenizer import PromptLangTokenizer
|
||||
|
||||
|
||||
class EmptyClass:
|
||||
pass
|
||||
|
||||
|
||||
class SpecialClipLoader:
|
||||
class MonacoPrompt:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"clip": ("CLIP",),
|
||||
"prompt": ("MONACO",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
FUNCTION = "do_crap"
|
||||
OUTPUT_IS_LIST = (False,)
|
||||
CATEGORY = "conditioning"
|
||||
|
||||
@staticmethod
|
||||
def do_crap(clip, prompt):
|
||||
return (clip,)
|
||||
|
||||
class SpecialClipLoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls): # type: ignore
|
||||
return {
|
||||
"required": {
|
||||
"source_clip": ("CLIP",),
|
||||
@@ -30,13 +51,12 @@ class SpecialClipLoader:
|
||||
CATEGORY = "conditioning"
|
||||
|
||||
@staticmethod
|
||||
def load_clip(source_clip):
|
||||
def load_clip(source_clip: comfy.sd.CLIP) -> Tuple[comfy.sd.CLIP]:
|
||||
clip_target = EmptyClass()
|
||||
clip_target.params = {}
|
||||
clip_target.clip = SD1FunClipModel
|
||||
clip_target.tokenizer = MyTokenizer
|
||||
clip_target.clip = PromptLangClipModel
|
||||
clip_target.tokenizer = PromptLangTokenizer
|
||||
|
||||
# TODO: Extract embedding directory from source_clip
|
||||
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.embedding_directory)
|
||||
comfy.sd.load_clip_weights(
|
||||
clip.cond_stage_model, source_clip.cond_stage_model.state_dict()
|
||||
@@ -44,65 +64,23 @@ class SpecialClipLoader:
|
||||
return (clip,)
|
||||
|
||||
|
||||
class KepAdvTextEncode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"multiline": True}),
|
||||
"clip": ("CLIP",),
|
||||
"nudge_start": ("INT", {}),
|
||||
"nudge_end": ("INT", {}),
|
||||
"split_newlines": ("BOOL", {"default": True}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
FUNCTION = "encode"
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
CATEGORY = "conditioning"
|
||||
|
||||
@staticmethod
|
||||
def encode(clip, text, nudge_start, nudge_end, split_newlines):
|
||||
ret = []
|
||||
if split_newlines:
|
||||
prompts = text.split("\n")
|
||||
else:
|
||||
prompts = [text]
|
||||
|
||||
for prompt in prompts:
|
||||
if prompt.strip() == "":
|
||||
continue
|
||||
tokens = clip.tokenizer.tokenize_with_weights(
|
||||
text,
|
||||
return_word_ids=False,
|
||||
nudge_start=nudge_start,
|
||||
nudge_end=nudge_end,
|
||||
)
|
||||
cond, pooled = clip.encode_from_tokens(
|
||||
tokens, return_pooled=True, position_ids=[0] * 77
|
||||
)
|
||||
cond = [[cond, {"pooled_output": pooled}]]
|
||||
ret.append(cond)
|
||||
return (ret,)
|
||||
|
||||
|
||||
def tensor2img(tensor_img):
|
||||
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"},
|
||||
@@ -111,46 +89,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])
|
||||
@@ -165,8 +150,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",
|
||||
@@ -176,15 +162,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",
|
||||
@@ -194,7 +189,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 } }
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
lark
|
||||
@@ -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.")
|
||||
@@ -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()
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
+12963
File diff suppressed because one or more lines are too long
Generated
+2047
File diff suppressed because it is too large
Load Diff
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Vendored
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user