diff --git a/comfy/sd1_clip.py b/comfy/sd1_clip.py index f8a7c2a1b..1bee4d110 100644 --- a/comfy/sd1_clip.py +++ b/comfy/sd1_clip.py @@ -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