T5 Output Embedding Support

This commit is contained in:
Koratahiu 2025-08-24 21:30:33 +03:00
parent 71ed4a399e
commit c614b7e045

View File

@ -1,4 +1,5 @@
import os
import logging
from transformers import CLIPTokenizer
import comfy.ops
@ -65,9 +66,29 @@ class ClipTokenWeightEncoder:
if (len(output) == 0):
r = (out[-1:].to(model_management.intermediate_device()), first_pooled)
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 = {}
for k in o[2]:
v = o[2][k]
@ -177,6 +198,8 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
embeds_out = []
attention_masks = []
num_tokens = []
output_embeds_info_batch = []
embeds_info = []
for x in tokens:
attention_mask = []
@ -184,6 +207,7 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
other_embeds = []
eos = False
index = 0
output_embeds_for_prompt = []
for y in x:
if isinstance(y, numbers.Integral):
if eos:
@ -196,6 +220,19 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
if end_token is None:
attention_mask[-1] = 0
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:
other_embeds.append((index, y))
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)
index = 0
pad_extra = 0
embeds_info = []
for o in other_embeds:
emb = o[1]
if torch.is_tensor(emb):
@ -220,6 +256,14 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
else:
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:
index += -1
continue
@ -242,15 +286,16 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
tokens_embed = torch.cat([tokens_embed, padd_embed], dim=1)
attention_mask = attention_mask + [0] * pad_extra
output_embeds_info_batch.append(output_embeds_for_prompt)
embeds_out.append(tokens_embed)
attention_masks.append(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):
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
if self.enable_attention_masks:
@ -283,9 +328,9 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
extra["attention_mask"] = attention_mask
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):
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))
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 'string_to_param' in embed:
values = embed['string_to_param'].values()
embed_out = next(iter(values))
embed_out = {"type": "embedding", "data": next(iter(values))}
elif isinstance(embed, list):
out_list = []
for x in range(len(embed)):
for k in embed[x]:
t = embed[x][k]
if t.shape[-1] != embedding_size:
continue
out_list.append(t.reshape(-1, t.shape[-1]))
embed_out = torch.cat(out_list, dim=0)
elif embed_key is not None and embed_key in embed:
embed_out = embed[embed_key]
else:
embed_out = bundled_embed(embed, 'bundle_emb.', '.string_to_param.*')
if embed_out is None:
embed_out = bundled_embed(embed, 'bundle_emb.', '.{}'.format(embed_key))
continue # Skip mismatched tensors
out_list.append(t.reshape(-1, t.shape[-1])) # Reshape to (num_tokens, embedding_dim)
embed_out = {"type": "embedding", "data": torch.cat(out_list, dim=0)}
elif embed_key is not None:
mapped_key = embed_key_map.get(embed_key, embed_key)
output_key = f"{mapped_key}_out" # output embeddings has _out tensors like "t5_out"
if output_key in embed:
embed_out = {"type": "output_embedding", "data": embed[output_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:
values = embed.values()
embed_out = next(iter(values))
embed_out = {"type": "embedding", "data": next(iter(values))}
return embed_out
class SDTokenizer:
@ -557,10 +619,14 @@ class SDTokenizer:
if embed is None:
logging.warning(f"warning, embedding:{embedding_name} does not exist, ignoring")
else:
if len(embed.shape) == 1:
tokens.append([(embed, weight)])
# The entire dictionary {type, data} is now the "token"
embed_tensor = embed["data"]
if len(embed_tensor.shape) == 1:
tokens.append([(embed, weight)]) # Single token embedding
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 leftover != "":
word = leftover