43 changed files with 688 additions and 3240 deletions
-10
View File
@@ -1,10 +0,0 @@
name: Test
on:
workflow_dispatch:
jobs:
build:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v2
- name: Setup upterm session
uses: lhotari/action-upterm@v1
-85
View File
@@ -1,85 +0,0 @@
name: Run Test Workflow
on: [workflow_dispatch]
jobs:
Test:
runs-on: ubuntu-latest
strategy:
fail-fast: false
matrix:
python-version: [ "3.7", "3.8", "3.9", "3.10", "3.11" ]
steps:
- name: Clone Upstream
uses: actions/checkout@v3
with:
repository: comfyanonymous/ComfyUI
ref: master
fetch-depth: 0
- name: Clone Node
uses: actions/checkout@v3
with:
ref: master
fetch-depth: 0
path: custom_nodes/KepPromptLang
- name: Setup Python
uses: actions/setup-python@v4
with:
# Version range or exact version of Python or PyPy to use, using SemVer's version range syntax. Reads from .python-version if unset.
python-version: ${{ matrix.python-version }}
- name: Cache virtualenv
uses: actions/cache@v3
id: cache-venv
with:
path: ./.venv/
key: ${{ runner.os }}-venv-${{ matrix.python-version }}-${{ hashFiles('**/requirements.txt') }}
restore-keys: |
${{ runner.os }}-venv-${{ matrix.python-version }}-
- name: Install Requirements
if: steps.cache-venv.outputs.cache-hit != 'true'
run: |
python -m venv ./.venv
source ./.venv/bin/activate
pip install torch --index-url https://download.pytorch.org/whl/cpu
pip install -r requirements.txt
pip install -r custom_nodes/KepPromptLang/requirements.txt
pip install huggingface_hub websocket-client
# - name: Cache SD Checkpoint
# uses: actions/cache@v3
# with:
# path: |
# models/checkpoints
# key: ${{ runner.os }}-sd-15-checkpoint
- name: Check and Download Model
run: |
source ./.venv/bin/activate
python custom_nodes/KepPromptLang/test_files/check_and_download_model.py
- name: Run in Background
env:
PYTHONUNBUFFERED: 1
run: |
source ./.venv/bin/activate
python main.py --cpu &> server.log &
sleep 10
# - name: Setup upterm session
# uses: lhotari/action-upterm@v1
- name: Run Workflow
run: |
source ./.venv/bin/activate
python custom_nodes/KepPromptLang/test_files/run_workflow.py
- name: Upload Comfy Server Log
if: always()
uses: actions/upload-artifact@v3
with:
name: comfy-server-log-${{ matrix.python-version }}
path: server.log
-80
View File
@@ -1,81 +1 @@
# 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
![Example Photo](assets/first_example.png)
## 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)
```
```
+2
View File
@@ -1,9 +1,11 @@
from .nodes import (
KepAdvTextEncode,
BuildGif,
SpecialClipLoader,
)
NODE_CLASS_MAPPINGS = {
"Kep Adv Text Encode": KepAdvTextEncode,
"Build Gif": BuildGif,
"Special CLIP Loader": SpecialClipLoader,
}
Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.8 MiB

