num_tokens should be used.

Apparently the original code doubled the class tokens for V2. If you were using the RepeatImageBatch node and want the same results as before, try halving the repeat amount.

Also related to #32
This commit is contained in:
shiimizu
2024-09-01 01:19:59 -07:00
parent 13987cf7b8
commit a9ea4cdc5a
2 changed files with 2 additions and 7 deletions
+1 -5
View File
@@ -120,13 +120,9 @@ class FuseModule(nn.Module):
valid_id_embeds = valid_id_embeds.view(-1, valid_id_embeds.shape[-1])
# slice out the image token embeddings
image_token_embeds = prompt_embeds[class_tokens_mask]
# add this due to self.num_tokens earlier causing to double id_embeds, and adjust the assert
# masked_scatter will have the same result if stacked_id_embeds is halved
rep = -(valid_id_embeds.shape[0] // -image_token_embeds.shape[0])
if rep != 1:
image_token_embeds = image_token_embeds.repeat(rep, 1)
stacked_id_embeds = self.fuse_fn(image_token_embeds, valid_id_embeds)
assert class_tokens_mask.sum() == stacked_id_embeds.shape[0] // rep, f"{class_tokens_mask.sum()} != {stacked_id_embeds.shape[0] // rep}"
assert class_tokens_mask.sum() == stacked_id_embeds.shape[0], f"{class_tokens_mask.sum()} != {stacked_id_embeds.shape[0] // rep}"
prompt_embeds.masked_scatter_(class_tokens_mask[:, None], stacked_id_embeds.to(prompt_embeds.dtype))
updated_prompt_embeds = prompt_embeds.view(batch_size, seq_length, -1)
return updated_prompt_embeds
+1 -2
View File
@@ -140,8 +140,7 @@ class PhotoMakerEncodePlus:
tokens = clip.tokenize(text)
class_tokens_mask = {}
out_tokens = {}
num_tokens = getattr(photomaker, 'num_tokens', 2)
num_tokens = 1
num_tokens = getattr(photomaker, 'num_tokens', 1)
for key, val in tokens.items():
clip_tokenizer = getattr(clip.tokenizer, f'clip_{key}', clip.tokenizer)
img_token = clip_tokenizer.tokenizer(trigger_word, truncation=False, add_special_tokens=False)["input_ids"][0] # only get the first token