6 Commits
7 changed files with 546 additions and 250 deletions
+121 -29
View File
@@ -1,39 +1,131 @@
from custom_nodes.ClipStuff.lib.actions.base import Action
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: str, ops_str: str):
def __init__(self, base_segment: PromptSegment | Action, ops: dict[str, list[PromptSegment | Action]]):
self.base_segment = base_segment
self.ops = self.process_ops_string(ops_str)
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 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
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>
# Add last op to dict
ops_dict[curr_op_char] += [buff]
return ops_dict
: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)
+82
View File
@@ -1,4 +1,11 @@
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):
@@ -11,3 +18,78 @@ class Action(ABC):
@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)
+118 -3
View File
@@ -1,13 +1,128 @@
from typing import Optional
from typing import Optional, Union, Callable
from custom_nodes.ClipStuff.lib.actions.base import Action
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=None, weight: Optional[float] = None, target=None):
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)
+55 -39
View File
@@ -1,5 +1,6 @@
import contextlib
import os
from typing import Union
import torch
from transformers import CLIPTextConfig, modeling_utils
@@ -7,6 +8,7 @@ from transformers import CLIPTextConfig, modeling_utils
from comfy import model_management
import comfy.ops
from comfy.sd import CLIP
from custom_nodes.ClipStuff.lib.actions.base import PromptSegment, Action
from custom_nodes.ClipStuff.lib.fun_clip_stuff import MyCLIPTextModel
from custom_nodes.ClipStuff.lib.tokenizer import TokenDict
@@ -66,50 +68,69 @@ 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 = []
def set_up_textual_embeddings(self, tokens: list[list[PromptSegment | Action]], current_embeds):
next_new_token = token_dict_size = current_embeds.weight.shape[0] - 1
embedding_weights = []
# 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
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
def forward(self, tokens, **kwargs):
backup_embeds = self.transformer.get_input_embeddings()
device = backup_embeds.weight.device
@@ -153,15 +174,10 @@ class SD1FunClipModel(torch.nn.Module):
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)
def encode_token_weights(self, prompt_segments: list[list[Union[PromptSegment | Action]]], **kwargs):
to_encode = [[PromptSegment(text="_Empty Batch_", tokens=self.empty_tokens[0])]]
for batch in prompt_segments:
to_encode.append(batch)
out, pooled = self.encode(to_encode, **kwargs)
z_empty = out[0:1]
@@ -173,10 +189,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):
+63 -36
View File
@@ -7,6 +7,7 @@ from transformers.modeling_outputs import BaseModelOutputWithPooling
from transformers.models.clip.modeling_clip import _expand_mask, CLIPTextEmbeddings, CLIPTextTransformer, \
CLIPTextModel
from custom_nodes.ClipStuff.lib.actions.base import PromptSegment, Action
from custom_nodes.ClipStuff.lib.tokenizer import TokenDict
def slerp(val, low, high):
@@ -30,47 +31,57 @@ class MyCLIPTextEmbeddings(CLIPTextEmbeddings):
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]
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)
# 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]
# 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):
@@ -80,7 +91,7 @@ class MyCLIPTextTransformer(CLIPTextTransformer):
def forward(
self,
input_ids: Optional[list[list[tuple[TokenDict]]]] = None,
input_ids: Optional[list[list[PromptSegment | Action]]] = None,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.Tensor] = None,
output_attentions: Optional[bool] = None,
@@ -103,9 +114,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 +142,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:
+106 -142
View File
@@ -1,3 +1,4 @@
import re
from typing import Union
from comfy.sd1_clip import SD1Tokenizer
@@ -6,14 +7,68 @@ from custom_nodes.ClipStuff.lib.actions import (
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
arith_action = r'(<[a-zA-Z0-9\-_]+:[a-zA-Z0-9\-_]+>)'
# TODO: Get embedding identifier from tokenizer
tokenizer_regex = re.compile(
fr"""
\d+\.\d+ # Capture decimals
|
(?:(?!embedding:)[\w\s]|embedding:[a-zA-Z0-9_]+)+ # Capture sequences of characters, including "embedding:"
|
\d+ # Capture whole numbers
|
[:+-{re.escape("".join(ALL_START_CHARS))}{re.escape("".join(ALL_END_CHARS))}] # Capture special characters including start and end characters
""",
re.VERBOSE
)
def tokenize(text: str) -> list[str]:
# Captures:
# 1. Words
# 2. Numbers(1.0, 1)
# 3. Special characters(ALL_START_CHARS, ALL_END_CHARS, :, +, -)
tokens = re.findall(tokenizer_regex, text)
print(tokens)
return [token.strip() for token in tokens]
def parse_special_tokens(string):
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 = ""
@@ -30,49 +85,10 @@ def parse_special_tokens(string):
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
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,
@@ -99,116 +115,64 @@ class MyTokenizer(SD1Tokenizer):
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[PromptSegment | Action]]:
if self.pad_with_end:
pad_token = self.end_token
else:
pad_token = 0
parsed_actions = parse_token_actions(text)
parsed_actions = parse_segment_actions(text, self)
nudge_start = kwargs.get("nudge_start")
nudge_end = kwargs.get("nudge_end")
if nudge_start is not None and nudge_end is not None:
nudge_start = int(nudge_start)
nudge_end = int(nudge_end)
# tokenize words
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]]
# 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:
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]])
print(segment.depth_repr())
# 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
for segment in parsed_actions:
num_tokens = segment.token_length()
# determine if we're going to try and keep the tokens in a single batch
is_large = 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
]
# If the segment is small enough to fit 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
+1 -1
View File
@@ -37,7 +37,7 @@ class SpecialClipLoader:
clip_target.tokenizer = MyTokenizer
# 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()
)