Initial commit.
This commit is contained in:
@@ -0,0 +1,25 @@
|
||||
# PhotoMaker for ComfyUI
|
||||
|
||||
Unofficial implementation of [PhotoMaker](https://github.com/TencentARC/PhotoMaker) for ComfyUI.
|
||||
|
||||
---
|
||||
|
||||
Download the [model](https://huggingface.co/TencentARC/PhotoMaker) and place it in your `models` folder such as `ComfyUI/models/photomaker`.
|
||||
|
||||
Extract the LoRa and place it in your `loras` folder:
|
||||
```python
|
||||
import torch
|
||||
from safetensors.torch import save_file
|
||||
sd = torch.load("/path/to/photomaker-v1.bin", map_location="cpu")
|
||||
save_file(sd["lora_weights"], "ComfyUI/models/loras/photomaker-v1-lora.safetensors")
|
||||
```
|
||||
|
||||
# Citation
|
||||
```bibtex
|
||||
@article{li2023photomaker,
|
||||
title={PhotoMaker: Customizing Realistic Human Photos via Stacked ID Embedding},
|
||||
author={Li, Zhen and Cao, Mingdeng and Wang, Xintao and Qi, Zhongang and Cheng, Ming-Ming and Shan, Ying},
|
||||
booktitle={arXiv preprint arxiv:2312.04461},
|
||||
year={2023}
|
||||
}
|
||||
```
|
||||
@@ -0,0 +1,2 @@
|
||||
from .photomaker import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
@@ -0,0 +1,151 @@
|
||||
# Merge image encoder and fuse module to create an ID Encoder
|
||||
# send multiple ID images, we can directly obtain the updated text encoder containing a stacked ID embedding
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
# from transformers.models.clip.modeling_clip import CLIPVisionModelWithProjection
|
||||
# from transformers.models.clip.configuration_clip import CLIPVisionConfig
|
||||
# from transformers import PretrainedConfig
|
||||
from comfy.clip_model import ACTIVATIONS
|
||||
from comfy.clip_model import CLIPVisionModelProjection
|
||||
from comfy.ops import manual_cast
|
||||
|
||||
VISION_CONFIG_DICT = {
|
||||
"hidden_size": 1024,
|
||||
"intermediate_size": 4096,
|
||||
"num_attention_heads": 16,
|
||||
"num_hidden_layers": 24,
|
||||
"patch_size": 14,
|
||||
"projection_dim": 768
|
||||
}
|
||||
act_fn = None
|
||||
class MLP(nn.Module):
|
||||
def __init__(self, in_dim, out_dim, hidden_dim, op: manual_cast, dtype, device, use_residual=True):
|
||||
super().__init__()
|
||||
if use_residual:
|
||||
assert in_dim == out_dim
|
||||
self.layernorm = op.LayerNorm(in_dim, device=device, dtype=dtype)
|
||||
self.fc1 = op.Linear(in_dim, hidden_dim, device=device, dtype=dtype)
|
||||
self.fc2 = op.Linear(hidden_dim, out_dim, device=device, dtype=dtype)
|
||||
self.use_residual = use_residual
|
||||
global act_fn
|
||||
self.act_fn = act_fn
|
||||
# self.layernorm = nn.LayerNorm(in_dim)
|
||||
# self.fc1 = nn.Linear(in_dim, hidden_dim)
|
||||
# self.fc2 = nn.Linear(hidden_dim, out_dim)
|
||||
# self.use_residual = use_residual
|
||||
# self.act_fn = nn.GELU()
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
x = self.layernorm(x)
|
||||
x = self.fc1(x)
|
||||
x = self.act_fn(x)
|
||||
x = self.fc2(x)
|
||||
if self.use_residual:
|
||||
x = x + residual
|
||||
return x
|
||||
|
||||
|
||||
class FuseModule(nn.Module):
|
||||
def __init__(self, embed_dim, op: manual_cast, dtype, device):
|
||||
super().__init__()
|
||||
self.mlp1 = MLP(embed_dim * 2, embed_dim, embed_dim, use_residual=False, op=op, device=device, dtype=dtype)
|
||||
self.mlp2 = MLP(embed_dim, embed_dim, embed_dim, use_residual=True, op=op, device=device, dtype=dtype)
|
||||
self.layer_norm = op.LayerNorm(embed_dim, device=device, dtype=dtype)
|
||||
# self.layer_norm = nn.LayerNorm(embed_dim)
|
||||
|
||||
def fuse_fn(self, prompt_embeds, id_embeds):
|
||||
stacked_id_embeds = torch.cat([prompt_embeds, id_embeds], dim=-1)
|
||||
stacked_id_embeds = self.mlp1(stacked_id_embeds) + prompt_embeds
|
||||
stacked_id_embeds = self.mlp2(stacked_id_embeds)
|
||||
stacked_id_embeds = self.layer_norm(stacked_id_embeds)
|
||||
return stacked_id_embeds
|
||||
|
||||
def forward(
|
||||
self,
|
||||
prompt_embeds,
|
||||
id_embeds,
|
||||
class_tokens_mask,
|
||||
) -> torch.Tensor:
|
||||
# id_embeds shape: [b, max_num_inputs, 1, 2048]
|
||||
id_embeds = id_embeds.to(prompt_embeds.dtype)
|
||||
num_inputs = class_tokens_mask.sum().unsqueeze(0) # TODO: check for training case
|
||||
batch_size, max_num_inputs = id_embeds.shape[:2]
|
||||
# seq_length: 77
|
||||
seq_length = prompt_embeds.shape[1]
|
||||
# flat_id_embeds shape: [b*max_num_inputs, 1, 2048]
|
||||
flat_id_embeds = id_embeds.view(
|
||||
-1, id_embeds.shape[-2], id_embeds.shape[-1]
|
||||
)
|
||||
# valid_id_mask [b*max_num_inputs]
|
||||
valid_id_mask = (
|
||||
torch.arange(max_num_inputs, device=flat_id_embeds.device)[None, :]
|
||||
< num_inputs[:, None]
|
||||
)
|
||||
valid_id_embeds = flat_id_embeds[valid_id_mask.flatten()]
|
||||
|
||||
prompt_embeds = prompt_embeds.view(-1, prompt_embeds.shape[-1])
|
||||
class_tokens_mask = class_tokens_mask.view(-1)
|
||||
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]
|
||||
stacked_id_embeds = self.fuse_fn(image_token_embeds, valid_id_embeds)
|
||||
assert class_tokens_mask.sum() == stacked_id_embeds.shape[0], f"{class_tokens_mask.sum()} != {stacked_id_embeds.shape[0]}"
|
||||
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
|
||||
|
||||
# class PhotoMakerIDEncoder(CLIPVisionModelWithProjection):
|
||||
# def __init__(self, config_dict, dtype, device, operations):
|
||||
# super().__init__(CLIPVisionConfig(**VISION_CONFIG_DICT))
|
||||
class PhotoMakerIDEncoder(CLIPVisionModelProjection):
|
||||
def __init__(self, config_dict, dtype, device, op: manual_cast):
|
||||
super().__init__(config_dict, dtype, device, op)
|
||||
intermediate_activation = config_dict["hidden_act"]
|
||||
global act_fn
|
||||
act_fn = ACTIVATIONS[intermediate_activation]
|
||||
self.visual_projection_2 = op.Linear(1024, 1280, bias=False, device=device, dtype=dtype)
|
||||
# self.visual_projection_2 = nn.Linear(1024, 1280, bias=False)
|
||||
# self.vision_model = self.vision_model.to(device=device,dtype=dtype)
|
||||
# self.visual_projection = self.visual_projection.to(device=device,dtype=dtype)
|
||||
# self.visual_projection_2 = self.visual_projection_2.to(device=device,dtype=dtype)
|
||||
self.fuse_module = FuseModule(2048, op, device=device, dtype=dtype)
|
||||
|
||||
# def forward(self, id_pixel_values, prompt_embeds, class_tokens_mask):
|
||||
def forward(self, id_pixel_values, prompt_embeds, class_tokens_mask, *args, **kwargs):
|
||||
b, num_inputs, c, h, w = id_pixel_values.shape
|
||||
id_pixel_values = id_pixel_values.view(b * num_inputs, c, h, w)
|
||||
|
||||
# CLIPVisionModelWithProjection
|
||||
# x = self.vision_model(id_pixel_values, output_hidden_states=True)
|
||||
# shared_id_embeds = x[1].to(device=self.device,dtype=self.dtype)
|
||||
|
||||
x = self.vision_model(id_pixel_values, **kwargs)
|
||||
shared_id_embeds = x[2]
|
||||
# shared_id_embeds = self.vision_model(id_pixel_values)[2]
|
||||
# shared_id_embeds = self.vision_model(id_pixel_values, intermediate_output=kwargs['intermediate_output'])[1]
|
||||
# shared_id_embeds = self.vision_model(id_pixel_values, kwargs['intermediate_output'])[1]
|
||||
# shared_id_embeds = self.vision_model(*args, **kwargs))[1]
|
||||
# shared_id_embeds = self.vision_model(id_pixel_values)[1]
|
||||
# shared_id_embeds = shared_id_embeds[2] if shared_id_embeds[2] is not None else shared_id_embeds[0]
|
||||
|
||||
id_embeds = self.visual_projection(shared_id_embeds)
|
||||
id_embeds_2 = self.visual_projection_2(shared_id_embeds)
|
||||
|
||||
id_embeds = id_embeds.view(b, num_inputs, 1, -1)
|
||||
id_embeds_2 = id_embeds_2.view(b, num_inputs, 1, -1)
|
||||
|
||||
id_embeds = torch.cat((id_embeds, id_embeds_2), dim=-1)
|
||||
prompt_embeds = prompt_embeds[0].to(device=id_embeds.device)
|
||||
class_tokens_mask = class_tokens_mask.to(device=id_embeds.device)
|
||||
updated_prompt_embeds = self.fuse_module(prompt_embeds, id_embeds, class_tokens_mask)
|
||||
|
||||
return (x[0], x[1], updated_prompt_embeds)
|
||||
|
||||
# CLIPVisionModelWithProjection
|
||||
# return (x.last_hidden_state, x.hidden_states[0], updated_prompt_embeds)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
PhotoMakerIDEncoder()
|
||||
+207
@@ -0,0 +1,207 @@
|
||||
import comfy.clip_vision
|
||||
import comfy.clip_model
|
||||
import comfy.model_management
|
||||
from comfy.sd import CLIP
|
||||
from comfy.clip_vision import ClipVisionModel
|
||||
import folder_paths
|
||||
from .model import PhotoMakerIDEncoder
|
||||
from copy import deepcopy
|
||||
from .utils import load_image, hook_all
|
||||
from transformers import CLIPImageProcessor
|
||||
from transformers.image_utils import PILImageResampling
|
||||
import torch
|
||||
import os
|
||||
from folder_paths import folder_names_and_paths, models_dir, supported_pt_extensions
|
||||
|
||||
folder_names_and_paths["photomaker"] = ([os.path.join(models_dir, "photomaker")], supported_pt_extensions)
|
||||
|
||||
class PhotoMakerLoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": { "clip_name": (folder_paths.get_filename_list("photomaker"), ),
|
||||
}}
|
||||
RETURN_TYPES = ("CLIP_VISION",)
|
||||
FUNCTION = "load_clip"
|
||||
|
||||
CATEGORY = "photomaker"
|
||||
|
||||
def load_clip(self, clip_name):
|
||||
hook_all()
|
||||
comfy.clip_model.CLIPVisionModelProjection_original = comfy.clip_model.CLIPVisionModelProjection
|
||||
comfy.clip_model.CLIPVisionModelProjection = PhotoMakerIDEncoder
|
||||
clip_path = folder_paths.get_full_path("photomaker", clip_name)
|
||||
# sd = comfy.clip_vision.load_torch_file(clip_path)
|
||||
sd = torch.load(clip_path, map_location="cpu")
|
||||
if 'id_encoder' in sd:
|
||||
sd = sd['id_encoder']
|
||||
# clip_vision = comfy.clip_vision.load_clipvision_from_sd(sd, "id_encoder.", True)
|
||||
clip_vision = comfy.clip_vision.load_clipvision_from_sd(sd)
|
||||
comfy.clip_model.CLIPVisionModelProjection = comfy.clip_model.CLIPVisionModelProjection_original
|
||||
hook_all(restore=True)
|
||||
return (clip_vision,)
|
||||
|
||||
|
||||
class PhotoMakerEncode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"clip": ("CLIP",),
|
||||
"clip_vision": ("CLIP_VISION",),
|
||||
"text": ("STRING", {"multiline": True, "forceInput": True}),
|
||||
"trigger_word": ("STRING", {"default": "img"}),
|
||||
"ref_images_path": ("STRING", {"multiline": False, "placeholder": "optional"}),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("CONDITIONING","CLIP_VISION_OUTPUT")
|
||||
FUNCTION = "encode"
|
||||
|
||||
CATEGORY = "photomaker"
|
||||
|
||||
def encode(self, clip: CLIP, clip_vision: ClipVisionModel, text: str, trigger_word: str, ref_images_path:str, image=None):
|
||||
|
||||
input_id_images = image
|
||||
if ref_images_path != '':
|
||||
input_id_images=ref_images_path
|
||||
if not isinstance(input_id_images, list):
|
||||
input_id_images = [input_id_images]
|
||||
|
||||
image_basename_list = os.listdir(ref_images_path)
|
||||
#image_path_list = sorted([os.path.join(ref_images_path, basename) for basename in image_basename_list])
|
||||
image_path_list = [
|
||||
os.path.join(ref_images_path, basename)
|
||||
for basename in image_basename_list
|
||||
if not basename.startswith('.') and basename.lower().endswith(('.png', '.jpg', '.jpeg', '.bmp', '.webp')) # 只包括有效的图像文件
|
||||
]
|
||||
|
||||
input_id_images = [load_image(image_path) for image_path in image_path_list]
|
||||
|
||||
if input_id_images is None:
|
||||
raise ValueError("Provide `image`. Cannot leave `image` undefined for PhotoMaker pipeline.")
|
||||
num_id_images = len(input_id_images)
|
||||
# Resize to 224x224
|
||||
comfy.model_management.load_model_gpu(clip_vision.patcher)
|
||||
clip_vision.id_image_processor = CLIPImageProcessor(resample=PILImageResampling.LANCZOS)
|
||||
if not isinstance(input_id_images[0], torch.Tensor):
|
||||
# clip_vision.id_image_processor = CLIPImageProcessor()
|
||||
id_pixel_values = clip_vision.id_image_processor(input_id_images, return_tensors="pt").pixel_values.float()
|
||||
# id_pixel_values = comfy.clip_vision.clip_preprocess(input_id_images[0].to(clip_vision.load_device)).float()
|
||||
# id_pixel_values = self.id_image_processor(input_id_images, return_tensors="pt").pixel_values
|
||||
else:
|
||||
# id_pixel_values = comfy.clip_vision.clip_preprocess(image).float()
|
||||
id_pixel_values = clip_vision.id_image_processor(input_id_images, return_tensors="pt").pixel_values.float()
|
||||
id_pixel_values = id_pixel_values.to(device=clip_vision.load_device)
|
||||
|
||||
clip=clip.clone()
|
||||
trigger_word_tokens = clip.tokenize(trigger_word)
|
||||
class_token = trigger_word_tokens['l'][0][1][0]
|
||||
tokens = clip.tokenize(text)
|
||||
class_tokens_mask = {}
|
||||
for key in tokens:
|
||||
clip_tokenizer = getattr(clip.tokenizer, f'clip_{key}', clip.tokenizer)
|
||||
clip_tokenizer.tokenizer.add_tokens([trigger_word], special_tokens=True)
|
||||
ls = clip_tokenizer.tokenize_with_weights(text, return_tokens=True)
|
||||
# e.g.: [49408]
|
||||
class_token = clip_tokenizer.tokenizer(trigger_word)["input_ids"][clip_tokenizer.tokens_start:-1]
|
||||
|
||||
# get trigger token indices
|
||||
trigger_indices = []
|
||||
for ix, v in enumerate(ls):
|
||||
ids = [tup[0] for tup in v]
|
||||
if ids == class_token:
|
||||
trigger_indices.append(ix)
|
||||
|
||||
# expand trigger tokens and mask
|
||||
ls2 = deepcopy(ls)
|
||||
mask_indices = list(map(lambda i: i-1, trigger_indices))
|
||||
for i, v in enumerate(ls2):
|
||||
if i in trigger_indices:
|
||||
for ii in trigger_indices:
|
||||
if ii-1 < 0: continue
|
||||
ls2[ii-1] = [ls2[ii-1][-1]] * num_id_images
|
||||
elif i not in mask_indices:
|
||||
ids = [(-1, tup[1]) for tup in v]
|
||||
# print(ls[i])
|
||||
ls2[i] = ids
|
||||
|
||||
# expand trigger tokens
|
||||
for ii in trigger_indices:
|
||||
if ii-1 < 0: continue
|
||||
ls[ii-1] = [ls[ii-1][-1]] * num_id_images
|
||||
|
||||
# remove trigger tokens
|
||||
token_weight_pairs = [i for j, i in enumerate(ls) if j not in trigger_indices]
|
||||
token_weight_pairs_mask = [i for j, i in enumerate(ls2) if j not in trigger_indices]
|
||||
# send it back to be batched evenly
|
||||
token_weight_pairs = clip_tokenizer.tokenize_with_weights(text, _tokens=token_weight_pairs)
|
||||
token_weight_pairs_mask = clip_tokenizer.tokenize_with_weights(text, _tokens=token_weight_pairs_mask)
|
||||
tokens[key] = token_weight_pairs
|
||||
|
||||
if clip_tokenizer.pad_with_end:
|
||||
pad_token = clip_tokenizer.end_token
|
||||
else:
|
||||
pad_token = 0
|
||||
|
||||
# Finalize the mask
|
||||
condition = lambda b: isinstance(b, tuple) and isinstance(b[0], int) and b[0] != pad_token and b[0] != clip_tokenizer.start_token and b[0] != -1
|
||||
class_tokens_mask[key] = list(map(lambda a: list(map(lambda b: condition(b), a)), token_weight_pairs_mask))
|
||||
|
||||
cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True)
|
||||
prompt_embeds = cond.to(device=clip_vision.load_device).unsqueeze(0)
|
||||
class_tokens_mask = torch.tensor(class_tokens_mask['l']).to(dtype=torch.bool, device=clip_vision.load_device)
|
||||
if (trigger_indices_len:=len(trigger_indices)) > 1:
|
||||
id_pixel_values = id_pixel_values.repeat([trigger_indices_len] + [1] * (len(id_pixel_values.shape) - 1))
|
||||
|
||||
# From clip_vision.encode_image
|
||||
out = clip_vision.model(id_pixel_values.unsqueeze(0), prompt_embeds, class_tokens_mask, intermediate_output=-2)
|
||||
|
||||
outputs = comfy.clip_vision.Output()
|
||||
outputs["last_hidden_state"] = out[0].to(comfy.model_management.intermediate_device())
|
||||
outputs["penultimate_hidden_states"] = out[1].to(comfy.model_management.intermediate_device())
|
||||
outputs["image_embeds"] = out[2].to(comfy.model_management.intermediate_device())
|
||||
|
||||
return ([[outputs.image_embeds, {"pooled_output": pooled}]], outputs)
|
||||
|
||||
|
||||
from .style_template import styles
|
||||
STYLE_NAMES = list(styles.keys())
|
||||
DEFAULT_STYLE_NAME = "Photographic (Default)"
|
||||
|
||||
def apply_style(style_name: str, positive: str, negative: str = "") -> tuple[str, str]:
|
||||
p, n = styles.get(style_name, styles[DEFAULT_STYLE_NAME])
|
||||
return p.replace("{prompt}", positive), n + ' ' + negative
|
||||
|
||||
class PhotoMakerStyles:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"style_name": (STYLE_NAMES, {"default": DEFAULT_STYLE_NAME}),
|
||||
},
|
||||
"optional": {
|
||||
"positive": ("STRING", {"multiline": False, "forceInput": True}),
|
||||
"negative": ("STRING", {"multiline": False, "forceInput": True}),
|
||||
},
|
||||
}
|
||||
RETURN_TYPES = ("STRING","STRING",)
|
||||
RETURN_NAMES = ("positive","negative",)
|
||||
FUNCTION = "apply"
|
||||
|
||||
CATEGORY = "photomaker"
|
||||
|
||||
def apply(self, style_name, positive: str = '', negative: str = ''):
|
||||
positive, negative = apply_style(style_name, positive, negative)
|
||||
return (positive, negative)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"PhotoMakerLoader": PhotoMakerLoader,
|
||||
"PhotoMakerEncode": PhotoMakerEncode,
|
||||
"PhotoMakerStyles": PhotoMakerStyles,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"PhotoMakerLoader": "PhotoMakerLoader",
|
||||
"PhotoMakerEncode": "PhotoMakerEncode",
|
||||
"PhotoMakerStyles": "PhotoMakerStyles",
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
style_list = [
|
||||
{
|
||||
"name": "(No style)",
|
||||
"prompt": "{prompt}",
|
||||
"negative_prompt": "",
|
||||
},
|
||||
{
|
||||
"name": "Cinematic",
|
||||
"prompt": "cinematic still {prompt} . emotional, harmonious, vignette, highly detailed, high budget, bokeh, cinemascope, moody, epic, gorgeous, film grain, grainy",
|
||||
"negative_prompt": "anime, cartoon, graphic, text, painting, crayon, graphite, abstract, glitch, deformed, mutated, ugly, disfigured",
|
||||
},
|
||||
{
|
||||
"name": "Disney Charactor",
|
||||
"prompt": "A Pixar animation character of {prompt} . pixar-style, studio anime, Disney, high-quality",
|
||||
"negative_prompt": "lowres, bad anatomy, bad hands, text, bad eyes, bad arms, bad legs, error, missing fingers, extra digit, fewer digits, cropped, worst quality, low quality, normal quality, jpeg artifacts, signature, watermark, blurry, grayscale, noisy, sloppy, messy, grainy, highly detailed, ultra textured, photo",
|
||||
},
|
||||
{
|
||||
"name": "Digital Art",
|
||||
"prompt": "concept art {prompt} . digital artwork, illustrative, painterly, matte painting, highly detailed",
|
||||
"negative_prompt": "photo, photorealistic, realism, ugly",
|
||||
},
|
||||
{
|
||||
"name": "Photographic (Default)",
|
||||
"prompt": "cinematic photo {prompt} . 35mm photograph, film, bokeh, professional, 4k, highly detailed",
|
||||
"negative_prompt": "drawing, painting, crayon, sketch, graphite, impressionist, noisy, blurry, soft, deformed, ugly",
|
||||
},
|
||||
{
|
||||
"name": "Fantasy art",
|
||||
"prompt": "ethereal fantasy concept art of {prompt} . magnificent, celestial, ethereal, painterly, epic, majestic, magical, fantasy art, cover art, dreamy",
|
||||
"negative_prompt": "photographic, realistic, realism, 35mm film, dslr, cropped, frame, text, deformed, glitch, noise, noisy, off-center, deformed, cross-eyed, closed eyes, bad anatomy, ugly, disfigured, sloppy, duplicate, mutated, black and white",
|
||||
},
|
||||
{
|
||||
"name": "Neonpunk",
|
||||
"prompt": "neonpunk style {prompt} . cyberpunk, vaporwave, neon, vibes, vibrant, stunningly beautiful, crisp, detailed, sleek, ultramodern, magenta highlights, dark purple shadows, high contrast, cinematic, ultra detailed, intricate, professional",
|
||||
"negative_prompt": "painting, drawing, illustration, glitch, deformed, mutated, cross-eyed, ugly, disfigured",
|
||||
},
|
||||
{
|
||||
"name": "Enhance",
|
||||
"prompt": "breathtaking {prompt} . award-winning, professional, highly detailed",
|
||||
"negative_prompt": "ugly, deformed, noisy, blurry, distorted, grainy",
|
||||
},
|
||||
{
|
||||
"name": "Comic book",
|
||||
"prompt": "comic {prompt} . graphic illustration, comic art, graphic novel art, vibrant, highly detailed",
|
||||
"negative_prompt": "photograph, deformed, glitch, noisy, realistic, stock photo",
|
||||
},
|
||||
{
|
||||
"name": "Lowpoly",
|
||||
"prompt": "low-poly style {prompt} . low-poly game art, polygon mesh, jagged, blocky, wireframe edges, centered composition",
|
||||
"negative_prompt": "noisy, sloppy, messy, grainy, highly detailed, ultra textured, photo",
|
||||
},
|
||||
{
|
||||
"name": "Line art",
|
||||
"prompt": "line art drawing {prompt} . professional, sleek, modern, minimalist, graphic, line art, vector graphics",
|
||||
"negative_prompt": "anime, photorealistic, 35mm film, deformed, glitch, blurry, noisy, off-center, deformed, cross-eyed, closed eyes, bad anatomy, ugly, disfigured, mutated, realism, realistic, impressionism, expressionism, oil, acrylic",
|
||||
}
|
||||
]
|
||||
|
||||
styles = {k["name"]: (k["prompt"], k["negative_prompt"]) for k in style_list}
|
||||
@@ -0,0 +1,210 @@
|
||||
import os
|
||||
import sys
|
||||
import PIL.Image
|
||||
import PIL.ImageOps
|
||||
import requests
|
||||
import torch
|
||||
from typing import List, Union
|
||||
from collections import namedtuple
|
||||
from .model import PhotoMakerIDEncoder
|
||||
from comfy.sd1_clip import escape_important, token_weights, unescape_important
|
||||
|
||||
Hook = namedtuple('Hook', ['fn', 'module_name', 'target', 'orig_key', 'module_name_nt', 'module_name_unix'])
|
||||
|
||||
def hook_clip_model_CLIPVisionModelProjection():
|
||||
return create_hook(PhotoMakerIDEncoder, 'comfy.clip_model', 'CLIPVisionModelProjection')
|
||||
|
||||
def hhok_unclip_adm():
|
||||
import comfy.model_base
|
||||
comfy.model_base.unclip_adm_original = comfy.model_base.unclip_adm
|
||||
return create_hook(unclip_adm, 'comfy.model_base')
|
||||
|
||||
def hook_tokenize_with_weights():
|
||||
import comfy.sd1_clip
|
||||
comfy.sd1_clip.SDTokenizer.tokenize_with_weights_original = comfy.sd1_clip.SDTokenizer.tokenize_with_weights
|
||||
comfy.sd1_clip.SDTokenizer.tokenize_with_weights = tokenize_with_weights
|
||||
return create_hook(tokenize_with_weights, 'comfy.sd1_clip', 'SDTokenizer.tokenize_with_weights')
|
||||
|
||||
def create_hook(fn, module_name, target = None, orig_key = None):
|
||||
if target is None: target = fn.__name__
|
||||
if orig_key is None: orig_key = f'{target}_original'
|
||||
module_name_nt = '\\'.join(module_name.split('.'))
|
||||
module_name_unix = '/'.join(module_name.split('.'))
|
||||
return Hook(fn, module_name, target, orig_key, module_name_nt, module_name_unix)
|
||||
|
||||
def hook_all(restore=False):
|
||||
hooks: List[Hook] = [
|
||||
hook_clip_model_CLIPVisionModelProjection(),
|
||||
hook_tokenize_with_weights(),
|
||||
hhok_unclip_adm(),
|
||||
]
|
||||
for m in sys.modules.keys():
|
||||
for hook in hooks:
|
||||
if hook.module_name == m or (os.name != 'nt' and m.endswith(hook.module_name_unix)) or (os.name == 'nt' and m.endswith(hook.module_name_nt)):
|
||||
if hasattr(sys.modules[m], hook.target):
|
||||
if not hasattr(sys.modules[m], hook.orig_key):
|
||||
if (orig_fn:=getattr(sys.modules[m], hook.target, None)) is not None:
|
||||
setattr(sys.modules[m], hook.orig_key, orig_fn)
|
||||
if restore:
|
||||
setattr(sys.modules[m], hook.target, getattr(sys.modules[m], hook.orig_key, None))
|
||||
else:
|
||||
setattr(sys.modules[m], hook.target, hook.fn)
|
||||
|
||||
|
||||
def unclip_adm(unclip_conditioning, device, noise_augmentor, noise_augment_merge=0.0):
|
||||
adm_inputs = []
|
||||
weights = []
|
||||
noise_aug = []
|
||||
for unclip_cond in unclip_conditioning:
|
||||
for adm_cond in unclip_cond["clip_vision_output"].image_embeds:
|
||||
weight = unclip_cond["strength"]
|
||||
noise_augment = unclip_cond["noise_augmentation"]
|
||||
noise_level = round((noise_augmentor.max_noise_level - 1) * noise_augment)
|
||||
# c_adm, noise_level_emb = noise_augmentor(adm_cond.to(device), noise_level=torch.tensor([noise_level], device=device))
|
||||
# adm_out = torch.cat((c_adm, noise_level_emb), 1) * weight
|
||||
|
||||
# if c_adm.shape[0] > 1:
|
||||
# c_adm = c_adm[:1, :]
|
||||
# c_adm, noise_level_emb = noise_augmentor(adm_cond.to(device), noise_level=torch.tensor([noise_level], device=device))
|
||||
# adm_out = torch.cat((c_adm, noise_level_emb), 1) * weight
|
||||
|
||||
i=1
|
||||
data_mean = noise_augmentor.data_mean.repeat(1, -(adm_cond.shape[i] // -noise_augmentor.data_mean.shape[1]))[:, :adm_cond.shape[i]]
|
||||
data_std = noise_augmentor.data_std.repeat(1, -(adm_cond.shape[i] // -noise_augmentor.data_std.shape[1]))[:, :adm_cond.shape[i]]
|
||||
noise_augmentor.register_buffer("data_mean", data_mean, persistent=False)
|
||||
noise_augmentor.register_buffer("data_std", data_std, persistent=False)
|
||||
# noise_augmentor.data_mean = data_mean
|
||||
# noise_augmentor.data_std = data_std
|
||||
c_adm, noise_level_emb = noise_augmentor(adm_cond.to(device), noise_level=torch.tensor([noise_level], device=device))
|
||||
tmp = noise_level_emb.repeat(-(c_adm.shape[0] // -noise_level_emb.shape[0]),1)[:c_adm.shape[0]]
|
||||
adm_out = torch.cat((c_adm, tmp), 1) * weight
|
||||
|
||||
weights.append(weight)
|
||||
noise_aug.append(noise_augment)
|
||||
adm_inputs.append(adm_out)
|
||||
|
||||
if len(noise_aug) > 1:
|
||||
adm_out = torch.stack(adm_inputs).sum(0)
|
||||
noise_augment = noise_augment_merge
|
||||
noise_level = round((noise_augmentor.max_noise_level - 1) * noise_augment)
|
||||
c_adm, noise_level_emb = noise_augmentor(adm_out[:, :noise_augmentor.time_embed.dim], noise_level=torch.tensor([noise_level], device=device))
|
||||
adm_out = torch.cat((c_adm, noise_level_emb), 1)
|
||||
|
||||
return adm_out
|
||||
|
||||
def tokenize_with_weights(self, text:str, return_word_ids=False, _tokens=[], return_tokens=False):
|
||||
'''
|
||||
Takes a prompt and converts it to a list of (token, weight, word id) elements.
|
||||
Tokens can both be integer tokens and pre computed CLIP tensors.
|
||||
Word id values are unique per word and embedding, where the id 0 is reserved for non word tokens.
|
||||
Returned list has the dimensions NxM where M is the input size of CLIP
|
||||
'''
|
||||
if self.pad_with_end:
|
||||
pad_token = self.end_token
|
||||
else:
|
||||
pad_token = 0
|
||||
|
||||
tokens = _tokens
|
||||
if len(_tokens) == 0:
|
||||
text = escape_important(text)
|
||||
parsed_weights = token_weights(text, 1.0)
|
||||
|
||||
#tokenize words
|
||||
tokens = []
|
||||
for weighted_segment, weight in parsed_weights:
|
||||
to_tokenize = unescape_important(weighted_segment).replace("\n", " ").split(' ')
|
||||
to_tokenize = [x for x in to_tokenize if x != ""]
|
||||
for word in to_tokenize:
|
||||
#if we find an embedding, deal with the embedding
|
||||
if word.startswith(self.embedding_identifier) and self.embedding_directory is not None:
|
||||
embedding_name = word[len(self.embedding_identifier):].strip('\n')
|
||||
embed, leftover = self._try_get_embedding(embedding_name)
|
||||
if embed is None:
|
||||
print(f"warning, embedding:{embedding_name} does not exist, ignoring")
|
||||
else:
|
||||
if len(embed.shape) == 1:
|
||||
tokens.append([(embed, weight)])
|
||||
else:
|
||||
tokens.append([(embed[x], weight) for x in range(embed.shape[0])])
|
||||
#if we accidentally have leftover text, continue parsing using leftover, else move on to next word
|
||||
if leftover != "":
|
||||
word = leftover
|
||||
else:
|
||||
continue
|
||||
#parse word
|
||||
tokens.append([(t, weight) for t in self.tokenizer(word)["input_ids"][self.tokens_start:-1]])
|
||||
if return_tokens: return tokens
|
||||
|
||||
#reshape token array to CLIP input size
|
||||
batched_tokens = []
|
||||
batch = []
|
||||
if self.start_token is not None:
|
||||
batch.append((self.start_token, 1.0, 0))
|
||||
batched_tokens.append(batch)
|
||||
for i, t_group in enumerate(tokens):
|
||||
#determine if we're going to try and keep the tokens in a single batch
|
||||
is_large = len(t_group) >= self.max_word_length
|
||||
|
||||
while len(t_group) > 0:
|
||||
if len(t_group) + len(batch) > self.max_length - 1:
|
||||
remaining_length = self.max_length - len(batch) - 1
|
||||
#break word in two and add end token
|
||||
if is_large:
|
||||
batch.extend([(t,w,i+1) for t,w in t_group[:remaining_length]])
|
||||
batch.append((self.end_token, 1.0, 0))
|
||||
t_group = t_group[remaining_length:]
|
||||
#add end token and pad
|
||||
else:
|
||||
batch.append((self.end_token, 1.0, 0))
|
||||
if self.pad_to_max_length:
|
||||
batch.extend([(pad_token, 1.0, 0)] * (remaining_length))
|
||||
#start new batch
|
||||
batch = []
|
||||
if self.start_token is not None:
|
||||
batch.append((self.start_token, 1.0, 0))
|
||||
batched_tokens.append(batch)
|
||||
else:
|
||||
batch.extend([(t,w,i+1) for t,w in t_group])
|
||||
t_group = []
|
||||
|
||||
#fill last batch
|
||||
batch.append((self.end_token, 1.0, 0))
|
||||
if self.pad_to_max_length:
|
||||
batch.extend([(pad_token, 1.0, 0)] * (self.max_length - len(batch)))
|
||||
|
||||
if not return_word_ids:
|
||||
batched_tokens = [[(t, w) for t, w,_ in x] for x in batched_tokens]
|
||||
|
||||
return batched_tokens
|
||||
|
||||
|
||||
# from diffusers.utils import load_image
|
||||
def load_image(image: Union[str, PIL.Image.Image]) -> PIL.Image.Image:
|
||||
"""
|
||||
Loads `image` to a PIL Image.
|
||||
|
||||
Args:
|
||||
image (`str` or `PIL.Image.Image`):
|
||||
The image to convert to the PIL Image format.
|
||||
Returns:
|
||||
`PIL.Image.Image`:
|
||||
A PIL Image.
|
||||
"""
|
||||
if isinstance(image, str):
|
||||
if image.startswith("http://") or image.startswith("https://"):
|
||||
image = PIL.Image.open(requests.get(image, stream=True).raw)
|
||||
elif os.path.isfile(image):
|
||||
image = PIL.Image.open(image)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Incorrect path or url, URLs must start with `http://` or `https://`, and {image} is not a valid path"
|
||||
)
|
||||
elif isinstance(image, PIL.Image.Image):
|
||||
image = image
|
||||
else:
|
||||
raise ValueError(
|
||||
"Incorrect format used for image. Should be an url linking to an image, a local path, or a PIL image."
|
||||
)
|
||||
image = PIL.ImageOps.exif_transpose(image)
|
||||
image = image.convert("RGB")
|
||||
return image
|
||||
Reference in New Issue
Block a user