Compare commits
6
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e0a15a657c | ||
|
|
dbf1fbf287 | ||
|
|
0413b0c1c8 | ||
|
|
6238776610 | ||
|
|
bb03530d5d | ||
|
|
447d04dffe |
@@ -1,10 +0,0 @@
|
|||||||
name: Test
|
|
||||||
on:
|
|
||||||
workflow_dispatch:
|
|
||||||
jobs:
|
|
||||||
build:
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
steps:
|
|
||||||
- uses: actions/checkout@v2
|
|
||||||
- name: Setup upterm session
|
|
||||||
uses: lhotari/action-upterm@v1
|
|
||||||
@@ -1,85 +0,0 @@
|
|||||||
name: Run Test Workflow
|
|
||||||
on: [workflow_dispatch]
|
|
||||||
|
|
||||||
jobs:
|
|
||||||
Test:
|
|
||||||
runs-on: ubuntu-latest
|
|
||||||
strategy:
|
|
||||||
fail-fast: false
|
|
||||||
matrix:
|
|
||||||
python-version: [ "3.7", "3.8", "3.9", "3.10", "3.11" ]
|
|
||||||
steps:
|
|
||||||
- name: Clone Upstream
|
|
||||||
uses: actions/checkout@v3
|
|
||||||
with:
|
|
||||||
repository: comfyanonymous/ComfyUI
|
|
||||||
ref: master
|
|
||||||
fetch-depth: 0
|
|
||||||
|
|
||||||
- name: Clone Node
|
|
||||||
uses: actions/checkout@v3
|
|
||||||
with:
|
|
||||||
ref: master
|
|
||||||
fetch-depth: 0
|
|
||||||
path: custom_nodes/KepPromptLang
|
|
||||||
|
|
||||||
- name: Setup Python
|
|
||||||
uses: actions/setup-python@v4
|
|
||||||
with:
|
|
||||||
# Version range or exact version of Python or PyPy to use, using SemVer's version range syntax. Reads from .python-version if unset.
|
|
||||||
python-version: ${{ matrix.python-version }}
|
|
||||||
|
|
||||||
- name: Cache virtualenv
|
|
||||||
uses: actions/cache@v3
|
|
||||||
id: cache-venv
|
|
||||||
with:
|
|
||||||
path: ./.venv/
|
|
||||||
key: ${{ runner.os }}-venv-${{ matrix.python-version }}-${{ hashFiles('**/requirements.txt') }}
|
|
||||||
restore-keys: |
|
|
||||||
${{ runner.os }}-venv-${{ matrix.python-version }}-
|
|
||||||
|
|
||||||
- name: Install Requirements
|
|
||||||
if: steps.cache-venv.outputs.cache-hit != 'true'
|
|
||||||
run: |
|
|
||||||
python -m venv ./.venv
|
|
||||||
source ./.venv/bin/activate
|
|
||||||
pip install torch --index-url https://download.pytorch.org/whl/cpu
|
|
||||||
pip install -r requirements.txt
|
|
||||||
pip install -r custom_nodes/KepPromptLang/requirements.txt
|
|
||||||
pip install huggingface_hub websocket-client
|
|
||||||
|
|
||||||
# - name: Cache SD Checkpoint
|
|
||||||
# uses: actions/cache@v3
|
|
||||||
# with:
|
|
||||||
# path: |
|
|
||||||
# models/checkpoints
|
|
||||||
# key: ${{ runner.os }}-sd-15-checkpoint
|
|
||||||
|
|
||||||
- name: Check and Download Model
|
|
||||||
run: |
|
|
||||||
source ./.venv/bin/activate
|
|
||||||
python custom_nodes/KepPromptLang/test_files/check_and_download_model.py
|
|
||||||
|
|
||||||
- name: Run in Background
|
|
||||||
env:
|
|
||||||
PYTHONUNBUFFERED: 1
|
|
||||||
run: |
|
|
||||||
source ./.venv/bin/activate
|
|
||||||
python main.py --cpu &> server.log &
|
|
||||||
sleep 10
|
|
||||||
|
|
||||||
# - name: Setup upterm session
|
|
||||||
# uses: lhotari/action-upterm@v1
|
|
||||||
|
|
||||||
- name: Run Workflow
|
|
||||||
run: |
|
|
||||||
source ./.venv/bin/activate
|
|
||||||
python custom_nodes/KepPromptLang/test_files/run_workflow.py
|
|
||||||
|
|
||||||
- name: Upload Comfy Server Log
|
|
||||||
if: always()
|
|
||||||
uses: actions/upload-artifact@v3
|
|
||||||
with:
|
|
||||||
name: comfy-server-log-${{ matrix.python-version }}
|
|
||||||
path: server.log
|
|
||||||
|
|
||||||
@@ -1,81 +1 @@
|
|||||||
# ClipStuff
|
# 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)
|
|
||||||
```
|
|
||||||
```
|
|
||||||
|
|||||||
@@ -1,9 +1,11 @@
|
|||||||
from .nodes import (
|
from .nodes import (
|
||||||
|
KepAdvTextEncode,
|
||||||
BuildGif,
|
BuildGif,
|
||||||
SpecialClipLoader,
|
SpecialClipLoader,
|
||||||
)
|
)
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
|
"Kep Adv Text Encode": KepAdvTextEncode,
|
||||||
"Build Gif": BuildGif,
|
"Build Gif": BuildGif,
|
||||||
"Special CLIP Loader": SpecialClipLoader,
|
"Special CLIP Loader": SpecialClipLoader,
|
||||||
}
|
}
|
||||||
|
|||||||
Binary file not shown.
|
Before Width: | Height: | Size: 1.8 MiB |
File diff suppressed because it is too large
Load Diff
@@ -1,22 +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.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)
|
|
||||||
|
|||||||
@@ -1,117 +0,0 @@
|
|||||||
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
|
|
||||||
@@ -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]
|
||||||
|
|||||||
@@ -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)
|
|
||||||
@@ -0,0 +1,131 @@
|
|||||||
|
from typing import Callable, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch.nn import Embedding
|
||||||
|
|
||||||
|
from comfy.sd1_clip import SD1Tokenizer
|
||||||
|
from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment
|
||||||
|
|
||||||
|
|
||||||
|
class ArithAction(Action):
|
||||||
|
START_CHAR = "<"
|
||||||
|
END_CHAR = ">"
|
||||||
|
|
||||||
|
def __init__(self, base_segment: PromptSegment | Action, ops: dict[str, list[PromptSegment | Action]]):
|
||||||
|
self.base_segment = base_segment
|
||||||
|
self.ops = ops
|
||||||
|
|
||||||
|
def __repr__(self):
|
||||||
|
return f"ArithAction(\n\tbase_segment={self.base_segment},\n\tops={self.ops}\n)"
|
||||||
|
|
||||||
|
def depth_repr(self, depth=1):
|
||||||
|
out = "ArithAction(\n"
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
base_segment_repr = self.base_segment.depth_repr(depth + 1)
|
||||||
|
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
|
||||||
|
elif isinstance(self.base_segment, PromptSegment):
|
||||||
|
out += "\t" * depth + f'base_segment={self.base_segment.depth_repr(depth)}'
|
||||||
|
else:
|
||||||
|
out += "\t" * depth + f'base_segment="{self.base_segment}",'
|
||||||
|
|
||||||
|
for op_key, ops in self.ops.items():
|
||||||
|
for op in ops:
|
||||||
|
out += "\n" + "\t" * depth + f'"{op_key}":[\n'
|
||||||
|
if isinstance(op, Action):
|
||||||
|
op_repr = op.depth_repr(depth + 2)
|
||||||
|
out += "\t" * (depth + 1) + f"{op_repr}\n"
|
||||||
|
else:
|
||||||
|
out += "\t" * (depth + 1) + f'{op.depth_repr()},\n'
|
||||||
|
out += "\t" * depth + "],"
|
||||||
|
out += "\n" + "\t" * (depth - 1) + ")"
|
||||||
|
return out
|
||||||
|
|
||||||
|
def token_length(self):
|
||||||
|
# ArithAction modifies the embeddings of the base segment, so the length is the length of the base segment
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
return self.base_segment.token_length()
|
||||||
|
|
||||||
|
return len(self.base_segment.tokens)
|
||||||
|
|
||||||
|
|
||||||
|
def get_all_segments(self):
|
||||||
|
segments = []
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
segments += self.base_segment.get_all_segments()
|
||||||
|
else:
|
||||||
|
segments.append(self.base_segment)
|
||||||
|
|
||||||
|
for op_key, ops in self.ops.items():
|
||||||
|
for op in ops:
|
||||||
|
if isinstance(op, Action):
|
||||||
|
segments += op.get_all_segments()
|
||||||
|
else:
|
||||||
|
segments.append(op)
|
||||||
|
|
||||||
|
return segments
|
||||||
|
|
||||||
|
def get_result(self, embedding_module: Embedding):
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
base_segment_result = self.base_segment.get_result(embedding_module)
|
||||||
|
else:
|
||||||
|
base_segment_result = self.base_segment.get_embeddings(embedding_module)
|
||||||
|
|
||||||
|
for op_key, ops in self.ops.items():
|
||||||
|
for op in ops:
|
||||||
|
if isinstance(op, Action):
|
||||||
|
op_result = op.get_result(embedding_module)
|
||||||
|
else:
|
||||||
|
op_result = op.get_embeddings(embedding_module)
|
||||||
|
|
||||||
|
|
||||||
|
if op_result.shape[1] > base_segment_result.shape[1]:
|
||||||
|
print('[WARN] ArithAction: op_result.shape[1] > base_segment_result.shape[1] - averaging op_result')
|
||||||
|
op_result = torch.mean(op_result, dim=1, keepdim=True)
|
||||||
|
|
||||||
|
if op_key == "+":
|
||||||
|
base_segment_result.add(op_result)
|
||||||
|
elif op_key == "-":
|
||||||
|
base_segment_result.subtract(op_result)
|
||||||
|
|
||||||
|
return base_segment_result
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def parse_segment(
|
||||||
|
cls,
|
||||||
|
tokens: list[str],
|
||||||
|
start_chars: list[str],
|
||||||
|
end_chars: list[str],
|
||||||
|
parent_parser: Callable[[list[str], SD1Tokenizer], Union[PromptSegment, 'Action']],
|
||||||
|
tokenizer: SD1Tokenizer,
|
||||||
|
) -> Action:
|
||||||
|
"""
|
||||||
|
Parse an arithmetic action from a list of tokens
|
||||||
|
Supported formats:
|
||||||
|
<base_segment:+op1-op2-op3>
|
||||||
|
|
||||||
|
:param tokens: List of tokens, will be modified
|
||||||
|
:param start_chars: List of start chars for all actions
|
||||||
|
:param end_chars: List of end chars for all actions
|
||||||
|
:param parent_parser: Function to parse segments to allow for nested actions
|
||||||
|
:return:
|
||||||
|
"""
|
||||||
|
token = tokens.pop(0)
|
||||||
|
assert token == cls.START_CHAR, "ArithAction must start with " + cls.START_CHAR + " but got " + token
|
||||||
|
|
||||||
|
# Parse base segment
|
||||||
|
base_segment = parent_parser(tokens, tokenizer)
|
||||||
|
|
||||||
|
token = tokens.pop(0)
|
||||||
|
assert token == ":", "ArithAction must have a ':' after the base segment" + " but got " + token
|
||||||
|
|
||||||
|
# Parse ops string
|
||||||
|
ops = {'+': [], '-': []}
|
||||||
|
while tokens[0] != cls.END_CHAR:
|
||||||
|
op_char = tokens.pop(0)
|
||||||
|
assert op_char in ["+", "-"], "ArithAction must have a '+' or '-' as an op char but got " + op_char
|
||||||
|
ops[op_char].append(parent_parser(tokens, tokenizer))
|
||||||
|
|
||||||
|
token = tokens.pop(0)
|
||||||
|
assert token == cls.END_CHAR, "ArithAction must end with " + cls.END_CHAR + " but got " + token
|
||||||
|
|
||||||
|
return cls(base_segment, ops)
|
||||||
@@ -1,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
|
|
||||||
@@ -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)
|
||||||
@@ -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
|
|
||||||
|
|
||||||
@@ -0,0 +1,17 @@
|
|||||||
|
from custom_nodes.ClipStuff.lib.actions import ALL_START_CHARS, ALL_END_CHARS
|
||||||
|
from custom_nodes.ClipStuff.lib.actions.base import Action
|
||||||
|
|
||||||
|
|
||||||
|
def is_action_segment(action_class: Action.__class__, segment: str):
|
||||||
|
if not issubclass(action_class, Action):
|
||||||
|
raise Exception(
|
||||||
|
f"action_class must be a subclass of Action, got {action_class}"
|
||||||
|
)
|
||||||
|
|
||||||
|
return (
|
||||||
|
segment[0] == action_class.START_CHAR and segment[-1] == action_class.END_CHAR
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def is_any_action_segment(segment: str):
|
||||||
|
return segment[0] in ALL_START_CHARS and segment[-1] in ALL_END_CHARS
|
||||||
@@ -1,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
|
|
||||||
|
|
||||||
@@ -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
|
|
||||||
|
|
||||||
@@ -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))
|
|
||||||
|
|
||||||
@@ -0,0 +1,128 @@
|
|||||||
|
from typing import Optional, Union, Callable
|
||||||
|
|
||||||
|
import torch
|
||||||
|
from torch.nn import Embedding
|
||||||
|
|
||||||
|
from comfy.sd1_clip import SD1Tokenizer
|
||||||
|
from custom_nodes.ClipStuff.lib.actions.base import Action, PromptSegment
|
||||||
|
|
||||||
|
|
||||||
|
class NudgeAction(Action):
|
||||||
|
def token_length(self):
|
||||||
|
# Nudge nudges the embeddings of the base segment, so the length is the length of the base segment
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
return self.base_segment.token_length()
|
||||||
|
|
||||||
|
return len(self.base_segment.tokens)
|
||||||
|
|
||||||
|
def get_all_segments(self):
|
||||||
|
segments = []
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
segments += self.base_segment.get_all_segments()
|
||||||
|
else:
|
||||||
|
segments.append(self.base_segment)
|
||||||
|
|
||||||
|
if isinstance(self.target, Action):
|
||||||
|
segments += self.target.get_all_segments()
|
||||||
|
else:
|
||||||
|
segments.append(self.target)
|
||||||
|
|
||||||
|
return segments
|
||||||
|
|
||||||
|
def get_result(self, embedding_module: Embedding):
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
base_segment_result = self.base_segment.get_result(embedding_module)
|
||||||
|
else:
|
||||||
|
base_segment_result = self.base_segment.get_embeddings(embedding_module)
|
||||||
|
|
||||||
|
if isinstance(self.target, Action):
|
||||||
|
target_segment_result = self.target.get_result(embedding_module)
|
||||||
|
else:
|
||||||
|
target_segment_result = self.target.get_embeddings(embedding_module)
|
||||||
|
|
||||||
|
base_mean = torch.mean(base_segment_result, dim=1, keepdim=True)
|
||||||
|
if target_segment_result.shape[1] == 1:
|
||||||
|
translation_vector = target_segment_result - base_mean
|
||||||
|
else:
|
||||||
|
translation_vector = torch.mean(target_segment_result, dim=1, keepdim=True) - base_mean
|
||||||
|
|
||||||
|
return base_segment_result.add(translation_vector, alpha=self.weight)
|
||||||
|
|
||||||
|
START_CHAR = "["
|
||||||
|
END_CHAR = "]"
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
base_segment: PromptSegment | Action,
|
||||||
|
target: Union[PromptSegment, Action],
|
||||||
|
weight: Optional[float] = None,
|
||||||
|
):
|
||||||
|
self.base_segment = base_segment
|
||||||
|
self.weight = weight
|
||||||
|
self.target = target
|
||||||
|
|
||||||
|
def __repr__(self):
|
||||||
|
return f"NudgeAction(\n\tbase_segment={self.base_segment},\n\ttarget={self.target},\n\tweight={self.weight}\n)"
|
||||||
|
|
||||||
|
def depth_repr(self, depth=1):
|
||||||
|
out = "NudgeAction(\n"
|
||||||
|
if isinstance(self.base_segment, Action):
|
||||||
|
base_segment_repr = self.base_segment.depth_repr(depth + 1)
|
||||||
|
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
|
||||||
|
else:
|
||||||
|
out += "\t" * depth + f'base_segment={self.base_segment.depth_repr()},\n'
|
||||||
|
|
||||||
|
if isinstance(self.target, Action):
|
||||||
|
target_repr = self.target.depth_repr(depth + 1)
|
||||||
|
out += "\t" * depth + f"target={target_repr},\n"
|
||||||
|
else:
|
||||||
|
out += "\t" * depth + f"target={self.target.depth_repr()},\n"
|
||||||
|
out += "\t" * depth + f"weight={self.weight},\n"
|
||||||
|
out += "\t" * (depth - 1) + ")"
|
||||||
|
return out
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def parse_segment(
|
||||||
|
cls,
|
||||||
|
tokens: list[str],
|
||||||
|
start_chars: list[str],
|
||||||
|
end_chars: list[str],
|
||||||
|
parent_parser: Callable[[list[str], SD1Tokenizer], PromptSegment | Action],
|
||||||
|
tokenizer: SD1Tokenizer,
|
||||||
|
) -> Action:
|
||||||
|
"""
|
||||||
|
Parse a nudge action from a list of tokens
|
||||||
|
Supported formats:
|
||||||
|
[base_segment:target_segment]
|
||||||
|
[base_segment:target_segment:weight]
|
||||||
|
|
||||||
|
Weight is optional, if not provided it will be None
|
||||||
|
:param tokens: List of tokens, will be modified
|
||||||
|
:param start_chars: List of start chars for all actions
|
||||||
|
:param end_chars: List of end chars for all actions
|
||||||
|
:param parent_parser: Function to parse segments to allow for nested actions
|
||||||
|
:return:
|
||||||
|
"""
|
||||||
|
token = tokens.pop(0)
|
||||||
|
assert token == cls.START_CHAR, "NudgeAction must start with " + cls.START_CHAR + " got " + token
|
||||||
|
|
||||||
|
# Parse base segment
|
||||||
|
base_segment = parent_parser(tokens, tokenizer)
|
||||||
|
|
||||||
|
token = tokens.pop(0)
|
||||||
|
assert token == ":", "NudgeAction must have a ':' after the base segment" + " but got " + token
|
||||||
|
|
||||||
|
# Parse target segment
|
||||||
|
target_segment = parent_parser(tokens, tokenizer)
|
||||||
|
|
||||||
|
# Parse weight if it exists
|
||||||
|
weight = None
|
||||||
|
if tokens[0] == ":":
|
||||||
|
# Parse weight
|
||||||
|
tokens.pop(0)
|
||||||
|
weight = float(tokens.pop(0))
|
||||||
|
|
||||||
|
token = tokens.pop(0)
|
||||||
|
assert token == cls.END_CHAR, "NudgeAction must end with " + cls.END_CHAR + " got " + token
|
||||||
|
|
||||||
|
return cls(base_segment, target_segment, weight)
|
||||||
@@ -1,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
|
|
||||||
|
|
||||||
@@ -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
|
|
||||||
@@ -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
|
|
||||||
@@ -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
|
|
||||||
@@ -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
|
|
||||||
|
|
||||||
@@ -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]
|
|
||||||
@@ -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
|
|
||||||
@@ -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
-61
@@ -1,21 +1,18 @@
|
|||||||
import contextlib
|
import contextlib
|
||||||
import os
|
import os
|
||||||
from typing import List
|
from typing import Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from transformers import CLIPTextConfig, modeling_utils
|
from transformers import CLIPTextConfig, modeling_utils
|
||||||
|
|
||||||
from comfy import model_management
|
from comfy import model_management
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
from comfy.sdxl_clip import SDXLClipModel
|
from comfy.sd import CLIP
|
||||||
from custom_nodes.KepPromptLang.lib.action.base import Action
|
from custom_nodes.ClipStuff.lib.actions.base import PromptSegment, Action
|
||||||
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
from custom_nodes.ClipStuff.lib.fun_clip_stuff import MyCLIPTextModel
|
||||||
from custom_nodes.KepPromptLang.lib.fun_clip_stuff import PromptLangTextModel
|
from custom_nodes.ClipStuff.lib.tokenizer import TokenDict
|
||||||
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)"""
|
"""Uses the CLIP transformer encoder for text (from huggingface)"""
|
||||||
LAYERS = [
|
LAYERS = [
|
||||||
"last",
|
"last",
|
||||||
@@ -25,37 +22,28 @@ class PromptLangClipModel(torch.nn.Module):
|
|||||||
|
|
||||||
def __init__(self, version="openai/clip-vit-large-patch14", device="cpu", max_length=77,
|
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,
|
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__()
|
super().__init__()
|
||||||
assert layer in self.LAYERS
|
assert layer in self.LAYERS
|
||||||
self.num_layers = 12
|
self.num_layers = 12
|
||||||
if textmodel_path is not None:
|
if textmodel_path is not None:
|
||||||
# Our transformer
|
self.transformer = MyCLIPTextModel.from_pretrained(textmodel_path)
|
||||||
self.transformer = PromptLangTextModel.from_pretrained(textmodel_path)
|
|
||||||
else:
|
else:
|
||||||
if textmodel_json_config is None:
|
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")
|
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "clip_config.json")
|
||||||
config = CLIPTextConfig.from_json_file(textmodel_json_config)
|
config = CLIPTextConfig.from_json_file(textmodel_json_config)
|
||||||
self.num_layers = config.num_hidden_layers
|
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():
|
with modeling_utils.no_init_weights():
|
||||||
# Our transformer
|
self.transformer = MyCLIPTextModel(config)
|
||||||
self.transformer = PromptLangTextModel(config)
|
|
||||||
|
|
||||||
if dtype is not None:
|
|
||||||
self.transformer.to(dtype)
|
|
||||||
self.max_length = max_length
|
self.max_length = max_length
|
||||||
if freeze:
|
if freeze:
|
||||||
self.freeze()
|
self.freeze()
|
||||||
self.layer = layer
|
self.layer = layer
|
||||||
self.layer_idx = None
|
self.layer_idx = None
|
||||||
self.empty_tokens = [[49406] + [49407] * 76]
|
self.empty_tokens = [[49406] + [49407] * 76]
|
||||||
self.text_projection = torch.nn.Parameter(torch.eye(self.transformer.get_input_embeddings().weight.shape[1]))
|
self.text_projection = None
|
||||||
self.logit_scale = torch.nn.Parameter(torch.tensor(4.6055))
|
|
||||||
|
|
||||||
self.layer_norm_hidden_state = True
|
self.layer_norm_hidden_state = True
|
||||||
if layer == "hidden":
|
if layer == "hidden":
|
||||||
assert layer_idx is not None
|
assert layer_idx is not None
|
||||||
@@ -80,8 +68,7 @@ class PromptLangClipModel(torch.nn.Module):
|
|||||||
self.layer = self.layer_default[0]
|
self.layer = self.layer_default[0]
|
||||||
self.layer_idx = self.layer_default[1]
|
self.layer_idx = self.layer_default[1]
|
||||||
|
|
||||||
# Completely changed to support Segments and actions
|
def set_up_textual_embeddings(self, tokens: list[list[PromptSegment | Action]], current_embeds):
|
||||||
def set_up_textual_embeddings(self, tokens: List[List[SegOrAction]], current_embeds):
|
|
||||||
next_new_token = token_dict_size = current_embeds.weight.shape[0] - 1
|
next_new_token = token_dict_size = current_embeds.weight.shape[0] - 1
|
||||||
embedding_weights = []
|
embedding_weights = []
|
||||||
|
|
||||||
@@ -144,8 +131,7 @@ class PromptLangClipModel(torch.nn.Module):
|
|||||||
if segment.tokens[tokenIdx] == -1:
|
if segment.tokens[tokenIdx] == -1:
|
||||||
segment.tokens[tokenIdx] = n
|
segment.tokens[tokenIdx] = n
|
||||||
|
|
||||||
# Support our set_up_textual_embeddings which modifies the input embeddings
|
def forward(self, tokens, **kwargs):
|
||||||
def forward(self, tokens):
|
|
||||||
backup_embeds = self.transformer.get_input_embeddings()
|
backup_embeds = self.transformer.get_input_embeddings()
|
||||||
device = backup_embeds.weight.device
|
device = backup_embeds.weight.device
|
||||||
self.set_up_textual_embeddings(tokens, backup_embeds)
|
self.set_up_textual_embeddings(tokens, backup_embeds)
|
||||||
@@ -156,8 +142,16 @@ class PromptLangClipModel(torch.nn.Module):
|
|||||||
else:
|
else:
|
||||||
precision_scope = contextlib.nullcontext
|
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)):
|
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)
|
self.transformer.set_input_embeddings(backup_embeds)
|
||||||
|
|
||||||
if self.layer == "last":
|
if self.layer == "last":
|
||||||
@@ -171,27 +165,21 @@ class PromptLangClipModel(torch.nn.Module):
|
|||||||
|
|
||||||
pooled_output = outputs.pooler_output
|
pooled_output = outputs.pooler_output
|
||||||
if self.text_projection is not None:
|
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()
|
return z.float(), pooled_output.float()
|
||||||
|
|
||||||
def encode(self, tokens):
|
def encode(self, tokens, **kwargs):
|
||||||
return self(tokens)
|
return self(tokens, **kwargs)
|
||||||
|
|
||||||
def load_sd(self, sd):
|
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)
|
return self.transformer.load_state_dict(sd, strict=False)
|
||||||
|
|
||||||
# Changed from comfy.sd1_clip.ClipTokenWeightEncoder
|
def encode_token_weights(self, prompt_segments: list[list[Union[PromptSegment | Action]]], **kwargs):
|
||||||
# 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])]]
|
to_encode = [[PromptSegment(text="_Empty Batch_", tokens=self.empty_tokens[0])]]
|
||||||
for batch in prompt_segments:
|
for batch in prompt_segments:
|
||||||
to_encode.append(batch)
|
to_encode.append(batch)
|
||||||
|
|
||||||
out, pooled = self.encode(to_encode)
|
out, pooled = self.encode(to_encode, **kwargs)
|
||||||
z_empty = out[0:1]
|
z_empty = out[0:1]
|
||||||
if pooled.shape[0] > 1:
|
if pooled.shape[0] > 1:
|
||||||
first_pooled = pooled[1:2]
|
first_pooled = pooled[1:2]
|
||||||
@@ -210,25 +198,3 @@ class PromptLangClipModel(torch.nn.Module):
|
|||||||
if (len(output) == 0):
|
if (len(output) == 0):
|
||||||
return z_empty.cpu(), first_pooled.cpu()
|
return z_empty.cpu(), first_pooled.cpu()
|
||||||
return torch.cat(output, dim=-2).cpu(), first_pooled.cpu()
|
return torch.cat(output, dim=-2).cpu(), first_pooled.cpu()
|
||||||
|
|
||||||
class PromptLangSDXLClipModel(SDXLClipModel):
|
|
||||||
def __init__(self, device="cpu", dtype=None) -> None:
|
|
||||||
# Skip SDXLClipModel's init
|
|
||||||
super(SDXLClipModel, self).__init__()
|
|
||||||
self.clip_l = PromptLangClipModel(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(PromptLangClipModel):
|
|
||||||
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)
|
|
||||||
|
|||||||
+58
-35
@@ -1,17 +1,14 @@
|
|||||||
from typing import Optional, Tuple, Union, List
|
from typing import Optional, Tuple, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
from torch import device
|
||||||
from transformers import CLIPTextConfig
|
from transformers import CLIPTextConfig
|
||||||
from transformers.modeling_outputs import BaseModelOutputWithPooling
|
from transformers.modeling_outputs import BaseModelOutputWithPooling
|
||||||
from transformers.models.clip.modeling_clip import (
|
from transformers.models.clip.modeling_clip import _expand_mask, CLIPTextEmbeddings, CLIPTextTransformer, \
|
||||||
_expand_mask,
|
CLIPTextModel
|
||||||
CLIPTextEmbeddings,
|
|
||||||
CLIPTextTransformer,
|
|
||||||
CLIPTextModel,
|
|
||||||
)
|
|
||||||
|
|
||||||
from custom_nodes.KepPromptLang.lib.action.base import Action
|
from custom_nodes.ClipStuff.lib.actions.base import PromptSegment, Action
|
||||||
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
from custom_nodes.ClipStuff.lib.tokenizer import TokenDict
|
||||||
|
|
||||||
def slerp(val, low, high):
|
def slerp(val, low, high):
|
||||||
low = low.unsqueeze(0)
|
low = low.unsqueeze(0)
|
||||||
@@ -23,20 +20,18 @@ def slerp(val, low, high):
|
|||||||
res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high
|
res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high
|
||||||
return res
|
return res
|
||||||
|
|
||||||
class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
class MyCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||||
def __init__(self, config: CLIPTextConfig):
|
def __init__(self, config: CLIPTextConfig):
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
input_dicts: Optional[List[List[SegOrAction]]] = None,
|
input_dicts: Optional[list[list[tuple[TokenDict]]]] = None,
|
||||||
input_ids: Optional[torch.LongTensor] = None,
|
input_ids: Optional[torch.LongTensor] = None,
|
||||||
position_ids: Optional[torch.LongTensor] = None,
|
position_ids: Optional[torch.LongTensor] = None,
|
||||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
|
||||||
if input_dicts is None:
|
|
||||||
raise ValueError("You have to specify input_dicts")
|
|
||||||
|
|
||||||
batches = []
|
batches = []
|
||||||
for batch_idx, batch in enumerate(input_dicts):
|
for batch_idx, batch in enumerate(input_dicts):
|
||||||
@@ -53,6 +48,29 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
|||||||
if position_ids is None:
|
if position_ids is None:
|
||||||
position_ids = self.position_ids[:, :seq_length]
|
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 = []
|
embeds = []
|
||||||
for batch in batches:
|
for batch in batches:
|
||||||
if len(batch) == 1:
|
if len(batch) == 1:
|
||||||
@@ -66,14 +84,14 @@ class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
|||||||
return embeddings
|
return embeddings
|
||||||
|
|
||||||
|
|
||||||
class PrompLangCLIPTextTransformer(CLIPTextTransformer):
|
class MyCLIPTextTransformer(CLIPTextTransformer):
|
||||||
def __init__(self, config: CLIPTextConfig):
|
def __init__(self, config: CLIPTextConfig):
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
self.embeddings = PromptLangCLIPTextEmbeddings(config)
|
self.embeddings = MyCLIPTextEmbeddings(config)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
input_ids: Optional[List[List[SegOrAction]]] = None,
|
input_ids: Optional[list[list[PromptSegment | Action]]] = None,
|
||||||
attention_mask: Optional[torch.Tensor] = None,
|
attention_mask: Optional[torch.Tensor] = None,
|
||||||
position_ids: Optional[torch.Tensor] = None,
|
position_ids: Optional[torch.Tensor] = None,
|
||||||
output_attentions: Optional[bool] = None,
|
output_attentions: Optional[bool] = None,
|
||||||
@@ -101,20 +119,12 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
|
|||||||
bsz = len(input_ids)
|
bsz = len(input_ids)
|
||||||
# TODO: Properly gather this
|
# TODO: Properly gather this
|
||||||
seq_len = 77
|
seq_len = 77
|
||||||
input_shape = torch.Size([bsz, seq_len])
|
# bsz, seq_len = input_shape
|
||||||
# CLIP's text model uses causal mask, prepare it here.
|
# CLIP's text model uses causal mask, prepare it here.
|
||||||
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
|
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
|
||||||
## VERSION DIFF ##
|
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(
|
||||||
# transformers < 4.30.0
|
hidden_states.device
|
||||||
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
|
# expand attention_mask
|
||||||
if attention_mask is not None:
|
if attention_mask is not None:
|
||||||
# [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
|
# [bsz, seq_len] -> [bsz, 1, tgt_seq_len, src_seq_len]
|
||||||
@@ -143,9 +153,6 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
|
|||||||
else:
|
else:
|
||||||
if seg_or_action.text == '__PAD__':
|
if seg_or_action.text == '__PAD__':
|
||||||
break
|
break
|
||||||
|
|
||||||
# Is a segment, and isn't the pad segment
|
|
||||||
idx += seg_or_action.token_length()
|
|
||||||
eot_idx.append(idx)
|
eot_idx.append(idx)
|
||||||
# text_embeds.shape = [batch_size, sequence_length, transformer.width]
|
# text_embeds.shape = [batch_size, sequence_length, transformer.width]
|
||||||
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
||||||
@@ -167,21 +174,37 @@ class PrompLangCLIPTextTransformer(CLIPTextTransformer):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
# This is necessary to pass the PromptLangCLIPTextTransformer
|
class MyCLIPTextModel(CLIPTextModel):
|
||||||
class PromptLangTextModel(CLIPTextModel):
|
|
||||||
def __init__(self, config: CLIPTextConfig):
|
def __init__(self, config: CLIPTextConfig):
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
self.text_model = PrompLangCLIPTextTransformer(config)
|
self.text_model = MyCLIPTextTransformer(config)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
input_ids: Optional[List[List[SegOrAction]]] = None,
|
input_ids: Optional[list[list[tuple[TokenDict]]]] = None,
|
||||||
attention_mask: Optional[torch.Tensor] = None,
|
attention_mask: Optional[torch.Tensor] = None,
|
||||||
position_ids: Optional[torch.Tensor] = None,
|
position_ids: Optional[torch.Tensor] = None,
|
||||||
output_attentions: Optional[bool] = None,
|
output_attentions: Optional[bool] = None,
|
||||||
output_hidden_states: Optional[bool] = None,
|
output_hidden_states: Optional[bool] = None,
|
||||||
return_dict: Optional[bool] = None,
|
return_dict: Optional[bool] = None,
|
||||||
) -> Union[Tuple, BaseModelOutputWithPooling]:
|
) -> 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_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||||
|
|
||||||
return self.text_model(
|
return self.text_model(
|
||||||
|
|||||||
@@ -1,5 +0,0 @@
|
|||||||
from lark import Lark
|
|
||||||
|
|
||||||
from .grammar import grammar
|
|
||||||
|
|
||||||
PromptParser = Lark(grammar, start="start", parser="earley")
|
|
||||||
@@ -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
|
|
||||||
"""
|
|
||||||
@@ -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
|
|
||||||
@@ -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]
|
|
||||||
@@ -1,63 +0,0 @@
|
|||||||
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}")
|
|
||||||
@@ -1,38 +0,0 @@
|
|||||||
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)
|
|
||||||
+130
-37
@@ -1,16 +1,117 @@
|
|||||||
from typing import List, Dict
|
import re
|
||||||
|
from typing import Union
|
||||||
from lark import Tree
|
|
||||||
|
|
||||||
from comfy.sd1_clip import SD1Tokenizer
|
from comfy.sd1_clip import SD1Tokenizer
|
||||||
from custom_nodes.KepPromptLang.lib.actions.types import SegOrAction
|
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.parser import PromptParser
|
arith_action = r'(<[a-zA-Z0-9\-_]+:[a-zA-Z0-9\-_]+>)'
|
||||||
from custom_nodes.KepPromptLang.lib.parser.transformer import PromptTransformer
|
|
||||||
from custom_nodes.KepPromptLang.lib.parser.prompt_segment import PromptSegment
|
|
||||||
|
|
||||||
class PromptLangTokenizer(SD1Tokenizer):
|
# TODO: Get embedding identifier from tokenizer
|
||||||
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) -> None:
|
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):
|
||||||
|
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)
|
super().__init__(tokenizer_path, max_length, pad_with_end, embedding_directory, embedding_size, embedding_key)
|
||||||
|
|
||||||
"""
|
"""
|
||||||
@@ -18,25 +119,34 @@ class PromptLangTokenizer(SD1Tokenizer):
|
|||||||
Returns batches of segments and actions
|
Returns batches of segments and actions
|
||||||
:return: List of list(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:
|
if self.pad_with_end:
|
||||||
pad_token = self.end_token
|
pad_token = self.end_token
|
||||||
else:
|
else:
|
||||||
pad_token = 0
|
pad_token = 0
|
||||||
|
|
||||||
parsed_prompt = PromptParser.parse(text)
|
parsed_actions = parse_segment_actions(text, self)
|
||||||
parsed_actions = PromptTransformer(self).transform(parsed_prompt)
|
|
||||||
|
# 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
|
# reshape token array to CLIP input size
|
||||||
batched_segments = []
|
batched_segments = []
|
||||||
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token])]
|
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token])]
|
||||||
# batched_segments.append(batch)
|
# batched_segments.append(batch)
|
||||||
batch_size = 1
|
batch_size = 1
|
||||||
if isinstance(parsed_actions, Tree):
|
for segment in parsed_actions:
|
||||||
segments_to_process = parsed_actions.children
|
|
||||||
else:
|
|
||||||
segments_to_process = [parsed_actions]
|
|
||||||
for segment in segments_to_process:
|
|
||||||
num_tokens = segment.token_length()
|
num_tokens = segment.token_length()
|
||||||
# determine if we're going to try and keep the tokens in a single batch
|
# determine if we're going to try and keep the tokens in a single batch
|
||||||
is_large = num_tokens >= self.max_word_length
|
is_large = num_tokens >= self.max_word_length
|
||||||
@@ -53,7 +163,7 @@ class PromptLangTokenizer(SD1Tokenizer):
|
|||||||
batch_size = num_tokens + 1 # +1 for start token
|
batch_size = num_tokens + 1 # +1 for start token
|
||||||
continue
|
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.append(segment)
|
||||||
batch_size += num_tokens
|
batch_size += num_tokens
|
||||||
|
|
||||||
@@ -62,24 +172,7 @@ class PromptLangTokenizer(SD1Tokenizer):
|
|||||||
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length))
|
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length))
|
||||||
batched_segments.append(batch)
|
batched_segments.append(batch)
|
||||||
|
|
||||||
# for batch in batched_segments:
|
for batch in batched_segments:
|
||||||
# batch_size_info(batch)
|
batch_size_info(batch)
|
||||||
|
|
||||||
return batched_segments
|
return batched_segments
|
||||||
|
|
||||||
|
|
||||||
class PromptLangSDXLClipGTokenizer(PromptLangTokenizer):
|
|
||||||
def __init__(self, tokenizer_path=None, embedding_directory=None) -> 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 = PromptLangTokenizer(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
|
|
||||||
|
|||||||
@@ -1,6 +1,4 @@
|
|||||||
import random
|
import random
|
||||||
import os
|
|
||||||
from typing import List, Tuple, Any
|
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
@@ -8,18 +6,9 @@ from PIL import Image
|
|||||||
import folder_paths
|
import folder_paths
|
||||||
import comfy.sd
|
import comfy.sd
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
from comfy.sd2_clip import SD2ClipModel
|
from custom_nodes.ClipStuff.lib.clip_model import SD1FunClipModel
|
||||||
from comfy.sdxl_clip import SDXLClipModel
|
|
||||||
from comfy.supported_models_base import ClipTarget
|
|
||||||
from custom_nodes.KepPromptLang.lib.clip_model import (
|
|
||||||
PromptLangClipModel,
|
|
||||||
PromptLangSDXLClipModel,
|
|
||||||
)
|
|
||||||
|
|
||||||
from custom_nodes.KepPromptLang.lib.tokenizer import (
|
from custom_nodes.ClipStuff.lib.tokenizer import MyTokenizer
|
||||||
PromptLangTokenizer,
|
|
||||||
PromptLangSDXLTokenizer,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class EmptyClass:
|
class EmptyClass:
|
||||||
@@ -28,7 +17,7 @@ class EmptyClass:
|
|||||||
|
|
||||||
class SpecialClipLoader:
|
class SpecialClipLoader:
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls): # type: ignore
|
def INPUT_TYPES(s):
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"source_clip": ("CLIP",),
|
"source_clip": ("CLIP",),
|
||||||
@@ -41,42 +30,79 @@ class SpecialClipLoader:
|
|||||||
CATEGORY = "conditioning"
|
CATEGORY = "conditioning"
|
||||||
|
|
||||||
@staticmethod
|
@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):
|
# TODO: Extract embedding directory from source_clip
|
||||||
clip_target = ClipTarget(PromptLangSDXLTokenizer, PromptLangSDXLClipModel)
|
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.embedding_directory)
|
||||||
clip = comfy.sd.CLIP(clip_target, embedding_directory=source_clip.tokenizer.clip_g.embedding_directory)
|
comfy.sd.load_clip_weights(
|
||||||
comfy.sd.load_clip_weights(
|
clip.cond_stage_model, source_clip.cond_stage_model.state_dict()
|
||||||
clip.cond_stage_model, source_clip.cond_stage_model.state_dict()
|
)
|
||||||
)
|
|
||||||
elif isinstance(source_clip, SD2ClipModel):
|
|
||||||
raise ValueError("SD2 Clip model is not supported.")
|
|
||||||
else:
|
|
||||||
clip_target = ClipTarget(PromptLangTokenizer, PromptLangClipModel)
|
|
||||||
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,)
|
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 = 255.0 * tensor_img.cpu().numpy()
|
||||||
i_np_arr = np.clip(i, 0, 255, out=i).astype(np.uint8, copy=False)
|
i_np_arr = np.clip(i, 0, 255, out=i).astype(np.uint8, copy=False)
|
||||||
return Image.fromarray(i_np_arr)
|
return Image.fromarray(i_np_arr)
|
||||||
|
|
||||||
|
|
||||||
class BuildGif:
|
class BuildGif:
|
||||||
def __init__(self) -> None:
|
def __init__(self):
|
||||||
self.output_dir = folder_paths.get_output_directory()
|
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls): # type: ignore
|
def INPUT_TYPES(cls):
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"images": ("IMAGE",),
|
"images": ("IMAGE",),
|
||||||
"split_every": ("INT", {"default": -1}),
|
"split_every": ("INT", {"default": -1}),
|
||||||
"frame_duration": ("INT", {"default": 125}),
|
|
||||||
"output_mode": (
|
"output_mode": (
|
||||||
["One Per Split", "Big Grid"],
|
["One Per Split", "Big Grid"],
|
||||||
{"default": "Big Grid"},
|
{"default": "Big Grid"},
|
||||||
@@ -85,53 +111,46 @@ class BuildGif:
|
|||||||
}
|
}
|
||||||
|
|
||||||
RELOAD_INST = True
|
RELOAD_INST = True
|
||||||
RETURN_TYPES = ()
|
RETURN_TYPES = ("IMAGE",)
|
||||||
# RETURN_NAMES = ("Gifs",)
|
RETURN_NAMES = ("Gifs",)
|
||||||
INPUT_IS_LIST = True
|
INPUT_IS_LIST = True
|
||||||
FUNCTION = "build_gif"
|
FUNCTION = "build_gif"
|
||||||
# OUTPUT_IS_LIST = (True,)
|
OUTPUT_IS_LIST = (True,)
|
||||||
OUTPUT_NODE = True
|
# OUTPUT_NODE = False
|
||||||
|
|
||||||
CATEGORY = "List Stuff"
|
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("Build GIF called!")
|
||||||
print(f"{type(images)}")
|
print(f"{type(images)}")
|
||||||
|
|
||||||
if len(split_every) > 1:
|
if len(split_every) > 1:
|
||||||
raise Exception("List input for split every is not supported.")
|
raise Exception("List input for split every is not supported.")
|
||||||
|
|
||||||
if len(output_mode) > 1:
|
split_every = split_every[0]
|
||||||
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]
|
|
||||||
batch_size = images[0].size()[0]
|
batch_size = images[0].size()[0]
|
||||||
if split_every_val == -1:
|
if split_every == -1:
|
||||||
split_chunks = 1
|
split_chunks = 1
|
||||||
split_every_val = len(images)
|
split_every = len(images)
|
||||||
else:
|
else:
|
||||||
split_chunks = int(len(images) / split_every_val)
|
split_chunks = int(len(images) / split_every)
|
||||||
|
|
||||||
|
out = []
|
||||||
|
|
||||||
num_wide = batch_size
|
num_wide = batch_size
|
||||||
num_tall = split_chunks
|
num_tall = split_chunks
|
||||||
|
|
||||||
chunked_batches = [
|
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)
|
for chunk_idx in range(split_chunks)
|
||||||
]
|
]
|
||||||
|
|
||||||
frames = []
|
frames = []
|
||||||
results = list()
|
|
||||||
|
|
||||||
if output_mode == "Big Grid":
|
if output_mode == "Big Grid":
|
||||||
# For every image in gif
|
# 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_shape = images[0][0].shape
|
||||||
img_frame = Image.new(
|
img_frame = Image.new(
|
||||||
"RGB", size=(num_wide * img_shape[0], num_tall * img_shape[1])
|
"RGB", size=(num_wide * img_shape[0], num_tall * img_shape[1])
|
||||||
@@ -146,9 +165,8 @@ class BuildGif:
|
|||||||
)
|
)
|
||||||
frames.append(img_frame)
|
frames.append(img_frame)
|
||||||
|
|
||||||
file = f"{filename}_{counter:05}_"
|
|
||||||
save_path = (
|
save_path = (
|
||||||
f"{os.path.join(full_output_folder, file)}"
|
f"{folder_paths.get_output_directory()}/{random.randint(1, 100)}"
|
||||||
)
|
)
|
||||||
frames[0].save(
|
frames[0].save(
|
||||||
f"{save_path}.webp",
|
f"{save_path}.webp",
|
||||||
@@ -158,24 +176,15 @@ class BuildGif:
|
|||||||
save_all=True,
|
save_all=True,
|
||||||
append_images=frames[1:],
|
append_images=frames[1:],
|
||||||
optimize=False,
|
optimize=False,
|
||||||
duration=frame_duration,
|
duration=125,
|
||||||
loop=0,
|
loop=0,
|
||||||
)
|
)
|
||||||
results.append({
|
|
||||||
"filename": f"{file}.webp",
|
|
||||||
"subfolder": subfolder,
|
|
||||||
"type": "output"
|
|
||||||
})
|
|
||||||
elif output_mode == "One Per Split":
|
elif output_mode == "One Per Split":
|
||||||
for split_idx in range(int(split_chunks)):
|
for split_idx in range(int(split_chunks)):
|
||||||
split_start = split_every_val * split_idx
|
split_start = split_every * split_idx
|
||||||
split_end = split_every_val * (split_idx + 1)
|
split_end = split_every * (split_idx + 1)
|
||||||
for batch_idx in range(batch_size):
|
for batch_idx in range(batch_size):
|
||||||
file = f"{filename}_{counter:05}_"
|
save_path = f"{folder_paths.get_output_directory()}/-{batch_idx}-{random.randint(1, 100)}"
|
||||||
save_path = (
|
|
||||||
f"{os.path.join(full_output_folder, file)}"
|
|
||||||
)
|
|
||||||
counter += 1
|
|
||||||
print(save_path)
|
print(save_path)
|
||||||
tensor2img(images[split_start][batch_idx]).save(
|
tensor2img(images[split_start][batch_idx]).save(
|
||||||
f"{save_path}.webp",
|
f"{save_path}.webp",
|
||||||
@@ -185,12 +194,7 @@ class BuildGif:
|
|||||||
for nested_batch in images[split_start + 1 : split_end]
|
for nested_batch in images[split_start + 1 : split_end]
|
||||||
],
|
],
|
||||||
optimize=False,
|
optimize=False,
|
||||||
duration=frame_duration,
|
duration=125,
|
||||||
loop=0,
|
loop=0,
|
||||||
)
|
)
|
||||||
results.append({
|
return (out,)
|
||||||
"filename": f"{file}.webp",
|
|
||||||
"subfolder": subfolder,
|
|
||||||
"type": "output"
|
|
||||||
})
|
|
||||||
return { "ui": { "images": results } }
|
|
||||||
|
|||||||
@@ -1 +0,0 @@
|
|||||||
lark
|
|
||||||
@@ -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.")
|
|
||||||
@@ -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()
|
|
||||||
|
|
||||||
@@ -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"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
Reference in New Issue
Block a user