Initial working commit
This commit is contained in:
@@ -1,5 +1,8 @@
|
||||
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
|
||||
|
||||
@@ -37,6 +40,54 @@ class ArithAction(Action):
|
||||
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(
|
||||
|
||||
+26
-1
@@ -1,7 +1,9 @@
|
||||
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
|
||||
|
||||
@@ -17,6 +19,18 @@ class Action(ABC):
|
||||
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(
|
||||
@@ -37,6 +51,14 @@ class PromptSegment:
|
||||
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}"('
|
||||
|
||||
@@ -59,7 +81,10 @@ def build_prompt_segment(text: str, tokenizer: SD1Tokenizer) -> PromptSegment:
|
||||
if embedding is None:
|
||||
print(f"warning, embedding:{embedding_name} does not exist, ignoring")
|
||||
else:
|
||||
tokens.append(embedding)
|
||||
if len(embedding.shape) == 1:
|
||||
tokens.append(embedding)
|
||||
else:
|
||||
tokens.extend(embedding)
|
||||
|
||||
if leftover != "":
|
||||
word = leftover
|
||||
|
||||
@@ -1,10 +1,53 @@
|
||||
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 = "]"
|
||||
|
||||
|
||||
+55
-39
@@ -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
@@ -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:
|
||||
|
||||
+34
-38
@@ -18,7 +18,7 @@ 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\-_]+>)'
|
||||
|
||||
@@ -115,9 +115,11 @@ 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) -> list[list[Action | int]]:
|
||||
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:
|
||||
@@ -139,44 +141,38 @@ class MyTokenizer(SD1Tokenizer):
|
||||
else:
|
||||
print(segment.depth_repr())
|
||||
|
||||
tokens: list[list[Action | int ]] = []
|
||||
|
||||
# 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
|
||||
|
||||
Reference in New Issue
Block a user