diff --git a/model_v2.py b/model_v2.py index 49bf217..5d61126 100644 --- a/model_v2.py +++ b/model_v2.py @@ -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 diff --git a/photomaker.py b/photomaker.py index 0de697d..a1c88bd 100644 --- a/photomaker.py +++ b/photomaker.py @@ -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