mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-26 19:24:25 +08:00
T5 Output Embedding Support
This commit is contained in:
parent
71ed4a399e
commit
c614b7e045
@ -1,4 +1,5 @@
|
|||||||
import os
|
import os
|
||||||
|
import logging
|
||||||
|
|
||||||
from transformers import CLIPTokenizer
|
from transformers import CLIPTokenizer
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
@ -65,9 +66,29 @@ class ClipTokenWeightEncoder:
|
|||||||
if (len(output) == 0):
|
if (len(output) == 0):
|
||||||
r = (out[-1:].to(model_management.intermediate_device()), first_pooled)
|
r = (out[-1:].to(model_management.intermediate_device()), first_pooled)
|
||||||
else:
|
else:
|
||||||
r = (torch.cat(output, dim=-2).to(model_management.intermediate_device()), first_pooled)
|
final_cond = torch.cat(output, dim=-2)
|
||||||
|
# o[3] is the output_embeds_info from the forward pass
|
||||||
|
if len(o) > 3 and o[3] is not None:
|
||||||
|
output_embeds_info_batch = o[3]
|
||||||
|
# Check if there are any output embeddings to apply across all batches
|
||||||
|
if any(output_embeds_info_batch):
|
||||||
|
final_cond = final_cond.clone()
|
||||||
|
logging.info("Applying output embeddings...")
|
||||||
|
# output_embeds_info_batch is a list, one entry per item in the batch
|
||||||
|
# final_cond is a single tensor with batch items concatenated on the sequence length dimension
|
||||||
|
seq_len = out.shape[1] # The length of a single prompt chunk (e.g., 77)
|
||||||
|
for batch_idx, prompt_embeds_list in enumerate(output_embeds_info_batch):
|
||||||
|
if not prompt_embeds_list:
|
||||||
|
continue
|
||||||
|
|
||||||
|
for seq_idx, embed_tensor in prompt_embeds_list:
|
||||||
|
final_seq_idx = batch_idx * seq_len + seq_idx
|
||||||
|
num_embed_tokens = embed_tensor.shape[0] if len(embed_tensor.shape) > 1 else 1
|
||||||
|
final_cond[0, final_seq_idx : final_seq_idx + num_embed_tokens] = embed_tensor.reshape(num_embed_tokens, -1).to(device=final_cond.device, dtype=final_cond.dtype)
|
||||||
|
|
||||||
if len(o) > 2:
|
r = (final_cond.to(model_management.intermediate_device()), first_pooled)
|
||||||
|
|
||||||
|
if len(o) > 2 and o[2] is not None:
|
||||||
extra = {}
|
extra = {}
|
||||||
for k in o[2]:
|
for k in o[2]:
|
||||||
v = o[2][k]
|
v = o[2][k]
|
||||||
@ -177,6 +198,8 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
|
|||||||
embeds_out = []
|
embeds_out = []
|
||||||
attention_masks = []
|
attention_masks = []
|
||||||
num_tokens = []
|
num_tokens = []
|
||||||
|
output_embeds_info_batch = []
|
||||||
|
embeds_info = []
|
||||||
|
|
||||||
for x in tokens:
|
for x in tokens:
|
||||||
attention_mask = []
|
attention_mask = []
|
||||||
@ -184,6 +207,7 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
|
|||||||
other_embeds = []
|
other_embeds = []
|
||||||
eos = False
|
eos = False
|
||||||
index = 0
|
index = 0
|
||||||
|
output_embeds_for_prompt = []
|
||||||
for y in x:
|
for y in x:
|
||||||
if isinstance(y, numbers.Integral):
|
if isinstance(y, numbers.Integral):
|
||||||
if eos:
|
if eos:
|
||||||
@ -196,6 +220,19 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
|
|||||||
if end_token is None:
|
if end_token is None:
|
||||||
attention_mask[-1] = 0
|
attention_mask[-1] = 0
|
||||||
eos = True
|
eos = True
|
||||||
|
# Check for the dictionary structure we created in load_embed
|
||||||
|
elif isinstance(y, dict) and "type" in y:
|
||||||
|
num_tokens_in_embed = y["data"].shape[0] if len(y["data"].shape) > 1 else 1
|
||||||
|
if y["type"] == "output_embedding":
|
||||||
|
# For output embeddings, store their position and data for later.
|
||||||
|
# Insert placeholder tokens to maintain sequence length.
|
||||||
|
output_embeds_for_prompt.append((index, y["data"]))
|
||||||
|
tokens_temp.extend([self.special_tokens["pad"]] * num_tokens_in_embed)
|
||||||
|
attention_mask.extend([1] * num_tokens_in_embed)
|
||||||
|
else: # Regular input embedding
|
||||||
|
other_embeds.append((index, y))
|
||||||
|
index += num_tokens_in_embed
|
||||||
|
continue # Skip the index+=1 at the end of the loop
|
||||||
else:
|
else:
|
||||||
other_embeds.append((index, y))
|
other_embeds.append((index, y))
|
||||||
index += 1
|
index += 1
|
||||||
@ -204,7 +241,6 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
|
|||||||
tokens_embed = self.transformer.get_input_embeddings()(tokens_embed, out_dtype=torch.float32)
|
tokens_embed = self.transformer.get_input_embeddings()(tokens_embed, out_dtype=torch.float32)
|
||||||
index = 0
|
index = 0
|
||||||
pad_extra = 0
|
pad_extra = 0
|
||||||
embeds_info = []
|
|
||||||
for o in other_embeds:
|
for o in other_embeds:
|
||||||
emb = o[1]
|
emb = o[1]
|
||||||
if torch.is_tensor(emb):
|
if torch.is_tensor(emb):
|
||||||
@ -220,6 +256,14 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
|
|||||||
else:
|
else:
|
||||||
emb = None
|
emb = None
|
||||||
|
|
||||||
|
# Adjust output embedding indices based on where
|
||||||
|
# multi-token INPUT embeddings were inserted.
|
||||||
|
if emb is not None:
|
||||||
|
num_new_tokens = emb.shape[1]
|
||||||
|
for i in range(len(output_embeds_for_prompt)):
|
||||||
|
if output_embeds_for_prompt[i][0] > o[0]:
|
||||||
|
output_embeds_for_prompt[i] = (output_embeds_for_prompt[i][0] + num_new_tokens - 1, output_embeds_for_prompt[i][1])
|
||||||
|
|
||||||
if emb is None:
|
if emb is None:
|
||||||
index += -1
|
index += -1
|
||||||
continue
|
continue
|
||||||
@ -242,15 +286,16 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
|
|||||||
tokens_embed = torch.cat([tokens_embed, padd_embed], dim=1)
|
tokens_embed = torch.cat([tokens_embed, padd_embed], dim=1)
|
||||||
attention_mask = attention_mask + [0] * pad_extra
|
attention_mask = attention_mask + [0] * pad_extra
|
||||||
|
|
||||||
|
output_embeds_info_batch.append(output_embeds_for_prompt)
|
||||||
embeds_out.append(tokens_embed)
|
embeds_out.append(tokens_embed)
|
||||||
attention_masks.append(attention_mask)
|
attention_masks.append(attention_mask)
|
||||||
num_tokens.append(sum(attention_mask))
|
num_tokens.append(sum(attention_mask))
|
||||||
|
|
||||||
return torch.cat(embeds_out), torch.tensor(attention_masks, device=device, dtype=torch.long), num_tokens, embeds_info
|
return torch.cat(embeds_out), torch.tensor(attention_masks, device=device, dtype=torch.long), num_tokens, output_embeds_info_batch, embeds_info
|
||||||
|
|
||||||
def forward(self, tokens):
|
def forward(self, tokens):
|
||||||
device = self.transformer.get_input_embeddings().weight.device
|
device = self.transformer.get_input_embeddings().weight.device
|
||||||
embeds, attention_mask, num_tokens, embeds_info = self.process_tokens(tokens, device)
|
embeds, attention_mask, num_tokens, output_embeds_info, embeds_info = self.process_tokens(tokens, device)
|
||||||
|
|
||||||
attention_mask_model = None
|
attention_mask_model = None
|
||||||
if self.enable_attention_masks:
|
if self.enable_attention_masks:
|
||||||
@ -283,9 +328,9 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
|
|||||||
extra["attention_mask"] = attention_mask
|
extra["attention_mask"] = attention_mask
|
||||||
|
|
||||||
if len(extra) > 0:
|
if len(extra) > 0:
|
||||||
return z, pooled_output, extra
|
return z, pooled_output, extra, output_embeds_info
|
||||||
|
|
||||||
return z, pooled_output
|
return z, pooled_output, None, output_embeds_info
|
||||||
|
|
||||||
def encode(self, tokens):
|
def encode(self, tokens):
|
||||||
return self(tokens)
|
return self(tokens)
|
||||||
@ -435,28 +480,45 @@ def load_embed(embedding_name, embedding_directory, embedding_size, embed_key=No
|
|||||||
logging.warning("{}\n\nerror loading embedding, skipping loading: {}".format(traceback.format_exc(), embedding_name))
|
logging.warning("{}\n\nerror loading embedding, skipping loading: {}".format(traceback.format_exc(), embedding_name))
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
# Map comfy embedding keys to keys used in the training framework's saved files
|
||||||
|
embed_key_map = {
|
||||||
|
't5xxl': 't5',
|
||||||
|
}
|
||||||
|
|
||||||
if embed_out is None:
|
if embed_out is None:
|
||||||
if 'string_to_param' in embed:
|
if 'string_to_param' in embed:
|
||||||
values = embed['string_to_param'].values()
|
values = embed['string_to_param'].values()
|
||||||
embed_out = next(iter(values))
|
embed_out = {"type": "embedding", "data": next(iter(values))}
|
||||||
elif isinstance(embed, list):
|
elif isinstance(embed, list):
|
||||||
out_list = []
|
out_list = []
|
||||||
for x in range(len(embed)):
|
for x in range(len(embed)):
|
||||||
for k in embed[x]:
|
for k in embed[x]:
|
||||||
t = embed[x][k]
|
t = embed[x][k]
|
||||||
if t.shape[-1] != embedding_size:
|
if t.shape[-1] != embedding_size:
|
||||||
continue
|
continue # Skip mismatched tensors
|
||||||
out_list.append(t.reshape(-1, t.shape[-1]))
|
out_list.append(t.reshape(-1, t.shape[-1])) # Reshape to (num_tokens, embedding_dim)
|
||||||
embed_out = torch.cat(out_list, dim=0)
|
embed_out = {"type": "embedding", "data": torch.cat(out_list, dim=0)}
|
||||||
elif embed_key is not None and embed_key in embed:
|
elif embed_key is not None:
|
||||||
embed_out = embed[embed_key]
|
mapped_key = embed_key_map.get(embed_key, embed_key)
|
||||||
else:
|
output_key = f"{mapped_key}_out" # output embeddings has _out tensors like "t5_out"
|
||||||
embed_out = bundled_embed(embed, 'bundle_emb.', '.string_to_param.*')
|
if output_key in embed:
|
||||||
if embed_out is None:
|
embed_out = {"type": "output_embedding", "data": embed[output_key]}
|
||||||
embed_out = bundled_embed(embed, 'bundle_emb.', '.{}'.format(embed_key))
|
elif mapped_key in embed:
|
||||||
|
embed_out = {"type": "embedding", "data": embed[mapped_key]}
|
||||||
|
elif embed_key in embed: # Fallback to original key if mapped keys fail
|
||||||
|
embed_out = {"type": "embedding", "data": embed[embed_key]}
|
||||||
|
|
||||||
|
if embed_out is None: # Fallback for other formats
|
||||||
|
bundled = bundled_embed(embed, 'bundle_emb.', '.string_to_param.*')
|
||||||
|
if bundled is not None:
|
||||||
|
embed_out = {"type": "embedding", "data": bundled}
|
||||||
|
if embed_out is None and embed_key is not None:
|
||||||
|
bundled = bundled_embed(embed, 'bundle_emb.', '.{}'.format(embed_key))
|
||||||
|
if bundled is not None:
|
||||||
|
embed_out = {"type": "embedding", "data": bundled}
|
||||||
if embed_out is None:
|
if embed_out is None:
|
||||||
values = embed.values()
|
values = embed.values()
|
||||||
embed_out = next(iter(values))
|
embed_out = {"type": "embedding", "data": next(iter(values))}
|
||||||
return embed_out
|
return embed_out
|
||||||
|
|
||||||
class SDTokenizer:
|
class SDTokenizer:
|
||||||
@ -557,10 +619,14 @@ class SDTokenizer:
|
|||||||
if embed is None:
|
if embed is None:
|
||||||
logging.warning(f"warning, embedding:{embedding_name} does not exist, ignoring")
|
logging.warning(f"warning, embedding:{embedding_name} does not exist, ignoring")
|
||||||
else:
|
else:
|
||||||
if len(embed.shape) == 1:
|
# The entire dictionary {type, data} is now the "token"
|
||||||
tokens.append([(embed, weight)])
|
embed_tensor = embed["data"]
|
||||||
|
if len(embed_tensor.shape) == 1:
|
||||||
|
tokens.append([(embed, weight)]) # Single token embedding
|
||||||
else:
|
else:
|
||||||
tokens.append([(embed[x], weight) for x in range(embed.shape[0])])
|
# For multi-token embeddings, we need to give each vector its own dictionary
|
||||||
|
# so that process_tokens can handle them individually.
|
||||||
|
tokens.append([({"type": embed["type"], "data": embed_tensor[x]}, weight) for x in range(embed_tensor.shape[0])])
|
||||||
#if we accidentally have leftover text, continue parsing using leftover, else move on to next word
|
#if we accidentally have leftover text, continue parsing using leftover, else move on to next word
|
||||||
if leftover != "":
|
if leftover != "":
|
||||||
word = leftover
|
word = leftover
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user