File diff suppressed because it is too large Load Diff
-24
View File
@@ -1,24 +0,0 @@
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.pos_scale import PosScaleAction
from custom_nodes.KepPromptLang.lib.actions.rand import RandAction
from custom_nodes.KepPromptLang.lib.actions.scale_dims import ScaleDims
from custom_nodes.KepPromptLang.lib.actions.set_dims import SetDims
from custom_nodes.KepPromptLang.lib.actions.slerp import SlerpAction
from custom_nodes.KepPromptLang.lib.actions.sum import SumAction
from custom_nodes.KepPromptLang.lib.parser.registration import register_action
register_action(DiffAction)
register_action(MultiplyAction)
register_action(NegAction)
register_action(NormAction)
register_action(RandAction)
register_action(SumAction)
register_action(SlerpAction)
register_action(AverageAction)
register_action(ScaleDims)
register_action(SetDims)
register_action(PosScaleAction)
View File
-124
View File
@@ -1,124 +0,0 @@
from abc import ABC, abstractmethod
from enum import Enum
from typing import Union, List, TypedDict, Tuple
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 PostModifiers(TypedDict):
"""
A dictionary of post modifiers for an action result.
"""
position_embed_scale: Union[float, None]
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) -> Union[Tensor, Tuple[Tensor, PostModifiers]]:
"""
Get the result of this action. This is called when the embeddings are being calculated.
:param embedding_module: The embedding module to use to get the base embeddings for tokens.
: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
+6
View File
@@ -0,0 +1,6 @@
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]
-11
View File
@@ -1,11 +0,0 @@
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)
+131
View File
@@ -0,0 +1,131 @@
from typing import Callable, Union
import torch
from torch.nn import Embedding
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment
class ArithAction(Action):
START_CHAR = "<"
END_CHAR = ">"
def __init__(self, base_segment: PromptSegment | Action, ops: dict[str, list[PromptSegment | Action]]):
self.base_segment = base_segment
self.ops = ops
def __repr__(self):
return f"ArithAction(\n\tbase_segment={self.base_segment},\n\tops={self.ops}\n)"
def depth_repr(self, depth=1):
out = "ArithAction(\n"
if isinstance(self.base_segment, Action):
base_segment_repr = self.base_segment.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
elif isinstance(self.base_segment, PromptSegment):
out += "\t" * depth + f'base_segment={self.base_segment.depth_repr(depth)}'
else:
out += "\t" * depth + f'base_segment="{self.base_segment}",'
for op_key, ops in self.ops.items():
for op in ops:
out += "\n" + "\t" * depth + f'"{op_key}":[\n'
if isinstance(op, Action):
op_repr = op.depth_repr(depth + 2)
out += "\t" * (depth + 1) + f"{op_repr}\n"
else:
out += "\t" * (depth + 1) + f'{op.depth_repr()},\n'
out += "\t" * depth + "],"
out += "\n" + "\t" * (depth - 1) + ")"
return out
def token_length(self):
# ArithAction modifies the embeddings of the base segment, so the length is the length of the base segment
if isinstance(self.base_segment, Action):
return self.base_segment.token_length()
return len(self.base_segment.tokens)
def get_all_segments(self):
segments = []
if isinstance(self.base_segment, Action):
segments += self.base_segment.get_all_segments()
else:
segments.append(self.base_segment)
for op_key, ops in self.ops.items():
for op in ops:
if isinstance(op, Action):
segments += op.get_all_segments()
else:
segments.append(op)
return segments
def get_result(self, embedding_module: Embedding):
if isinstance(self.base_segment, Action):
base_segment_result = self.base_segment.get_result(embedding_module)
else:
base_segment_result = self.base_segment.get_embeddings(embedding_module)
for op_key, ops in self.ops.items():
for op in ops:
if isinstance(op, Action):
op_result = op.get_result(embedding_module)
else:
op_result = op.get_embeddings(embedding_module)
if op_result.shape[1] > base_segment_result.shape[1]:
print('[WARN] ArithAction: op_result.shape[1] > base_segment_result.shape[1] - averaging op_result')
op_result = torch.mean(op_result, dim=1, keepdim=True)
if op_key == "+":
base_segment_result.add(op_result)
elif op_key == "-":
base_segment_result.subtract(op_result)
return base_segment_result
@classmethod
def parse_segment(
cls,
tokens: list[str],
start_chars: list[str],
end_chars: list[str],
parent_parser: Callable[[list[str], SD1Tokenizer], Union[PromptSegment, 'Action']],
tokenizer: SD1Tokenizer,
) -> Action:
"""
Parse an arithmetic action from a list of tokens
Supported formats:
<base_segment:+op1-op2-op3>
:param tokens: List of tokens, will be modified
:param start_chars: List of start chars for all actions
:param end_chars: List of end chars for all actions
:param parent_parser: Function to parse segments to allow for nested actions
:return:
"""
token = tokens.pop(0)
assert token == cls.START_CHAR, "ArithAction must start with " + cls.START_CHAR + " but got " + token
# Parse base segment
base_segment = parent_parser(tokens, tokenizer)
token = tokens.pop(0)
assert token == ":", "ArithAction must have a ':' after the base segment" + " but got " + token
# Parse ops string
ops = {'+': [], '-': []}
while tokens[0] != cls.END_CHAR:
op_char = tokens.pop(0)
assert op_char in ["+", "-"], "ArithAction must have a '+' or '-' as an op char but got " + op_char
ops[op_char].append(parent_parser(tokens, tokenizer))
token = tokens.pop(0)
assert token == cls.END_CHAR, "ArithAction must end with " + cls.END_CHAR + " but got " + token
return cls(base_segment, ops)
-100
View File
@@ -1,100 +0,0 @@
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
+95
View File
@@ -0,0 +1,95 @@
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)
-77
View File
@@ -1,77 +0,0 @@
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
+17
View File
@@ -0,0 +1,17 @@
from custom_nodes.ClipStuff.lib.actions import ALL_START_CHARS, ALL_END_CHARS
from custom_nodes.ClipStuff.lib.actions.base import Action
def is_action_segment(action_class: Action.__class__, segment: str):
if not issubclass(action_class, Action):
raise Exception(
f"action_class must be a subclass of Action, got {action_class}"
)
return (
segment[0] == action_class.START_CHAR and segment[-1] == action_class.END_CHAR
)
def is_any_action_segment(segment: str):
return segment[0] in ALL_START_CHARS and segment[-1] in ALL_END_CHARS
-62
View File
@@ -1,62 +0,0 @@
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
-34
View File
@@ -1,34 +0,0 @@
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
-38
View File
@@ -1,38 +0,0 @@
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))
+128
View File
@@ -0,0 +1,128 @@
from typing import Optional, Union, Callable
import torch
from torch.nn import Embedding
from comfy.sd1_clip import SD1Tokenizer
from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment
class NudgeAction(Action):
def token_length(self):
# Nudge nudges the embeddings of the base segment, so the length is the length of the base segment
if isinstance(self.base_segment, Action):
return self.base_segment.token_length()
return len(self.base_segment.tokens)
def get_all_segments(self):
segments = []
if isinstance(self.base_segment, Action):
segments += self.base_segment.get_all_segments()
else:
segments.append(self.base_segment)
if isinstance(self.target, Action):
segments += self.target.get_all_segments()
else:
segments.append(self.target)
return segments
def get_result(self, embedding_module: Embedding):
if isinstance(self.base_segment, Action):
base_segment_result = self.base_segment.get_result(embedding_module)
else:
base_segment_result = self.base_segment.get_embeddings(embedding_module)
if isinstance(self.target, Action):
target_segment_result = self.target.get_result(embedding_module)
else:
target_segment_result = self.target.get_embeddings(embedding_module)
base_mean = torch.mean(base_segment_result, dim=1, keepdim=True)
if target_segment_result.shape[1] == 1:
translation_vector = target_segment_result - base_mean
else:
translation_vector = torch.mean(target_segment_result, dim=1, keepdim=True) - base_mean
return base_segment_result.add(translation_vector, alpha=self.weight)
START_CHAR = "["
END_CHAR = "]"
def __init__(
self,
base_segment: PromptSegment | Action,
target: Union[PromptSegment, Action],
weight: Optional[float] = None,
):
self.base_segment = base_segment
self.weight = weight
self.target = target
def __repr__(self):
return f"NudgeAction(\n\tbase_segment={self.base_segment},\n\ttarget={self.target},\n\tweight={self.weight}\n)"
def depth_repr(self, depth=1):
out = "NudgeAction(\n"
if isinstance(self.base_segment, Action):
base_segment_repr = self.base_segment.depth_repr(depth + 1)
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
else:
out += "\t" * depth + f'base_segment={self.base_segment.depth_repr()},\n'
if isinstance(self.target, Action):
target_repr = self.target.depth_repr(depth + 1)
out += "\t" * depth + f"target={target_repr},\n"
else:
out += "\t" * depth + f"target={self.target.depth_repr()},\n"
out += "\t" * depth + f"weight={self.weight},\n"
out += "\t" * (depth - 1) + ")"
return out
@classmethod
def parse_segment(
cls,
tokens: list[str],
start_chars: list[str],
end_chars: list[str],
parent_parser: Callable[[list[str], SD1Tokenizer], PromptSegment | Action],
tokenizer: SD1Tokenizer,
) -> Action:
"""
Parse a nudge action from a list of tokens
Supported formats:
[base_segment:target_segment]
[base_segment:target_segment:weight]
Weight is optional, if not provided it will be None
:param tokens: List of tokens, will be modified
:param start_chars: List of start chars for all actions
:param end_chars: List of end chars for all actions
:param parent_parser: Function to parse segments to allow for nested actions
:return:
"""
token = tokens.pop(0)
assert token == cls.START_CHAR, "NudgeAction must start with " + cls.START_CHAR + " got " + token
# Parse base segment
base_segment = parent_parser(tokens, tokenizer)
token = tokens.pop(0)
assert token == ":", "NudgeAction must have a ':' after the base segment" + " but got " + token
# Parse target segment
target_segment = parent_parser(tokens, tokenizer)
# Parse weight if it exists
weight = None
if tokens[0] == ":":
# Parse weight
tokens.pop(0)
weight = float(tokens.pop(0))
token = tokens.pop(0)
assert token == cls.END_CHAR, "NudgeAction must end with " + cls.END_CHAR + " got " + token
return cls(base_segment, target_segment, weight)
-65
View File
@@ -1,65 +0,0 @@
from typing import Tuple, List
import torch
from torch.nn import Embedding
from custom_nodes.KepPromptLang.lib.action.base import (
Action,
PostModifiers,
MultiArgAction,
)
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
class PosScaleAction(MultiArgAction):
grammar = 'posScale(" arg+ ")"'
name = "posScale"
chars = ["[", "]"]
def __init__(self, args: List[List[SegOrAction]]) -> None:
super().__init__(args)
if len(args) != 2:
raise ValueError("PosScale action should have exactly two arguments")
self.target_arg = args[0]
self._parse_multiplier(args[1])
def _parse_multiplier(self, arg: List[SegOrAction]) -> None:
if len(arg) != 1:
raise ValueError(
"PosScale actions multiplier should have exactly one segment"
)
multiplier_seg_or_action = arg[0]
if isinstance(multiplier_seg_or_action, Action):
raise ValueError("PosScale actions multiplier must be a number")
try:
self.parsed_multiplier = float(multiplier_seg_or_action.text)
except ValueError:
raise ValueError(
"PosScale action should have an integer/float as the multiplier"
)
def token_length(self) -> int:
"""
PosScale modifies the posional embeddings of the base segment, so the length is the length of the base segment
:return:
"""
total_length = 0
for seg_or_action in self.target_arg:
total_length += seg_or_action.token_length()
return total_length
def get_result(self, embedding_module: Embedding) -> Tuple[torch.Tensor, PostModifiers]:
all_embeddings = []
for seg_or_action in self.target_arg:
if isinstance(seg_or_action, Action):
all_embeddings.append(seg_or_action.get_result(embedding_module))
else:
all_embeddings.append(seg_or_action.get_embeddings(embedding_module))
target_embeddings = torch.cat(all_embeddings, dim=1)
return target_embeddings, {"position_embed_scale": self.parsed_multiplier}
-87
View File
@@ -1,87 +0,0 @@
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
-72
View File
@@ -1,72 +0,0 @@
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
-72
View File
@@ -1,72 +0,0 @@
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
-100
View File
@@ -1,100 +0,0 @@
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
-79
View File
@@ -1,79 +0,0 @@
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
-6
View File
@@ -1,6 +0,0 @@
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]
-38
View File
@@ -1,38 +0,0 @@
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
-23
View File
@@ -1,23 +0,0 @@
{
"architectures": [
"CLIPTextModel"
],
"attention_dropout": 0.0,
"bos_token_id": 0,
"dropout": 0.0,
"eos_token_id": 2,
"hidden_act": "gelu",
"hidden_size": 1280,
"initializer_factor": 1.0,
"initializer_range": 0.02,
"intermediate_size": 5120,
"layer_norm_eps": 1e-05,
"max_position_embeddings": 77,
"model_type": "clip_text_model",
"num_attention_heads": 20,
"num_hidden_layers": 32,
"pad_token_id": 1,
"projection_dim": 1280,
"torch_dtype": "float32",
"vocab_size": 49408
}
+27 -70
View File
@@ -1,22 +1,18 @@
import contextlib
import os
from typing import List
from typing import Union
import torch
from transformers import CLIPTextConfig, modeling_utils
from comfy import model_management
import comfy.ops
from comfy.sd1_clip import SD1ClipModel
from comfy.sdxl_clip import SDXLClipModel
from custom_nodes.KepPromptLang.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.KepPromptLang.lib.fun_clip_stuff import PromptLangTextModel
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
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
# Methods with no comment can be assumed to be the same as comfy.sd1_clip.SD1ClipModel
class PromptLangSDClipModel(torch.nn.Module):
class SD1FunClipModel(torch.nn.Module):
"""Uses the CLIP transformer encoder for text (from huggingface)"""
LAYERS = [
"last",
@@ -26,37 +22,28 @@ class PromptLangSDClipModel(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, dtype=None): # clip-vit-base-patch32
textmodel_path=None): # clip-vit-base-patch32
super().__init__()
assert layer in self.LAYERS
self.num_layers = 12
if textmodel_path is not None:
# Our transformer
self.transformer = PromptLangTextModel.from_pretrained(textmodel_path)
self.transformer = MyCLIPTextModel.from_pretrained(textmodel_path)
else:
if textmodel_json_config is None:
# TODO: Maybe re-use clip config?
# Config could come from cond_stage_model.transformer.config
# Copied clip_config
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "clip_config.json")
config = CLIPTextConfig.from_json_file(textmodel_json_config)
self.num_layers = config.num_hidden_layers
with comfy.ops.use_comfy_ops(device, dtype):
with comfy.ops.use_comfy_ops():
with modeling_utils.no_init_weights():
# Our transformer
self.transformer = PromptLangTextModel(config)
self.transformer = MyCLIPTextModel(config)
if dtype is not None:
self.transformer.to(dtype)
self.max_length = max_length
if freeze:
self.freeze()
self.layer = layer
self.layer_idx = None
self.empty_tokens = [[49406] + [49407] * 76]
self.text_projection = torch.nn.Parameter(torch.eye(self.transformer.get_input_embeddings().weight.shape[1]))
self.logit_scale = torch.nn.Parameter(torch.tensor(4.6055))
self.text_projection = None
self.layer_norm_hidden_state = True
if layer == "hidden":
assert layer_idx is not None
@@ -81,8 +68,7 @@ class PromptLangSDClipModel(torch.nn.Module):
self.layer = self.layer_default[0]
self.layer_idx = self.layer_default[1]
# Completely changed to support Segments and actions
def set_up_textual_embeddings(self, tokens: List[List[SegOrAction]], current_embeds):
def set_up_textual_embeddings(self, tokens: list[list[PromptSegment | Action]], current_embeds):
next_new_token = token_dict_size = current_embeds.weight.shape[0] - 1
embedding_weights = []
@@ -145,8 +131,7 @@ class PromptLangSDClipModel(torch.nn.Module):
if segment.tokens[tokenIdx] == -1:
segment.tokens[tokenIdx] = n
# Support our set_up_textual_embeddings which modifies the input embeddings
def forward(self, tokens):
def forward(self, tokens, **kwargs):
backup_embeds = self.transformer.get_input_embeddings()
device = backup_embeds.weight.device
self.set_up_textual_embeddings(tokens, backup_embeds)
@@ -157,8 +142,16 @@ class PromptLangSDClipModel(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")
outputs = self.transformer(input_ids=tokens, output_hidden_states=self.layer == "hidden",
position_ids=position_ids)
self.transformer.set_input_embeddings(backup_embeds)
if self.layer == "last":
@@ -172,27 +165,21 @@ class PromptLangSDClipModel(torch.nn.Module):
pooled_output = outputs.pooler_output
if self.text_projection is not None:
pooled_output = pooled_output.float().to(self.text_projection.device) @ self.text_projection.float()
pooled_output = pooled_output.to(self.text_projection.device) @ self.text_projection
return z.float(), pooled_output.float()
def encode(self, tokens):
return self(tokens)
def encode(self, tokens, **kwargs):
return self(tokens, **kwargs)
def load_sd(self, sd):
if "text_projection" in sd:
self.text_projection[:] = sd.pop("text_projection")
if "text_projection.weight" in sd:
self.text_projection[:] = sd.pop("text_projection.weight").transpose(0, 1)
return self.transformer.load_state_dict(sd, strict=False)
# Changed from comfy.sd1_clip.ClipTokenWeightEncoder
# Changed to use PromptSegments
def encode_token_weights(self, prompt_segments: List[List[SegOrAction]]):
def encode_token_weights(self, prompt_segments: list[list[Union[PromptSegment | Action]]], **kwargs):
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)
out, pooled = self.encode(to_encode, **kwargs)
z_empty = out[0:1]
if pooled.shape[0] > 1:
first_pooled = pooled[1:2]
@@ -211,33 +198,3 @@ class PromptLangSDClipModel(torch.nn.Module):
if (len(output) == 0):
return z_empty.cpu(), first_pooled.cpu()
return torch.cat(output, dim=-2).cpu(), first_pooled.cpu()
class PromptLangSD1ClipModel(SD1ClipModel):
def __init__(self, device="cpu", dtype=None, clip_name="l", clip_model=PromptLangSDClipModel):
super().__init__()
self.clip_name = clip_name
self.clip = "clip_{}".format(self.clip_name)
setattr(self, self.clip, clip_model(device=device, dtype=dtype))
class PromptLangSDXLClipModel(SDXLClipModel):
def __init__(self, device="cpu", dtype=None) -> None:
# Skip SDXLClipModel's init
super(SDXLClipModel, self).__init__()
self.clip_l = PromptLangSDClipModel(layer="hidden", layer_idx=11, device=device, dtype=dtype)
self.clip_l.layer_norm_hidden_state = False
self.clip_g = PromptLangSDXLClipG(device, dtype)
class PromptLangSDXLClipG(PromptLangSDClipModel):
def __init__(self, device="cpu", max_length=77, freeze=True, layer="penultimate", layer_idx=None, textmodel_path=None, dtype=None):
if layer == "penultimate":
layer="hidden"
layer_idx=-2
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "clip_config_bigg.json")
super().__init__(device=device, freeze=freeze, layer=layer, layer_idx=layer_idx, textmodel_json_config=textmodel_json_config, textmodel_path=textmodel_path, dtype=dtype)
self.empty_tokens = [[49406] + [49407] + [0] * 75]
self.layer_norm_hidden_state = False
def load_sd(self, sd):
return super().load_sd(sd)
+68 -99
View File
@@ -1,18 +1,14 @@
from typing import Optional, Tuple, Union, List, TypedDict
from importlib_metadata import version as import_version
from packaging import version
from typing import Optional, Tuple, Union
import torch
from torch import device
from transformers import CLIPTextConfig
from transformers.modeling_outputs import BaseModelOutputWithPooling
from transformers.models.clip.modeling_clip import (
CLIPTextEmbeddings,
CLIPTextTransformer,
CLIPTextModel,
)
from transformers.models.clip.modeling_clip import _expand_mask, CLIPTextEmbeddings, CLIPTextTransformer, \
CLIPTextModel
from custom_nodes.KepPromptLang.lib.action.base import Action
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
from custom_nodes.ClipStuff.lib.actions.base import PromptSegment, Action
from custom_nodes.ClipStuff.lib.tokenizer import TokenDict
def slerp(val, low, high):
low = low.unsqueeze(0)
@@ -24,63 +20,57 @@ def slerp(val, low, high):
res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high
return res
class PosModifier(TypedDict):
"""
A dictionary of post modifiers for an action result.
"""
position_embed_scale: Union[float]
start_idx: Union[int]
end_idx: Union[int]
class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
class MyCLIPTextEmbeddings(CLIPTextEmbeddings):
def __init__(self, config: CLIPTextConfig):
super().__init__(config)
def forward(
self,
input_dicts: Optional[List[List[SegOrAction]]] = None,
input_dicts: Optional[list[list[tuple[TokenDict]]]] = None,
input_ids: Optional[torch.LongTensor] = None,
position_ids: Optional[torch.LongTensor] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
) -> torch.Tensor:
if input_dicts is None:
raise ValueError("You have to specify input_dicts")
batches = []
pos_modifiers: List[List[PosModifier]] = []
for batch_idx, batch in enumerate(input_dicts):
results = []
batch_pos_modifiers = []
token_idx = 0
for seg_or_action in batch:
if isinstance(seg_or_action, Action):
action_result = seg_or_action.get_result(self.token_embedding)
if isinstance(action_result, tuple):
result, post_modifiers = action_result
if post_modifiers["position_embed_scale"] is not None:
post_modifiers["start_idx"] = token_idx
post_modifiers["end_idx"] = (
token_idx + seg_or_action.token_length()
)
batch_pos_modifiers.append(post_modifiers)
else:
result = action_result
results.append(seg_or_action.get_result(self.token_embedding))
else:
result = seg_or_action.get_embeddings(self.token_embedding)
results.append(result)
token_idx += seg_or_action.token_length()
results.append(seg_or_action.get_embeddings(self.token_embedding))
batches.append(results)
pos_modifiers.append(batch_pos_modifiers)
seq_length = batches[0][0].shape[-2]
if position_ids is None:
position_ids = self.position_ids[:, :seq_length]
# if inputs_embeds is None:
# inputs_embeds = self.token_embedding(input_ids)
# for batch_idx, batch in enumerate(input_dicts):
# for token_idx, token in enumerate(batch):
# if token[0].nudge_id is not None:
# nudged_embed = inputs_embeds[batch_idx, token_idx][:] + self.token_embedding(torch.LongTensor([token[0].nudge_id]).to(torch.device('cpu')))[0]
# if token[0].nudge_index_start is not None and token[0].nudge_index_stop is not None:
# nudge_start = token[0].nudge_index_start
# nudge_end = token[0].nudge_index_stop
# else:
# nudge_start = 0
# nudge_end = 768
# inputs_embeds[batch_idx, token_idx][nudge_start:nudge_end] = (slerp(token[0].nudge_weight, inputs_embeds[batch_idx, token_idx][:], nudged_embed)[0][nudge_start:nudge_end])
# elif token[0].arith_ops is not None:
# for op, id_list in token[0].arith_ops.items():
# if op == '+':
# for this_id in id_list:
# inputs_embeds[batch_idx, token_idx] += self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0]
# elif op == '-':
# for this_id in id_list:
# inputs_embeds[batch_idx, token_idx] -= self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0]
embeds = []
for batch in batches:
if len(batch) == 1:
@@ -88,62 +78,20 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
else:
embeds.append(torch.cat(batch, dim=-2))
for idx, batch_pos_modifiers in enumerate(pos_modifiers):
position_embeddings = self.position_embedding(position_ids)
if len(batch_pos_modifiers) > 0:
print(f"Found {len(batch_pos_modifiers)} pos modifiers for batch {idx}")
for post_modifier in batch_pos_modifiers:
position_embeddings[
0, post_modifier["start_idx"] : post_modifier["end_idx"]
] *= post_modifier["position_embed_scale"]
embeds[idx] = embeds[idx] + position_embeddings
embeddings = torch.cat(embeds, dim=0)
position_embeddings = self.position_embedding(position_ids)
embeddings = torch.cat(embeds, dim=0) + position_embeddings
return embeddings
class PrompLangCLIPTextTransformer(CLIPTextTransformer):
class MyCLIPTextTransformer(CLIPTextTransformer):
def __init__(self, config: CLIPTextConfig):
super().__init__(config)
self.embeddings = PromptLangCLIPTextEmbeddings(config)
self.transformers_version = version.parse(import_version('transformers'))
def process_attention_mask(self, hidden_states, attention_mask, bsz, seq_len):
# Parse the transformer version
input_shape = torch.Size([bsz, seq_len])
v4_30 = version.parse('4.30.0')
v4_35 = version.parse('4.35')
if self.transformers_version < v4_30:
print("Using transformers < 4.30.0")
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(
hidden_states.device)
elif v4_30 <= self.transformers_version < v4_35:
print("Using transformers >= 4.30.0 and <= 4.34.*")
from transformers.models.clip.modeling_clip import _make_causal_mask
causal_attention_mask = _make_causal_mask(input_shape, hidden_states.dtype, device=hidden_states.device)
else:
print("Using transformers >= 4.35")
from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask
causal_attention_mask = _create_4d_causal_attention_mask(
input_shape, hidden_states.dtype, device=hidden_states.device
)
# Expand attention_mask if it exists
if attention_mask is not None:
# Import _expand_mask or _prepare_4d_attention_mask based on version
if self.transformers_version < v4_35:
from transformers.models.clip.modeling_clip import _expand_mask
attention_mask = _expand_mask(attention_mask, hidden_states.dtype)
else:
from transformers.modeling_attn_mask_utils import _prepare_4d_attention_mask
attention_mask = _prepare_4d_attention_mask(attention_mask, hidden_states.dtype)
return causal_attention_mask, attention_mask
self.embeddings = MyCLIPTextEmbeddings(config)
def forward(
self,
input_ids: Optional[List[List[SegOrAction]]] = None,
input_ids: Optional[list[list[PromptSegment | Action]]] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
@@ -171,8 +119,16 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
bsz = len(input_ids)
# TODO: Properly gather this
seq_len = 77
causal_attention_mask, attention_mask = self.process_attention_mask(hidden_states, attention_mask, bsz, seq_len)
# bsz, seq_len = input_shape
# CLIP's text model uses causal mask, prepare it here.
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(
hidden_states.device
)
# expand attention_mask
if attention_mask is not None:
# [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
attention_mask = _expand_mask(attention_mask, hidden_states.dtype)
encoder_outputs = self.encoder(
inputs_embeds=hidden_states,
@@ -197,9 +153,6 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
else:
if seg_or_action.text == '__PAD__':
break
# Is a segment, and isn't the pad segment
idx += seg_or_action.token_length()
eot_idx.append(idx)
# text_embeds.shape = [batch_size, sequence_length, transformer.width]
# take features from the eot embedding (eot_token is the highest number in each sequence)
@@ -221,21 +174,37 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
)
# This is necessary to pass the PromptLangCLIPTextTransformer
class PromptLangTextModel(CLIPTextModel):
class MyCLIPTextModel(CLIPTextModel):
def __init__(self, config: CLIPTextConfig):
super().__init__(config)
self.text_model = PrompLangCLIPTextTransformer(config)
self.text_model = MyCLIPTextTransformer(config)
def forward(
self,
input_ids: Optional[List[List[SegOrAction]]] = None,
input_ids: Optional[list[list[tuple[TokenDict]]]] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
output_hidden_states: Optional[bool] = None,
return_dict: Optional[bool] = None,
) -> Union[Tuple, BaseModelOutputWithPooling]:
r"""
Returns:
Examples:
```python
>>> from transformers import AutoTokenizer, CLIPTextModel
>>> model = CLIPTextModel.from_pretrained("openai/clip-vit-base-patch32")
>>> tokenizer = AutoTokenizer.from_pretrained("openai/clip-vit-base-patch32")
>>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="pt")
>>> outputs = model(**inputs)
>>> last_hidden_state = outputs.last_hidden_state
>>> pooled_output = outputs.pooler_output # pooled (EOS token) states
```"""
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
return self.text_model(
-5
View File
@@ -1,5 +0,0 @@
from lark import Lark
from .grammar import grammar
PromptParser = Lark(grammar, start="start", parser="earley")
-20
View File
@@ -1,20 +0,0 @@
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
"""
-31
View File
@@ -1,31 +0,0 @@
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
-20
View File
@@ -1,20 +0,0 @@
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]
-58
View File
@@ -1,58 +0,0 @@
from typing import List
from lark import Transformer, Token
from comfy.sd1_clip import SDTokenizer
from custom_nodes.KepPromptLang.lib.action.base import Action, ActionArity
from custom_nodes.KepPromptLang.lib.parser.registration import get_action_by_name
from custom_nodes.KepPromptLang.lib.parser.utils import build_prompt_segment
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
class PromptTransformer(Transformer):
# def WORD(self, items):
# return items
def __init__(self, tokenizer: SDTokenizer):
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}")
-38
View File
@@ -1,38 +0,0 @@
from lark import Token
from comfy.sd1_clip import SDTokenizer
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: SDTokenizer) -> 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)
+132 -42
View File
@@ -1,51 +1,161 @@
from typing import List, Dict
import re
from typing import Union
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 comfy.sd1_clip import SD1Tokenizer, SDTokenizer
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]
class PromptLangSDTokenizer(SDTokenizer):
def __init__(self, tokenizer_path=None, max_length=77, pad_with_end=True, embedding_directory=None, embedding_size=768, embedding_key='clip_l'):
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):
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)
"""
Doesn't actually tokenize...
Returns batches of segments and actions
:return: List of list(batches) of segments and actions
"""
def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs) -> List[List[SegOrAction]]:
def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs) -> list[list[PromptSegment | Action]]:
if self.pad_with_end:
pad_token = self.end_token
else:
pad_token = 0
parsed_prompt = PromptParser.parse(text)
parsed_actions = PromptTransformer(self).transform(parsed_prompt)
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())
# reshape token array to CLIP input size
batched_segments = []
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token])]
# batched_segments.append(batch)
batch_size = 1
if isinstance(parsed_actions, Tree):
segments_to_process = parsed_actions.children
else:
segments_to_process = [parsed_actions]
for segment in segments_to_process:
for segment in parsed_actions:
num_tokens = segment.token_length()
# determine if we're going to try and keep the tokens in a single batch
is_large = num_tokens >= self.max_word_length
# If the segment is too large to fit in a single batch, pad the current batch and start a new one
if num_tokens + batch_size > self.max_length - 1:
remaining_length = self.max_length - batch_size
remaining_length = self.max_length - batch_size - 1 # -1 for end token
# Pad batch
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * (remaining_length - 1))) # -1 for end token
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length - 1))
batched_segments.append(batch)
# start new batch
@@ -53,7 +163,7 @@ class PromptLangSDTokenizer(SDTokenizer):
batch_size = num_tokens + 1 # +1 for start token
continue
# Since the segment fits in the current batch, add it
# If the segment is small enough to fit in the current batch, add it
batch.append(segment)
batch_size += num_tokens
@@ -62,27 +172,7 @@ class PromptLangSDTokenizer(SDTokenizer):
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
class PromptLangSD1Tokenizer(SD1Tokenizer):
def __init__(self, embedding_directory=None, clip_name='l', tokenizer=PromptLangSDTokenizer) -> None:
super().__init__(embedding_directory, clip_name, tokenizer)
class PromptLangSDXLClipGTokenizer(PromptLangSDTokenizer):
def __init__(self, tokenizer_path=None, embedding_directory=None):
super().__init__(tokenizer_path, pad_with_end=False, embedding_directory=embedding_directory, embedding_size=1280, embedding_key='clip_g')
class PromptLangSDXLTokenizer(SD1Tokenizer):
def __init__(self, embedding_directory=None) -> None:
self.clip_l = PromptLangSDTokenizer(embedding_directory=embedding_directory)
self.clip_g = PromptLangSDXLClipGTokenizer(embedding_directory=embedding_directory)
def tokenize_with_weights(self, text:str, return_word_ids=False) -> Dict[str, List[List[SegOrAction]]]:
out = {}
out["g"] = self.clip_g.tokenize_with_weights(text, return_word_ids)
out["l"] = self.clip_l.tokenize_with_weights(text, return_word_ids)
return out
+82 -78
View File
@@ -1,5 +1,4 @@
import os
from typing import List, Tuple, Any
import random
import numpy as np
from PIL import Image
@@ -7,18 +6,9 @@ from PIL import Image
import folder_paths
import comfy.sd
import comfy.ops
from comfy.sd2_clip import SD2ClipModel
from comfy.sdxl_clip import SDXLClipModel
from comfy.supported_models_base import ClipTarget
from custom_nodes.KepPromptLang.lib.clip_model import (
PromptLangSDXLClipModel,
PromptLangSD1ClipModel,
)
from custom_nodes.ClipStuff.lib.clip_model import SD1FunClipModel
from custom_nodes.KepPromptLang.lib.tokenizer import (
PromptLangSDXLTokenizer,
PromptLangSD1Tokenizer,
)
from custom_nodes.ClipStuff.lib.tokenizer import MyTokenizer
class EmptyClass:
@@ -27,7 +17,7 @@ class EmptyClass:
class SpecialClipLoader:
@classmethod
def INPUT_TYPES(cls): # type: ignore
def INPUT_TYPES(s):
return {
"required": {
"source_clip": ("CLIP",),
@@ -40,43 +30,79 @@ class SpecialClipLoader:
CATEGORY = "conditioning"
@staticmethod
def load_clip(source_clip: comfy.sd.CLIP) -> Tuple[comfy.sd.CLIP]:
def load_clip(source_clip):
clip_target = EmptyClass()
clip_target.params = {}
clip_target.clip = SD1FunClipModel
clip_target.tokenizer = MyTokenizer
if isinstance(source_clip.cond_stage_model, SDXLClipModel):
clip_target = ClipTarget(PromptLangSDXLTokenizer, PromptLangSDXLClipModel)
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.clip_g.embedding_directory)
comfy.sd.load_clip_weights(clip.cond_stage_model.clip_g,source_clip.cond_stage_model.clip_g.state_dict())
comfy.sd.load_clip_weights(
clip.cond_stage_model.clip_l, source_clip.cond_stage_model.clip_l.state_dict()
)
elif isinstance(source_clip, SD2ClipModel):
raise ValueError("SD2 Clip model is not supported.")
else:
clip_target = ClipTarget(PromptLangSD1Tokenizer, PromptLangSD1ClipModel)
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.clip_l.embedding_directory)
comfy.sd.load_clip_weights(
clip.cond_stage_model, source_clip.cond_stage_model.state_dict()
)
# 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()
)
return (clip,)
def tensor2img(tensor_img) -> Image.Image:
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):
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) -> None:
self.output_dir = folder_paths.get_output_directory()
def __init__(self):
pass
@classmethod
def INPUT_TYPES(cls): # type: ignore
def INPUT_TYPES(cls):
return {
"required": {
"images": ("IMAGE",),
"split_every": ("INT", {"default": -1}),
"frame_duration": ("INT", {"default": 125}),
"output_mode": (
["One Per Split", "Big Grid"],
{"default": "Big Grid"},
@@ -85,53 +111,46 @@ class BuildGif:
}
RELOAD_INST = True
RETURN_TYPES = ()
# RETURN_NAMES = ("Gifs",)
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("Gifs",)
INPUT_IS_LIST = True
FUNCTION = "build_gif"
# OUTPUT_IS_LIST = (True,)
OUTPUT_NODE = True
OUTPUT_IS_LIST = (True,)
# OUTPUT_NODE = False
CATEGORY = "List Stuff"
def build_gif(self, images: List[Any], split_every: List[int], frame_duration: List[int], output_mode: List[str]):
@staticmethod
def build_gif(images: list, split_every: list[int], output_mode: str):
print("Build GIF called!")
print(f"{type(images)}")
if len(split_every) > 1:
raise Exception("List input for split every is not supported.")
if len(output_mode) > 1:
raise Exception("List input for output_mode is not supported.")
output_mode = output_mode[0]
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]
split_every = split_every[0]
batch_size = images[0].size()[0]
if split_every_val == -1:
if split_every == -1:
split_chunks = 1
split_every_val = len(images)
split_every = len(images)
else:
split_chunks = int(len(images) / split_every_val)
split_chunks = int(len(images) / split_every)
out = []
num_wide = batch_size
num_tall = split_chunks
chunked_batches = [
images[split_every_val * chunk_idx : split_every_val * (chunk_idx + 1)]
images[split_every * chunk_idx : split_every * (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_val):
for idx_in_chunk in range(split_every):
img_shape = images[0][0].shape
img_frame = Image.new(
"RGB", size=(num_wide * img_shape[0], num_tall * img_shape[1])
@@ -146,9 +165,8 @@ class BuildGif:
)
frames.append(img_frame)
file = f"{filename}_{counter:05}_"
save_path = (
f"{os.path.join(full_output_folder, file)}"
f"{folder_paths.get_output_directory()}/{random.randint(1, 100)}"
)
frames[0].save(
f"{save_path}.webp",
@@ -158,24 +176,15 @@ class BuildGif:
save_all=True,
append_images=frames[1:],
optimize=False,
duration=frame_duration,
duration=125,
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_val * split_idx
split_end = split_every_val * (split_idx + 1)
split_start = split_every * split_idx
split_end = split_every * (split_idx + 1)
for batch_idx in range(batch_size):
file = f"{filename}_{counter:05}_"
save_path = (
f"{os.path.join(full_output_folder, file)}"
)
counter += 1
save_path = f"{folder_paths.get_output_directory()}/-{batch_idx}-{random.randint(1, 100)}"
print(save_path)
tensor2img(images[split_start][batch_idx]).save(
f"{save_path}.webp",
@@ -185,12 +194,7 @@ class BuildGif:
for nested_batch in images[split_start + 1 : split_end]
],
optimize=False,
duration=frame_duration,
duration=125,
loop=0,
)
results.append({
"filename": f"{file}.webp",
"subfolder": subfolder,
"type": "output"
})
return { "ui": { "images": results } }
return (out,)
-2
View File
@@ -1,2 +0,0 @@
lark
packaging
-12
View File
@@ -1,12 +0,0 @@
# If the models/checkpoints folder does not have test.txt, then download the model.
import os
from huggingface_hub import hf_hub_download
FILE = "v1-5-pruned-emaonly.safetensors"
REPO_ID = "runwayml/stable-diffusion-v1-5"
if not os.path.exists(f"models/checkpoints/{FILE}"):
print("Downloading model...")
hf_hub_download(repo_id=REPO_ID, filename=FILE, local_dir="models/checkpoints", local_dir_use_symlinks=False)
else:
print("Model already downloaded.")
-109
View File
@@ -1,109 +0,0 @@
#This is an example that uses the websockets api to know when a prompt execution is done
#Once the prompt execution is done it downloads the images using the /history endpoint
import os
import websocket #NOTE: websocket-client (https://github.com/websocket-client/websocket-client)
import uuid
import json
import urllib.request
import urllib.parse
server_address = "127.0.0.1:8188"
client_id = str(uuid.uuid4())
def queue_prompt(prompt):
p = {"prompt": prompt, "client_id": client_id}
data = json.dumps(p).encode('utf-8')
req = urllib.request.Request("http://{}/prompt".format(server_address), data=data)
try:
response = urllib.request.urlopen(req)
return json.loads(response.read())
except urllib.error.HTTPError as e:
print(f"HTTP Error {e.code}: {e.reason}")
error_body = e.read()
# Attempt to read and print the JSON error body
try:
json_error_body = json.loads(error_body)
print(json_error_body)
raise e
except json.JSONDecodeError:
print("Failed to decode error response as JSON.")
print(error_body)
raise e
except Exception as e:
print(f"An unexpected error occurred: {e}")
raise e
def get_image(filename, subfolder, folder_type):
data = {"filename": filename, "subfolder": subfolder, "type": folder_type}
url_values = urllib.parse.urlencode(data)
with urllib.request.urlopen("http://{}/view?{}".format(server_address, url_values)) as response:
return response.read()
def get_history(prompt_id):
with urllib.request.urlopen("http://{}/history/{}".format(server_address, prompt_id)) as response:
return json.loads(response.read())
def get_images(ws, prompt):
prompt_id = queue_prompt(prompt)['prompt_id']
output_images = {}
while True:
out = ws.recv()
if isinstance(out, str):
message = json.loads(out)
if message['type'] == 'executing':
data = message['data']
if data['node'] is None and data['prompt_id'] == prompt_id:
break #Execution is done
else:
continue #previews are binary data
history = get_history(prompt_id)[prompt_id]
for o in history['outputs']:
for node_id in history['outputs']:
node_output = history['outputs'][node_id]
if 'images' in node_output:
images_output = []
for image in node_output['images']:
image_data = get_image(image['filename'], image['subfolder'], image['type'])
images_output.append(image_data)
output_images[node_id] = images_output
return output_images
# Load json from file relative to this script
prompt = json.load(open(os.path.join(os.path.dirname(os.path.realpath(__file__)), "workflow_api.json")))
#set the text prompt for our positive CLIPTextEncode
# prompt["6"]["inputs"]["text"] = "masterpiece best quality man"
#set the seed for our KSampler node
print(queue_prompt(prompt))
ws = websocket.WebSocket()
ws.connect("ws://{}/ws?clientId={}".format(server_address, client_id))
while True:
out = ws.recv()
if isinstance(out, str):
message = json.loads(out)
# print(message)
if message["type"] == "executing" and message["data"]["node"] is None:
print("Execution is done")
break
if message["type"] == "execution_error":
print("Execution error")
print(json.dumps(message["data"], indent=4))
raise Exception("Execution error")
# images = get_images(ws, prompt)
#Commented out code to display the output images:
# for node_id in images:
# for image_data in images[node_id]:
# from PIL import Image
# import io
# image = Image.open(io.BytesIO(image_data))
# image.show()
-84
View File
@@ -1,84 +0,0 @@
{
"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"
}
}