Compare commits
30
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
919f2dbebf | ||
|
|
bd0093f64e | ||
|
|
bdb4d00910 | ||
|
|
c1cbdbe9df | ||
|
|
c28099021b | ||
|
|
2d2d29bd5a | ||
|
|
d0719cb9bc | ||
|
|
3fc1a7d7b9 | ||
|
|
f66eeb8117 | ||
|
|
d5a2cba894 | ||
|
|
0eaa2c1c41 | ||
|
|
553f90691e | ||
|
|
c4769a797b | ||
|
|
66eef6d7d6 | ||
|
|
0fdc4244f1 | ||
|
|
4057ff8347 | ||
|
|
5f84a4b530 | ||
|
|
fc3560be65 | ||
|
|
833a027486 | ||
|
|
6b8a1ea243 | ||
|
|
07cec3f68f | ||
|
|
6a516d093c | ||
|
|
c9f7895f62 | ||
|
|
b125677e6d | ||
|
|
7ced8b7533 | ||
|
|
d883444fdd | ||
|
|
c6dc33a9ae | ||
|
|
110da543ed | ||
|
|
21eecda007 | ||
|
|
235a681f70 |
@@ -0,0 +1,10 @@
|
||||
name: Test
|
||||
on:
|
||||
workflow_dispatch:
|
||||
jobs:
|
||||
build:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v2
|
||||
- name: Setup upterm session
|
||||
uses: lhotari/action-upterm@v1
|
||||
@@ -1 +1,81 @@
|
||||
# ClipStuff
|
||||
## Basic Instructions.
|
||||
Clone repo into custom_nodes folder.
|
||||
Install the requirements.txt file via pip.
|
||||
|
||||
Pass CLIP output from Load Checkpoint into SpecialClipLoader node, then use the outputted clip with standard Clip Text Encode.
|
||||
|
||||
See example workflow in examples folder.
|
||||
|
||||
### Example Photo
|
||||

