Initial commit.

This commit is contained in:
shiimizu
2024-01-18 03:35:47 -08:00
parent bfa92a2892
commit 5b0c054889
6 changed files with 654 additions and 0 deletions
+25
View File
@@ -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}
}
```
+2
View File
@@ -0,0 +1,2 @@
from .photomaker import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+151
View File
@@ -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
View File
@@ -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",
}
+59
View File
@@ -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}
+210
View File
@@ -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