|
||||
|
||||
## Functions
|
||||
|
||||
## Syntax Elements
|
||||
|
||||
1. **Embedding**:
|
||||
- Syntax: `embedding:WORD`
|
||||
- Example: `embedding:face_vector`
|
||||
- Represents a named vector embedding(Textual Inversion).
|
||||
|
||||
2. **Word**:
|
||||
- Syntax: Any alphanumeric word including characters such as `,`, `_`, and `-`.
|
||||
- Example: `cat, dog_face, id_123`
|
||||
- Represents simple words or identifiers.
|
||||
|
||||
3. **Quoted String**:
|
||||
- Syntax: A string enclosed within double or single quotes. You can escape quotes inside the string using a backslash (`\`).
|
||||
- Example: `"Hello World"`, `'It\'s a sunny day'`
|
||||
- Represents string literals.
|
||||
|
||||
## Functions
|
||||
|
||||
Here are the available functions and their usage:
|
||||
|
||||
1. **Sum Function**:
|
||||
- Syntax: `sum(arg1 | arg2 | ... | argN)`
|
||||
- Adds together multiple embeddings.
|
||||
- Example: `sum(embedding:face1 | dog)`
|
||||
|
||||
2. **Negation Function**:
|
||||
- Syntax: `neg(arg)`
|
||||
- Negates the output.
|
||||
- Example: `neg(A embedding:happycats outside)`
|
||||
|
||||
3. **Normalization Function**:
|
||||
- Syntax: `norm(arg)`
|
||||
- Normalizes the given vector embedding.
|
||||
- Example: `norm(sum(embedding:face1 | embedding:face2))`
|
||||
|
||||
4. **Difference Function**:
|
||||
- Syntax: `diff(arg1 | arg2 | ... | argN)`
|
||||
- Computes the difference between multiple vector embeddings.
|
||||
- Example: `diff(embedding:face1 | embedding:face2)`
|
||||
|
||||
### Notes on Arguments:
|
||||
- Each function takes one or more arguments.
|
||||
- An argument (`arg`) can be an embedding, a word, another function, or a quoted string.
|
||||
- For functions that accept multiple arguments, they are separated by the `|` symbol.
|
||||
|
||||
## Examples
|
||||
|
||||
1. Add two embeddings and normalize the result:
|
||||
```
|
||||
norm(sum(cat | dog | horse | parrot))
|
||||
```
|
||||
|
||||
2. Negate an embedding:
|
||||
```
|
||||
neg(embedding:body_vector)
|
||||
```
|
||||
|
||||
3. King - Man + Woman = Queen:
|
||||
```
|
||||
sum(diff(king|man)|woman)
|
||||
```
|
||||
or
|
||||
```
|
||||
sum(king|neg(man)|woman)
|
||||
```
|
||||
```
|
||||
|
||||
@@ -1,11 +1,9 @@
|
||||
from .nodes import (
|
||||
KepAdvTextEncode,
|
||||
BuildGif,
|
||||
SpecialClipLoader,
|
||||
)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"Kep Adv Text Encode": KepAdvTextEncode,
|
||||
"Build Gif": BuildGif,
|
||||
"Special CLIP Loader": SpecialClipLoader,
|
||||
}
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 1.8 MiB |
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,100 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Union
|
||||
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
|
||||
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
|
||||
class Action(ABC):
|
||||
@property
|
||||
@abstractmethod
|
||||
def chars(self) -> list[str] | None:
|
||||
pass
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def name(self) -> str:
|
||||
pass
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def grammar(self) -> str:
|
||||
"""
|
||||
The grammar for this action. This is used to parse the action from the prompt.
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def token_length(self) -> int:
|
||||
"""
|
||||
The length of the tokens that this action will add to the prompt.
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_all_segments(self) -> list[PromptSegment]:
|
||||
"""
|
||||
Get all segments, including nested segments.
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_result(self, embedding_module: Embedding) -> Tensor:
|
||||
"""
|
||||
Get the result of this action. This is called when the embeddings are being calculated.
|
||||
:param embedding_module: The embedding module to use to get the base embeddings for tokens.
|
||||
:return:
|
||||
"""
|
||||
pass
|
||||
|
||||
def depth_repr(self, depth: int = 1) -> str:
|
||||
raise NotImplementedError()
|
||||
|
||||
|
||||
class SingleArgAction(Action, ABC):
|
||||
def get_all_segments(self) -> list[PromptSegment]:
|
||||
segments = []
|
||||
for seg_or_action in self.arg:
|
||||
if isinstance(seg_or_action, Action):
|
||||
segments.extend(seg_or_action.get_all_segments())
|
||||
else:
|
||||
segments.append(seg_or_action)
|
||||
return segments
|
||||
|
||||
def __init__(self, arg: list[PromptSegment | Action]):
|
||||
# TODO: Target is a list now... what does this mean for us..
|
||||
self.arg = arg
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"{self.name}({self.arg})"
|
||||
|
||||
class MultiArgAction(Action, ABC):
|
||||
def get_all_segments(self) -> list[PromptSegment]:
|
||||
segments = []
|
||||
for seg_or_action in self.base_segment:
|
||||
if isinstance(seg_or_action, Action):
|
||||
segments.extend(seg_or_action.get_all_segments())
|
||||
else:
|
||||
segments.append(seg_or_action)
|
||||
|
||||
for arg in self.args:
|
||||
for seg_or_action in arg:
|
||||
if isinstance(seg_or_action, Action):
|
||||
segments.extend(seg_or_action.get_all_segments())
|
||||
else:
|
||||
segments.append(seg_or_action)
|
||||
|
||||
return segments
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_segment: list[PromptSegment | Action],
|
||||
args: list[list[Union[PromptSegment, Action]]],
|
||||
):
|
||||
self.base_segment = base_segment
|
||||
self.args = args
|
||||
@@ -1,6 +0,0 @@
|
||||
from custom_nodes.ClipStuff.lib.actions.arith import ArithAction
|
||||
from custom_nodes.ClipStuff.lib.actions.nudge import NudgeAction
|
||||
|
||||
ALL_ACTIONS = [NudgeAction, ArithAction]
|
||||
ALL_START_CHARS = [action.START_CHAR for action in ALL_ACTIONS]
|
||||
ALL_END_CHARS = [action.END_CHAR for action in ALL_ACTIONS]
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
from torch import Tensor
|
||||
from torch.nn import Embedding
|
||||
|
||||
from custom_nodes.ClipStuff.lib.action.base import Action
|
||||
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
|
||||
|
||||
|
||||
def get_embedding(seg_or_action: SegOrAction, embedding_module: Embedding) -> Tensor:
|
||||
if isinstance(seg_or_action, Action):
|
||||
return seg_or_action.get_result(embedding_module)
|
||||
return seg_or_action.get_embeddings(embedding_module)
|
||||
@@ -1,39 +0,0 @@
|
||||
from custom_nodes.ClipStuff.lib.actions.base import Action
|
||||
|
||||
|
||||
class ArithAction(Action):
|
||||
START_CHAR = "<"
|
||||
END_CHAR = ">"
|
||||
|
||||
def __init__(self, base_segment: str, ops_str: str):
|
||||
self.base_segment = base_segment
|
||||
self.ops = self.process_ops_string(ops_str)
|
||||
|
||||
@classmethod
|
||||
def process_ops_string(cls, ops_string):
|
||||
supported_ops = ["+", "-"]
|
||||
# dict[[Union[Literal['add'], Literal['subtract']]], str]
|
||||
ops_dict = {"+": [], "-": []}
|
||||
buff = ""
|
||||
curr_op_char = ""
|
||||
for char in ops_string:
|
||||
if char in supported_ops:
|
||||
# We have a buffer
|
||||
if buff != "":
|
||||
# Add op string
|
||||
ops_dict[curr_op_char] += [buff]
|
||||
# Reset buffer
|
||||
buff = ""
|
||||
# Set new current op char
|
||||
curr_op_char = char
|
||||
continue
|
||||
else:
|
||||
# No buffer, the start of processing
|
||||
curr_op_char = char
|
||||
else:
|
||||
# Append char to buffer
|
||||
buff += char
|
||||
|
||||
# Add last op to dict
|
||||
ops_dict[curr_op_char] += [buff]
|
||||
return ops_dict
|
||||
@@ -1,13 +0,0 @@
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
|
||||
class Action(ABC):
|
||||
@property
|
||||
@abstractmethod
|
||||
def START_CHAR(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def END_CHAR(self):
|
||||
pass
|
||||
@@ -0,0 +1,66 @@
|
||||
import torch
|
||||
from torch.nn import Embedding
|
||||
|
||||
from custom_nodes.ClipStuff.lib.action.base import Action, MultiArgAction
|
||||
from custom_nodes.ClipStuff.lib.actions.action_utils import get_embedding
|
||||
|
||||
|
||||
class DiffAction(MultiArgAction):
|
||||
grammar = 'diff(" arg ("|" arg)* ")"'
|
||||
name = "diff"
|
||||
chars = ["-", "-"]
|
||||
|
||||
def token_length(self) -> int:
|
||||
# Sum adds to the embeddings of the base segment, so the length is the length of the base segment
|
||||
return sum(seg_or_action.token_length() for seg_or_action in self.base_segment)
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
|
||||
# Calculate the embeddings for the base segment
|
||||
all_base_embeddings = [
|
||||
get_embedding(seg_or_action, embedding_module)
|
||||
for seg_or_action in self.base_segment
|
||||
]
|
||||
|
||||
result = torch.cat(all_base_embeddings, dim=1)
|
||||
|
||||
for arg in self.args:
|
||||
all_arg_embeddings = [
|
||||
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
|
||||
]
|
||||
|
||||
arg_embedding = torch.cat(all_arg_embeddings, dim=1)
|
||||
|
||||
if (
|
||||
arg_embedding.shape[-2] == 1
|
||||
or result.shape[-2] == arg_embedding.shape[-2]
|
||||
):
|
||||
result = result.sub(arg_embedding)
|
||||
else:
|
||||
print(
|
||||
"WARNING: shape mismatch when trying to apply sum, arg will be averaged"
|
||||
)
|
||||
result = result.sub(torch.mean(arg_embedding, dim=1, keepdim=True))
|
||||
|
||||
return result
|
||||
|
||||
# def __repr__(self):
|
||||
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
|
||||
def __repr__(self) -> str:
|
||||
return f"sum({', '.join(map(str, self.args))})"
|
||||
|
||||
def depth_repr(self, depth=1):
|
||||
out = "NudgeAction(\n"
|
||||
if isinstance(self.base_segment, Action):
|
||||
base_segment_repr = self.base_segment.depth_repr(depth + 1)
|
||||
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
|
||||
else:
|
||||
out += "\t" * depth + f"base_segment={self.base_segment.depth_repr()},\n"
|
||||
|
||||
if isinstance(self.args, Action):
|
||||
target_repr = self.args.depth_repr(depth + 1)
|
||||
out += "\t" * depth + f"target={target_repr},\n"
|
||||
else:
|
||||
out += "\t" * depth + f"target={self.args.depth_repr()},\n"
|
||||
out += "\t" * depth + f"weight={self.weight},\n"
|
||||
out += "\t" * (depth - 1) + ")"
|
||||
return out
|
||||
@@ -1,17 +0,0 @@
|
||||
from custom_nodes.ClipStuff.lib.actions import ALL_START_CHARS, ALL_END_CHARS
|
||||
from custom_nodes.ClipStuff.lib.actions.base import Action
|
||||
|
||||
|
||||
def is_action_segment(action_class: Action.__class__, segment: str):
|
||||
if not issubclass(action_class, Action):
|
||||
raise Exception(
|
||||
f"action_class must be a subclass of Action, got {action_class}"
|
||||
)
|
||||
|
||||
return (
|
||||
segment[0] == action_class.START_CHAR and segment[-1] == action_class.END_CHAR
|
||||
)
|
||||
|
||||
|
||||
def is_any_action_segment(segment: str):
|
||||
return segment[0] in ALL_START_CHARS and segment[-1] in ALL_END_CHARS
|
||||
@@ -0,0 +1,32 @@
|
||||
import torch
|
||||
from torch.nn import Embedding
|
||||
|
||||
from custom_nodes.ClipStuff.lib.action.base import Action, SingleArgAction
|
||||
|
||||
|
||||
class NegAction(SingleArgAction):
|
||||
grammar = 'neg(" arg+ ")"'
|
||||
name = "neg"
|
||||
chars = ["[", "]"]
|
||||
|
||||
def token_length(self) -> int:
|
||||
"""
|
||||
Neg negates the embeddings of the base segment, so the length is the length of the base segment
|
||||
:return:
|
||||
"""
|
||||
total_length = 0
|
||||
for seg_or_action in self.arg:
|
||||
total_length += seg_or_action.token_length()
|
||||
|
||||
return total_length
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
|
||||
all_embeddings = []
|
||||
for seg_or_action in self.arg:
|
||||
if isinstance(seg_or_action, Action):
|
||||
all_embeddings.append(seg_or_action.get_result(embedding_module))
|
||||
else:
|
||||
all_embeddings.append(seg_or_action.get_embeddings(embedding_module))
|
||||
|
||||
target_embeddings = torch.cat(all_embeddings, dim=1)
|
||||
return target_embeddings * -1
|
||||
@@ -0,0 +1,37 @@
|
||||
import torch
|
||||
from torch.nn import Embedding
|
||||
|
||||
from custom_nodes.ClipStuff.lib.action.base import (
|
||||
Action,
|
||||
SingleArgAction,
|
||||
)
|
||||
|
||||
|
||||
class NormAction(SingleArgAction):
|
||||
grammar = 'norm(" arg+ ")"'
|
||||
name = "norm"
|
||||
chars = None
|
||||
|
||||
def token_length(self) -> int:
|
||||
"""
|
||||
Norm normalizes the embeddings of the base segment, so the length is the length of the base segment
|
||||
:return:
|
||||
"""
|
||||
total_length = 0
|
||||
for seg_or_action in self.arg:
|
||||
total_length += seg_or_action.token_length()
|
||||
|
||||
return total_length
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
|
||||
all_embeddings = []
|
||||
for seg_or_action in self.arg:
|
||||
if isinstance(seg_or_action, Action):
|
||||
target_embeddings = seg_or_action.get_result(embedding_module)
|
||||
else:
|
||||
target_embeddings = seg_or_action.get_embeddings(embedding_module)
|
||||
all_embeddings.append(target_embeddings)
|
||||
|
||||
target_embeddings = torch.cat(all_embeddings, dim=1)
|
||||
return torch.div(target_embeddings, torch.norm(target_embeddings, dim=-1, keepdim=True))
|
||||
|
||||
@@ -1,13 +0,0 @@
|
||||
from typing import Optional
|
||||
|
||||
from custom_nodes.ClipStuff.lib.actions.base import Action
|
||||
|
||||
|
||||
class NudgeAction(Action):
|
||||
START_CHAR = "["
|
||||
END_CHAR = "]"
|
||||
|
||||
def __init__(self, base_segment=None, weight: Optional[float] = None, target=None):
|
||||
self.base_segment = base_segment
|
||||
self.weight = weight
|
||||
self.target = target
|
||||
@@ -0,0 +1,96 @@
|
||||
from typing import Union
|
||||
|
||||
import torch
|
||||
from torch.nn import Embedding
|
||||
|
||||
from custom_nodes.ClipStuff.lib.actions.action_utils import get_embedding
|
||||
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
|
||||
from custom_nodes.ClipStuff.lib.action.base import Action
|
||||
|
||||
|
||||
class SumAction(Action):
|
||||
grammar = 'sum(" arg ("|" arg)* ")"'
|
||||
name = "sum"
|
||||
chars = ["+", "+"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_segment: list[PromptSegment | Action],
|
||||
args: list[list[Union[PromptSegment, Action]]],
|
||||
):
|
||||
self.base_segment = base_segment
|
||||
self.args = args
|
||||
|
||||
def token_length(self) -> int:
|
||||
# Sum adds to the embeddings of the base segment, so the length is the length of the base segment
|
||||
return sum(seg_or_action.token_length() for seg_or_action in self.base_segment)
|
||||
|
||||
def get_all_segments(self) -> list[PromptSegment]:
|
||||
segments = []
|
||||
for seg_or_action in self.base_segment:
|
||||
if isinstance(seg_or_action, Action):
|
||||
segments.extend(seg_or_action.get_all_segments())
|
||||
else:
|
||||
segments.append(seg_or_action)
|
||||
|
||||
for arg in self.args:
|
||||
for seg_or_action in arg:
|
||||
if isinstance(seg_or_action, Action):
|
||||
segments.extend(seg_or_action.get_all_segments())
|
||||
else:
|
||||
segments.append(seg_or_action)
|
||||
|
||||
return segments
|
||||
|
||||
def get_result(self, embedding_module: Embedding) -> torch.Tensor:
|
||||
# Calculate the embeddings for the base segment
|
||||
all_base_embeddings = [
|
||||
get_embedding(seg_or_action, embedding_module)
|
||||
for seg_or_action in self.base_segment
|
||||
]
|
||||
|
||||
result = torch.cat(all_base_embeddings, dim=1)
|
||||
|
||||
for arg in self.args:
|
||||
all_arg_embeddings = [
|
||||
get_embedding(seg_or_action, embedding_module) for seg_or_action in arg
|
||||
]
|
||||
|
||||
arg_embedding = torch.cat(all_arg_embeddings, dim=1)
|
||||
|
||||
if (
|
||||
arg_embedding.shape[-2] == 1
|
||||
or result.shape[-2] == arg_embedding.shape[-2]
|
||||
):
|
||||
result = result.add(arg_embedding)
|
||||
else:
|
||||
print(
|
||||
"WARNING: shape mismatch when trying to apply sum, arg will be averaged"
|
||||
)
|
||||
result = result.add(
|
||||
torch.mean(arg_embedding, dim=1, keepdim=True)
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
# def __repr__(self):
|
||||
# return f"sum(\n\tbase_segment={self.base_segment},\n\targs={self.args}\n)"
|
||||
def __repr__(self) -> str:
|
||||
return f"sum({', '.join(map(str, self.args))})"
|
||||
|
||||
def depth_repr(self, depth=1):
|
||||
out = "NudgeAction(\n"
|
||||
if isinstance(self.base_segment, Action):
|
||||
base_segment_repr = self.base_segment.depth_repr(depth + 1)
|
||||
out += "\t" * depth + f"base_segment={base_segment_repr}\n"
|
||||
else:
|
||||
out += "\t" * depth + f"base_segment={self.base_segment.depth_repr()},\n"
|
||||
|
||||
if isinstance(self.args, Action):
|
||||
target_repr = self.args.depth_repr(depth + 1)
|
||||
out += "\t" * depth + f"target={target_repr},\n"
|
||||
else:
|
||||
out += "\t" * depth + f"target={self.args.depth_repr()},\n"
|
||||
out += "\t" * depth + f"weight={self.weight},\n"
|
||||
out += "\t" * (depth - 1) + ")"
|
||||
return out
|
||||
@@ -0,0 +1,6 @@
|
||||
from typing import Union
|
||||
|
||||
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
|
||||
from custom_nodes.ClipStuff.lib.action.base import Action
|
||||
|
||||
SegOrAction = Union[PromptSegment, Action]
|
||||
@@ -0,0 +1,6 @@
|
||||
from custom_nodes.ClipStuff.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())
|
||||
+76
-58
@@ -6,11 +6,14 @@ from transformers import CLIPTextConfig, modeling_utils
|
||||
|
||||
from comfy import model_management
|
||||
import comfy.ops
|
||||
from comfy.sd import CLIP
|
||||
from custom_nodes.ClipStuff.lib.fun_clip_stuff import MyCLIPTextModel
|
||||
from custom_nodes.ClipStuff.lib.tokenizer import TokenDict
|
||||
from custom_nodes.ClipStuff.lib.action.base import Action
|
||||
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
|
||||
from custom_nodes.ClipStuff.lib.fun_clip_stuff import PromptLangTextModel
|
||||
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
class SD1FunClipModel(torch.nn.Module):
|
||||
|
||||
# Methods with no comment can be assumed to be the same as comfy.sd1_clip.SD1ClipModel
|
||||
class PromptLangClipModel(torch.nn.Module):
|
||||
"""Uses the CLIP transformer encoder for text (from huggingface)"""
|
||||
LAYERS = [
|
||||
"last",
|
||||
@@ -25,15 +28,20 @@ class SD1FunClipModel(torch.nn.Module):
|
||||
assert layer in self.LAYERS
|
||||
self.num_layers = 12
|
||||
if textmodel_path is not None:
|
||||
self.transformer = MyCLIPTextModel.from_pretrained(textmodel_path)
|
||||
# Our transformer
|
||||
self.transformer = PromptLangTextModel.from_pretrained(textmodel_path)
|
||||
else:
|
||||
if textmodel_json_config is None:
|
||||
# TODO: Maybe re-use clip config?
|
||||
# Config could come from cond_stage_model.transformer.config
|
||||
# Copied clip_config
|
||||
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "clip_config.json")
|
||||
config = CLIPTextConfig.from_json_file(textmodel_json_config)
|
||||
self.num_layers = config.num_hidden_layers
|
||||
with comfy.ops.use_comfy_ops():
|
||||
with modeling_utils.no_init_weights():
|
||||
self.transformer = MyCLIPTextModel(config)
|
||||
# Our transformer
|
||||
self.transformer = PromptLangTextModel(config)
|
||||
|
||||
self.max_length = max_length
|
||||
if freeze:
|
||||
@@ -66,51 +74,72 @@ class SD1FunClipModel(torch.nn.Module):
|
||||
self.layer = self.layer_default[0]
|
||||
self.layer_idx = self.layer_default[1]
|
||||
|
||||
def set_up_textual_embeddings(self, tokens: list[list[tuple[TokenDict]]], current_embeds):
|
||||
out_tokens = []
|
||||
# Completely changed to support Segments and actions
|
||||
def set_up_textual_embeddings(self, tokens: list[list[SegOrAction]], current_embeds):
|
||||
next_new_token = token_dict_size = current_embeds.weight.shape[0] - 1
|
||||
embedding_weights = []
|
||||
|
||||
# For each batch
|
||||
for batch in tokens:
|
||||
tokens_temp = []
|
||||
for tokenDict in batch:
|
||||
y = tokenDict[0].token_id
|
||||
if isinstance(y, int):
|
||||
if y == token_dict_size: # EOS token
|
||||
y = -1
|
||||
tokens_temp += [y]
|
||||
for seg_or_action in batch:
|
||||
if isinstance(seg_or_action, Action):
|
||||
segments = seg_or_action.get_all_segments()
|
||||
else:
|
||||
if y.shape[0] == current_embeds.weight.shape[1]:
|
||||
embedding_weights += [y]
|
||||
tokens_temp += [next_new_token]
|
||||
next_new_token += 1
|
||||
else:
|
||||
print("WARNING: shape mismatch when trying to apply embedding, embedding will be ignored",
|
||||
y.shape[0], current_embeds.weight.shape[1])
|
||||
while len(tokens_temp) < len(batch):
|
||||
tokens_temp += [self.empty_tokens[0][-1]]
|
||||
out_tokens += [tokens_temp]
|
||||
segments = [seg_or_action]
|
||||
|
||||
for segment in segments:
|
||||
tokens_temp = []
|
||||
segment_length = segment.token_length()
|
||||
for tid_or_tensor in segment.tokens:
|
||||
if isinstance(tid_or_tensor, int):
|
||||
if tid_or_tensor == token_dict_size: # Is EOS token
|
||||
tid_or_tensor = -1 # Set to -1 so that it can be replaced with the EOS token later
|
||||
tokens_temp += [tid_or_tensor]
|
||||
else:
|
||||
if tid_or_tensor.shape[0] == current_embeds.weight.shape[1]:
|
||||
embedding_weights += [tid_or_tensor]
|
||||
tokens_temp += [next_new_token]
|
||||
next_new_token += 1
|
||||
else:
|
||||
print("WARNING: shape mismatch when trying to apply embedding, embedding will be ignored",
|
||||
tid_or_tensor.shape[0], current_embeds.weight.shape[1])
|
||||
if len(tokens_temp) < segment_length:
|
||||
# Pretty sure this is only needed if the embedding is not the same size as the CLIP embedding
|
||||
print("WARNING: segment length mismatch, padding with EOS token")
|
||||
tokens_temp.extend([self.empty_tokens[0][-1] * (segment_length - len(tokens_temp))])
|
||||
segment.tokens = tokens_temp
|
||||
|
||||
n = token_dict_size
|
||||
if len(embedding_weights) > 0:
|
||||
# Create new embedding, with size of current embedding + number of new embeddings
|
||||
new_embedding = torch.nn.Embedding(next_new_token + 1, current_embeds.weight.shape[1],
|
||||
device=current_embeds.weight.device, dtype=current_embeds.weight.dtype)
|
||||
# Copy current embedding weights to new embedding
|
||||
new_embedding.weight[:token_dict_size] = current_embeds.weight[:-1]
|
||||
# Add new embeddings
|
||||
for embed in embedding_weights:
|
||||
new_embedding.weight[n] = embed
|
||||
n += 1
|
||||
|
||||
# Set re-add the EOS token
|
||||
new_embedding.weight[n] = current_embeds.weight[-1] # EOS embedding
|
||||
self.transformer.set_input_embeddings(new_embedding)
|
||||
|
||||
for i, out_batch in enumerate(out_tokens):
|
||||
for tokenIdx in range(len(out_batch)):
|
||||
if out_batch[tokenIdx] == -1:
|
||||
tokens[i][tokenIdx][0].token_id = n # The EOS token should always be the largest one
|
||||
else:
|
||||
tokens[i][tokenIdx][0].token_id = out_batch[tokenIdx]
|
||||
|
||||
# return processed_tokens
|
||||
def forward(self, tokens, **kwargs):
|
||||
for batch in tokens:
|
||||
for seg_or_action in batch:
|
||||
if isinstance(seg_or_action, Action):
|
||||
segments = seg_or_action.get_all_segments()
|
||||
else:
|
||||
segments = [seg_or_action]
|
||||
|
||||
for segment in segments:
|
||||
for tokenIdx in range(len(segment.tokens)):
|
||||
if segment.tokens[tokenIdx] == -1:
|
||||
segment.tokens[tokenIdx] = n
|
||||
|
||||
# Support our set_up_textual_embeddings which modifies the input embeddings
|
||||
def forward(self, tokens):
|
||||
backup_embeds = self.transformer.get_input_embeddings()
|
||||
device = backup_embeds.weight.device
|
||||
self.set_up_textual_embeddings(tokens, backup_embeds)
|
||||
@@ -121,16 +150,8 @@ class SD1FunClipModel(torch.nn.Module):
|
||||
else:
|
||||
precision_scope = contextlib.nullcontext
|
||||
|
||||
|
||||
if (kwargs.get("position_ids", None) is not None):
|
||||
position_ids = torch.LongTensor(kwargs["position_ids"]).to(device)
|
||||
else:
|
||||
position_ids = None
|
||||
|
||||
|
||||
with precision_scope(model_management.get_autocast_device(device)):
|
||||
outputs = self.transformer(input_ids=tokens, output_hidden_states=self.layer == "hidden",
|
||||
position_ids=position_ids)
|
||||
outputs = self.transformer(input_ids=tokens, output_hidden_states=self.layer == "hidden")
|
||||
self.transformer.set_input_embeddings(backup_embeds)
|
||||
|
||||
if self.layer == "last":
|
||||
@@ -147,23 +168,20 @@ class SD1FunClipModel(torch.nn.Module):
|
||||
pooled_output = pooled_output.to(self.text_projection.device) @ self.text_projection
|
||||
return z.float(), pooled_output.float()
|
||||
|
||||
def encode(self, tokens, **kwargs):
|
||||
return self(tokens, **kwargs)
|
||||
def encode(self, tokens):
|
||||
return self(tokens)
|
||||
|
||||
def load_sd(self, sd):
|
||||
return self.transformer.load_state_dict(sd, strict=False)
|
||||
|
||||
def encode_token_weights(self, token_dicts: list[list[tuple[TokenDict]]], **kwargs):
|
||||
to_encode = [list(
|
||||
map(
|
||||
lambda id: (TokenDict(token_id=id, weight=1.0, nudge_id=None),),
|
||||
self.empty_tokens[0]
|
||||
)
|
||||
)]
|
||||
for x in token_dicts:
|
||||
to_encode.append(x)
|
||||
# Changed from comfy.sd1_clip.ClipTokenWeightEncoder
|
||||
# Changed to use PromptSegments
|
||||
def encode_token_weights(self, prompt_segments: list[list[SegOrAction]]):
|
||||
to_encode = [[PromptSegment(text="_Empty Batch_", tokens=self.empty_tokens[0])]]
|
||||
for batch in prompt_segments:
|
||||
to_encode.append(batch)
|
||||
|
||||
out, pooled = self.encode(to_encode, **kwargs)
|
||||
out, pooled = self.encode(to_encode)
|
||||
z_empty = out[0:1]
|
||||
if pooled.shape[0] > 1:
|
||||
first_pooled = pooled[1:2]
|
||||
@@ -173,10 +191,10 @@ class SD1FunClipModel(torch.nn.Module):
|
||||
output = []
|
||||
for k in range(1, out.shape[0]):
|
||||
z = out[k:k + 1]
|
||||
for i in range(len(z)):
|
||||
for j in range(len(z[i])):
|
||||
weight = token_dicts[k - 1][j][0].weight
|
||||
z[i][j] = (z[i][j] - z_empty[0][j]) * weight + z_empty[0][j]
|
||||
# for i in range(len(z)):
|
||||
# for j in range(len(z[i])):
|
||||
# weight = token_dicts[k - 1][j][0].weight
|
||||
# z[i][j] = (z[i][j] - z_empty[0][j]) * weight + z_empty[0][j]
|
||||
output.append(z)
|
||||
|
||||
if (len(output) == 0):
|
||||
|
||||
+52
-63
@@ -1,13 +1,13 @@
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch import device
|
||||
from transformers import CLIPTextConfig
|
||||
from transformers.modeling_outputs import BaseModelOutputWithPooling
|
||||
from transformers.models.clip.modeling_clip import _expand_mask, CLIPTextEmbeddings, CLIPTextTransformer, \
|
||||
CLIPTextModel
|
||||
|
||||
from custom_nodes.ClipStuff.lib.tokenizer import TokenDict
|
||||
from custom_nodes.ClipStuff.lib.action.base import Action
|
||||
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
|
||||
|
||||
def slerp(val, low, high):
|
||||
low = low.unsqueeze(0)
|
||||
@@ -19,68 +19,57 @@ def slerp(val, low, high):
|
||||
res = (torch.sin((1.0-val)*omega)/so).unsqueeze(1)*low + (torch.sin(val*omega)/so).unsqueeze(1) * high
|
||||
return res
|
||||
|
||||
class MyCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||
class PromptLangCLIPTextEmbeddings(CLIPTextEmbeddings):
|
||||
def __init__(self, config: CLIPTextConfig):
|
||||
super().__init__(config)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_dicts: Optional[list[list[tuple[TokenDict]]]] = None,
|
||||
input_dicts: Optional[list[list[SegOrAction]]] = None,
|
||||
input_ids: Optional[torch.LongTensor] = None,
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||
) -> torch.Tensor:
|
||||
input_ids = [
|
||||
[
|
||||
tokenDict[0].token_id for tokenDict in batch
|
||||
] for batch in input_dicts
|
||||
]
|
||||
tokens = torch.LongTensor(input_ids).to(torch.device('cpu'))
|
||||
input_shape = tokens.size()
|
||||
input_ids = tokens.view(-1, input_shape[-1])
|
||||
|
||||
seq_length = input_ids.shape[-1] if input_ids is not None else inputs_embeds.shape[-2]
|
||||
if input_dicts is None:
|
||||
raise ValueError("You have to specify input_dicts")
|
||||
|
||||
batches = []
|
||||
for batch_idx, batch in enumerate(input_dicts):
|
||||
results = []
|
||||
for seg_or_action in batch:
|
||||
if isinstance(seg_or_action, Action):
|
||||
results.append(seg_or_action.get_result(self.token_embedding))
|
||||
else:
|
||||
results.append(seg_or_action.get_embeddings(self.token_embedding))
|
||||
batches.append(results)
|
||||
|
||||
seq_length = batches[0][0].shape[-2]
|
||||
|
||||
if position_ids is None:
|
||||
position_ids = self.position_ids[:, :seq_length]
|
||||
|
||||
if inputs_embeds is None:
|
||||
inputs_embeds = self.token_embedding(input_ids)
|
||||
|
||||
for batch_idx, batch in enumerate(input_dicts):
|
||||
for token_idx, token in enumerate(batch):
|
||||
if token[0].nudge_id is not None:
|
||||
nudged_embed = inputs_embeds[batch_idx, token_idx][:] + self.token_embedding(torch.LongTensor([token[0].nudge_id]).to(torch.device('cpu')))[0]
|
||||
if token[0].nudge_index_start is not None and token[0].nudge_index_stop is not None:
|
||||
nudge_start = token[0].nudge_index_start
|
||||
nudge_end = token[0].nudge_index_stop
|
||||
else:
|
||||
nudge_start = 0
|
||||
nudge_end = 768
|
||||
inputs_embeds[batch_idx, token_idx][nudge_start:nudge_end] = (slerp(token[0].nudge_weight, inputs_embeds[batch_idx, token_idx][:], nudged_embed)[0][nudge_start:nudge_end])
|
||||
elif token[0].arith_ops is not None:
|
||||
for op, id_list in token[0].arith_ops.items():
|
||||
if op == '+':
|
||||
for this_id in id_list:
|
||||
inputs_embeds[batch_idx, token_idx] += self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0]
|
||||
elif op == '-':
|
||||
for this_id in id_list:
|
||||
inputs_embeds[batch_idx, token_idx] -= self.token_embedding(torch.LongTensor([this_id]).to(torch.device('cpu')))[0]
|
||||
embeds = []
|
||||
for batch in batches:
|
||||
if len(batch) == 1:
|
||||
embeds.append(batch[0])
|
||||
else:
|
||||
embeds.append(torch.cat(batch, dim=-2))
|
||||
|
||||
position_embeddings = self.position_embedding(position_ids)
|
||||
embeddings = inputs_embeds + position_embeddings
|
||||
embeddings = torch.cat(embeds, dim=0) + position_embeddings
|
||||
|
||||
return embeddings, input_ids, input_shape
|
||||
return embeddings
|
||||
|
||||
|
||||
class MyCLIPTextTransformer(CLIPTextTransformer):
|
||||
class PrompLangCLIPTextTransformer(CLIPTextTransformer):
|
||||
def __init__(self, config: CLIPTextConfig):
|
||||
super().__init__(config)
|
||||
self.embeddings = MyCLIPTextEmbeddings(config)
|
||||
self.embeddings = PromptLangCLIPTextEmbeddings(config)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[list[list[tuple[TokenDict]]]] = None,
|
||||
input_ids: Optional[list[list[SegOrAction]]] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
output_attentions: Optional[bool] = None,
|
||||
@@ -103,9 +92,12 @@ class MyCLIPTextTransformer(CLIPTextTransformer):
|
||||
# input_shape = input_ids.size()
|
||||
# input_ids = input_ids.view(-1, input_shape[-1])
|
||||
|
||||
hidden_states, input_ids, input_shape = self.embeddings(input_dicts=input_ids)
|
||||
hidden_states = self.embeddings(input_dicts=input_ids)
|
||||
|
||||
bsz, seq_len = input_shape
|
||||
bsz = len(input_ids)
|
||||
# TODO: Properly gather this
|
||||
seq_len = 77
|
||||
# bsz, seq_len = input_shape
|
||||
# CLIP's text model uses causal mask, prepare it here.
|
||||
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
|
||||
causal_attention_mask = self._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(
|
||||
@@ -128,12 +120,25 @@ class MyCLIPTextTransformer(CLIPTextTransformer):
|
||||
last_hidden_state = encoder_outputs[0]
|
||||
last_hidden_state = self.final_layer_norm(last_hidden_state)
|
||||
|
||||
|
||||
# Hacky way to get idx of first EOT token
|
||||
eot_idx = [1]
|
||||
for batch in input_ids[1:]:
|
||||
idx = 0
|
||||
for seg_or_action in batch:
|
||||
if isinstance(seg_or_action, Action):
|
||||
idx += seg_or_action.token_length()
|
||||
else:
|
||||
if seg_or_action.text == '__PAD__':
|
||||
break
|
||||
eot_idx.append(idx)
|
||||
# text_embeds.shape = [batch_size, sequence_length, transformer.width]
|
||||
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
||||
# casting to torch.int for onnx compatibility: argmax doesn't support int64 inputs with opset 14
|
||||
# TODO: Get the index of the first EOT token
|
||||
pooled_output = last_hidden_state[
|
||||
torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device),
|
||||
input_ids.to(dtype=torch.int, device=last_hidden_state.device).argmax(dim=-1),
|
||||
eot_idx
|
||||
]
|
||||
|
||||
if not return_dict:
|
||||
@@ -147,37 +152,21 @@ class MyCLIPTextTransformer(CLIPTextTransformer):
|
||||
)
|
||||
|
||||
|
||||
class MyCLIPTextModel(CLIPTextModel):
|
||||
# This is necessary to pass the PromptLangCLIPTextTransformer
|
||||
class PromptLangTextModel(CLIPTextModel):
|
||||
def __init__(self, config: CLIPTextConfig):
|
||||
super().__init__(config)
|
||||
self.text_model = MyCLIPTextTransformer(config)
|
||||
self.text_model = PrompLangCLIPTextTransformer(config)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[list[list[tuple[TokenDict]]]] = None,
|
||||
input_ids: Optional[list[list[SegOrAction]]] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
position_ids: Optional[torch.Tensor] = None,
|
||||
output_attentions: Optional[bool] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
return_dict: Optional[bool] = None,
|
||||
) -> Union[Tuple, BaseModelOutputWithPooling]:
|
||||
r"""
|
||||
Returns:
|
||||
|
||||
Examples:
|
||||
|
||||
```python
|
||||
>>> from transformers import AutoTokenizer, CLIPTextModel
|
||||
|
||||
>>> model = CLIPTextModel.from_pretrained("openai/clip-vit-base-patch32")
|
||||
>>> tokenizer = AutoTokenizer.from_pretrained("openai/clip-vit-base-patch32")
|
||||
|
||||
>>> inputs = tokenizer(["a photo of a cat", "a photo of a dog"], padding=True, return_tensors="pt")
|
||||
|
||||
>>> outputs = model(**inputs)
|
||||
>>> last_hidden_state = outputs.last_hidden_state
|
||||
>>> pooled_output = outputs.pooler_output # pooled (EOS token) states
|
||||
```"""
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
|
||||
return self.text_model(
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
from lark import Lark
|
||||
|
||||
from .grammar import grammar
|
||||
|
||||
PromptParser = Lark(grammar, start="start", parser="earley")
|
||||
@@ -0,0 +1,28 @@
|
||||
grammar = """
|
||||
?start: item+
|
||||
|
||||
item: embedding
|
||||
| WORD
|
||||
| function
|
||||
| QUOTED_STRING
|
||||
|
||||
function: sum_function
|
||||
| neg_function
|
||||
| norm_function
|
||||
| diff_function
|
||||
|
||||
sum_function: "sum(" arg ("|" arg)* ")"
|
||||
neg_function: "neg(" arg ")"
|
||||
norm_function: "norm(" arg ")"
|
||||
diff_function: "diff(" arg ("|" arg)* ")"
|
||||
|
||||
arg: item+
|
||||
|
||||
embedding: "embedding:" WORD
|
||||
|
||||
WORD: /[A-Za-z0-9,_-]+/
|
||||
QUOTED_STRING: /"([^"\\\]*(\\\.[^"\\\]*)*)"|'([^'\\\]*(\\\.[^'\\\]*)*)'/
|
||||
|
||||
%import common.WS
|
||||
%ignore WS
|
||||
"""
|
||||
@@ -0,0 +1,31 @@
|
||||
from typing import Union
|
||||
|
||||
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(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
|
||||
@@ -0,0 +1,61 @@
|
||||
from lark import Transformer, Token
|
||||
|
||||
from comfy.sd1_clip import SD1Tokenizer
|
||||
from custom_nodes.ClipStuff.lib.action.base import Action
|
||||
from custom_nodes.ClipStuff.lib.actions.diff import DiffAction
|
||||
from custom_nodes.ClipStuff.lib.parser.utils import build_prompt_segment
|
||||
from custom_nodes.ClipStuff.lib.actions.neg import NegAction
|
||||
from custom_nodes.ClipStuff.lib.actions.norm import NormAction
|
||||
from custom_nodes.ClipStuff.lib.actions.sum import SumAction
|
||||
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
|
||||
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 function(self, items):
|
||||
for item in items:
|
||||
if item.data == 'sum_function':
|
||||
return SumAction(item.children[0][:], item.children[1:][:])
|
||||
elif item.data == 'neg_function':
|
||||
return NegAction(item.children[0])
|
||||
elif item.data == 'norm_function':
|
||||
return NormAction(item.children[0])
|
||||
elif item.data == 'diff_function':
|
||||
return DiffAction(item.children[0][:], item.children[1:][:])
|
||||
else:
|
||||
raise Exception("Unknown function type: " + str(item.data))
|
||||
@@ -0,0 +1,38 @@
|
||||
from lark import Token
|
||||
|
||||
from comfy.sd1_clip import SD1Tokenizer
|
||||
from custom_nodes.ClipStuff.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)
|
||||
+45
-193
@@ -1,214 +1,66 @@
|
||||
from typing import Union
|
||||
from lark import Tree
|
||||
|
||||
from comfy.sd1_clip import SD1Tokenizer
|
||||
from custom_nodes.ClipStuff.lib.actions import (
|
||||
NudgeAction,
|
||||
ArithAction,
|
||||
ALL_START_CHARS,
|
||||
ALL_END_CHARS,
|
||||
)
|
||||
from custom_nodes.ClipStuff.lib.actions.lib import (
|
||||
is_any_action_segment,
|
||||
is_action_segment,
|
||||
)
|
||||
from custom_nodes.ClipStuff.lib.actions.types import SegOrAction
|
||||
|
||||
from custom_nodes.ClipStuff.lib.parser import PromptParser
|
||||
from custom_nodes.ClipStuff.lib.parser.transformer import PromptTransformer
|
||||
from custom_nodes.ClipStuff.lib.parser.prompt_segment import PromptSegment
|
||||
|
||||
def parse_special_tokens(string):
|
||||
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_token_actions(string) -> list[Union[str, NudgeAction, ArithAction]]:
|
||||
out: list[Union[str, NudgeAction, ArithAction]] = []
|
||||
for prompt_segment in parse_special_tokens(string):
|
||||
if prompt_segment == "":
|
||||
continue
|
||||
|
||||
if not is_any_action_segment(prompt_segment):
|
||||
out += [prompt_segment]
|
||||
continue
|
||||
|
||||
is_nudge = is_action_segment(NudgeAction, prompt_segment)
|
||||
is_arith = is_action_segment(ArithAction, prompt_segment)
|
||||
|
||||
prompt_segment = prompt_segment[1:-1]
|
||||
word_sep_idx = prompt_segment.find(":")
|
||||
|
||||
# No word seperator, add whole segment
|
||||
if word_sep_idx < 0:
|
||||
out += [prompt_segment]
|
||||
continue
|
||||
|
||||
base_segment = prompt_segment[:word_sep_idx]
|
||||
|
||||
if is_nudge:
|
||||
trailing_segment = prompt_segment[word_sep_idx + 1 :]
|
||||
|
||||
weight_sep_idx = trailing_segment.find(":")
|
||||
# Has a weight(base_word:nudge_to:1.4)
|
||||
if weight_sep_idx >= 0:
|
||||
[nudge_to, weight] = trailing_segment.split(":")
|
||||
weight = float(weight)
|
||||
else:
|
||||
# No weight(base_word:trailing_segment)
|
||||
nudge_to = trailing_segment
|
||||
weight = None
|
||||
|
||||
out += [NudgeAction(base_segment, weight, nudge_to)]
|
||||
elif is_arith:
|
||||
arith_op_string = prompt_segment[word_sep_idx + 1 :]
|
||||
out += [ArithAction(base_segment, arith_op_string)]
|
||||
|
||||
return out
|
||||
|
||||
|
||||
class TokenDict:
|
||||
def __init__(self,
|
||||
token_id: int,
|
||||
weight: float = None,
|
||||
nudge_id=None, nudge_weight=None, nudge_start: int = None, nudge_end: int = None,
|
||||
arith_ops: dict[str, list[str]] = None):
|
||||
if weight is None:
|
||||
self.weight = 1.0
|
||||
else:
|
||||
self.weight = weight
|
||||
|
||||
self.token_id = token_id
|
||||
self.nudge_id = nudge_id
|
||||
self.nudge_weight = nudge_weight
|
||||
self.nudge_index_start = nudge_start
|
||||
self.nudge_index_stop = nudge_end
|
||||
|
||||
self.arith_ops = arith_ops
|
||||
|
||||
|
||||
class MyTokenizer(SD1Tokenizer):
|
||||
class PromptLangTokenizer(SD1Tokenizer):
|
||||
def __init__(self, tokenizer_path=None, max_length=77, pad_with_end=True, embedding_directory=None, embedding_size=768, embedding_key='clip_l', special_tokens=None):
|
||||
super().__init__(tokenizer_path, max_length, pad_with_end, embedding_directory, embedding_size, embedding_key)
|
||||
|
||||
"""
|
||||
:return: list of tuples (tokenDict, word_id?)
|
||||
Doesn't actually tokenize...
|
||||
Returns batches of segments and actions
|
||||
:return: List of list(batches) of segments and actions
|
||||
"""
|
||||
def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs):
|
||||
def tokenize_with_weights(self, text:str, return_word_ids=False, **kwargs) -> list[list[SegOrAction]]:
|
||||
if self.pad_with_end:
|
||||
pad_token = self.end_token
|
||||
else:
|
||||
pad_token = 0
|
||||
|
||||
parsed_actions = parse_token_actions(text)
|
||||
|
||||
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
|
||||
tokens: list[list[TokenDict]] = []
|
||||
|
||||
for action in parsed_actions:
|
||||
nudge_weight = None
|
||||
nudge_to_id = None
|
||||
arith_ops = None
|
||||
if isinstance(action, str):
|
||||
token_segment = action
|
||||
elif isinstance(action, NudgeAction):
|
||||
token_segment = action.base_segment
|
||||
nudge_to_id = self.tokenizer(action.target)["input_ids"][1:-1][0]
|
||||
|
||||
nudge_weight = action.weight
|
||||
if nudge_weight is None:
|
||||
nudge_weight = 0.5
|
||||
elif isinstance(action, ArithAction):
|
||||
token_segment = action.base_segment
|
||||
arith_ops = action.ops
|
||||
for op in arith_ops:
|
||||
arith_ops[op] = [self.tokenizer(word)["input_ids"][1:-1][0] for word in arith_ops[op]]
|
||||
else:
|
||||
raise Exception(f"Unexpected action type: {type(action)}")
|
||||
|
||||
to_tokenize = token_segment.split(' ')
|
||||
to_tokenize = [x for x in to_tokenize if x != ""]
|
||||
|
||||
for word in to_tokenize:
|
||||
# if we find an embedding, deal with the embedding
|
||||
if word.startswith(self.embedding_identifier) and self.embedding_directory is not None:
|
||||
embedding_name = word[len(self.embedding_identifier):].strip('\n')
|
||||
embed, leftover = self._try_get_embedding(embedding_name)
|
||||
if embed is None:
|
||||
print(f"warning, embedding:{embedding_name} does not exist, ignoring")
|
||||
else:
|
||||
if len(embed.shape) == 1:
|
||||
tokens.append([TokenDict(token_id=embed)])
|
||||
else:
|
||||
tokens.append([
|
||||
TokenDict(token_id=embed[x])
|
||||
for x in range(embed.shape[0])
|
||||
])
|
||||
# if we accidentally have leftover text, continue parsing using leftover, else move on to next word
|
||||
if leftover != "":
|
||||
word = leftover
|
||||
else:
|
||||
continue
|
||||
# parse word
|
||||
tokens.append([TokenDict(
|
||||
token_id=t,
|
||||
nudge_id=nudge_to_id,
|
||||
nudge_weight=nudge_weight,
|
||||
nudge_start=nudge_start,
|
||||
nudge_end=nudge_end,
|
||||
arith_ops=arith_ops
|
||||
) for t in self.tokenizer(word)["input_ids"][1:-1]])
|
||||
parsed_prompt = PromptParser.parse(text)
|
||||
parsed_actions = PromptTransformer(self).transform(parsed_prompt)
|
||||
|
||||
# reshape token array to CLIP input size
|
||||
batched_tokens = []
|
||||
batch = [(TokenDict(token_id=self.start_token), 0)]
|
||||
batched_tokens.append(batch)
|
||||
for i, t_group in enumerate(tokens):
|
||||
batched_segments = []
|
||||
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token])]
|
||||
# batched_segments.append(batch)
|
||||
batch_size = 1
|
||||
if isinstance(parsed_actions, Tree):
|
||||
segments_to_process = parsed_actions.children
|
||||
else:
|
||||
segments_to_process = [parsed_actions]
|
||||
for segment in segments_to_process:
|
||||
num_tokens = segment.token_length()
|
||||
# determine if we're going to try and keep the tokens in a single batch
|
||||
is_large = len(t_group) >= self.max_word_length
|
||||
is_large = num_tokens >= self.max_word_length
|
||||
|
||||
while len(t_group) > 0:
|
||||
if len(t_group) + len(batch) > self.max_length - 1:
|
||||
remaining_length = self.max_length - len(batch) - 1
|
||||
# break word in two and add end token
|
||||
if is_large:
|
||||
batch.extend([(tokenDict, i+1) for tokenDict in t_group[:remaining_length]])
|
||||
batch.append((TokenDict(token_id=self.end_token), 0))
|
||||
t_group = t_group[remaining_length:]
|
||||
# add end token and pad
|
||||
else:
|
||||
batch.append((TokenDict(token_id=self.end_token), 0))
|
||||
batch.extend([(TokenDict(token_id=pad_token), 0)] * remaining_length)
|
||||
# start new batch
|
||||
batch = [(TokenDict(token_id=self.start_token), 1.0, 0)]
|
||||
batched_tokens.append(batch)
|
||||
else:
|
||||
batch.extend([(tokenDict, i+1) for tokenDict in t_group])
|
||||
t_group = []
|
||||
# If the segment is too large to fit in a single batch, pad the current batch and start a new one
|
||||
if num_tokens + batch_size > self.max_length - 1:
|
||||
remaining_length = self.max_length - batch_size - 1 # -1 for end token
|
||||
# Pad batch
|
||||
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length - 1))
|
||||
batched_segments.append(batch)
|
||||
|
||||
# fill last batch
|
||||
batch.extend([(TokenDict(token_id=self.end_token), 0)] + [
|
||||
(TokenDict(token_id=pad_token), 0)] * (self.max_length - len(batch) - 1))
|
||||
# start new batch
|
||||
batch = [PromptSegment(text="[SOT]", tokens=[self.start_token]), segment]
|
||||
batch_size = num_tokens + 1 # +1 for start token
|
||||
continue
|
||||
|
||||
if not return_word_ids:
|
||||
batched_tokens = [
|
||||
[
|
||||
(tokenInfo[0],) for tokenInfo in batch
|
||||
] for batch in batched_tokens
|
||||
]
|
||||
# Since the segment fits in the current batch, add it
|
||||
batch.append(segment)
|
||||
batch_size += num_tokens
|
||||
|
||||
return batched_tokens
|
||||
# Pad the last batch
|
||||
remaining_length = self.max_length - batch_size - 1 # -1 for end token
|
||||
batch.append(PromptSegment("__PAD__", [self.end_token] + [pad_token] * remaining_length))
|
||||
batched_segments.append(batch)
|
||||
|
||||
# for batch in batched_segments:
|
||||
# batch_size_info(batch)
|
||||
|
||||
return batched_segments
|
||||
|
||||
@@ -6,9 +6,9 @@ from PIL import Image
|
||||
import folder_paths
|
||||
import comfy.sd
|
||||
import comfy.ops
|
||||
from custom_nodes.ClipStuff.lib.clip_model import SD1FunClipModel
|
||||
from custom_nodes.ClipStuff.lib.clip_model import PromptLangClipModel
|
||||
|
||||
from custom_nodes.ClipStuff.lib.tokenizer import MyTokenizer
|
||||
from custom_nodes.ClipStuff.lib.tokenizer import PromptLangTokenizer
|
||||
|
||||
|
||||
class EmptyClass:
|
||||
@@ -33,60 +33,16 @@ class SpecialClipLoader:
|
||||
def load_clip(source_clip):
|
||||
clip_target = EmptyClass()
|
||||
clip_target.params = {}
|
||||
clip_target.clip = SD1FunClipModel
|
||||
clip_target.tokenizer = MyTokenizer
|
||||
clip_target.clip = PromptLangClipModel
|
||||
clip_target.tokenizer = PromptLangTokenizer
|
||||
|
||||
# TODO: Extract embedding directory from source_clip
|
||||
clip = comfy.sd.CLIP(clip_target, embedding_directory=None)
|
||||
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,)
|
||||
|
||||
|
||||
class KepAdvTextEncode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"multiline": True}),
|
||||
"clip": ("CLIP",),
|
||||
"nudge_start": ("INT", {}),
|
||||
"nudge_end": ("INT", {}),
|
||||
"split_newlines": ("BOOL", {"default": True}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
FUNCTION = "encode"
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
CATEGORY = "conditioning"
|
||||
|
||||
@staticmethod
|
||||
def encode(clip, text, nudge_start, nudge_end, split_newlines):
|
||||
ret = []
|
||||
if split_newlines:
|
||||
prompts = text.split("\n")
|
||||
else:
|
||||
prompts = [text]
|
||||
|
||||
for prompt in prompts:
|
||||
if prompt.strip() == "":
|
||||
continue
|
||||
tokens = clip.tokenizer.tokenize_with_weights(
|
||||
text,
|
||||
return_word_ids=False,
|
||||
nudge_start=nudge_start,
|
||||
nudge_end=nudge_end,
|
||||
)
|
||||
cond, pooled = clip.encode_from_tokens(
|
||||
tokens, return_pooled=True, position_ids=[0] * 77
|
||||
)
|
||||
cond = [[cond, {"pooled_output": pooled}]]
|
||||
ret.append(cond)
|
||||
return (ret,)
|
||||
|
||||
|
||||
def tensor2img(tensor_img):
|
||||
i = 255.0 * tensor_img.cpu().numpy()
|
||||
i_np_arr = np.clip(i, 0, 255, out=i).astype(np.uint8, copy=False)
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
lark
|
||||
Reference in New Issue
Block a user