diff --git a/models/put_models_here.txt b/models/put_models_here.txt new file mode 100644 index 0000000..eb5ef87 --- /dev/null +++ b/models/put_models_here.txt @@ -0,0 +1 @@ +chinchin \ No newline at end of file diff --git a/pfg.py b/pfg.py new file mode 100644 index 0000000..bf6021f --- /dev/null +++ b/pfg.py @@ -0,0 +1,89 @@ +import os +import torch +import numpy as np +from PIL import Image +from .pfg_model import ViT +from .pfg_utils import download, preprocess_image, TAGGER_FILE + +CURRENT_DIR = os.path.dirname(os.path.realpath(__file__)) + +def get_file_list(path): + return [file for file in os.listdir(path) if file != "put_models_here.txt"] + +class PFG: + def __init__(self): + download(CURRENT_DIR) + self.tagger = ViT(3, 448, 9083) + self.tagger.load_state_dict(torch.load(os.path.join(CURRENT_DIR, TAGGER_FILE))) + self.tagger.eval() + + # wd-14-taggerの推論関数 + @torch.no_grad() + def infer(self, img: Image): + img = preprocess_image(img) + img = torch.tensor(img).permute(2, 0, 1).unsqueeze(0) + print("inferencing by torch model.") + probs = self.tagger(img).squeeze(0) + return probs + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "positive": ("CONDITIONING",), + "negative": ("CONDITIONING",), + "pfg_scale": ("FLOAT", { + "default": 1, + "min": 0, #Minimum value + "max": 2, #Maximum value + "step": 0.05 #Slider's step + }), + "attention_dim": ((768, 1024, 2048), ), + "image": ("IMAGE", ), + "model_name": (get_file_list(os.path.join(CURRENT_DIR,"models")), ), + } + } + RETURN_TYPES = ("CONDITIONING", "CONDITIONING") + FUNCTION = "add_pfg" + CATEGORY = "loaders" + + def add_pfg(self, positive, negative, pfg_scale, image, attention_dim, model_name): + # load weight + pfg_weight = torch.load(os.path.join(CURRENT_DIR, "models/" + model_name)) + weight = pfg_weight["pfg_linear.weight"].cpu() + bias = pfg_weight["pfg_linear.bias"].cpu() + + # comfyのload imageはtensorを返すので一度pillowに戻す + tensor = image*255 + tensor = np.array(tensor, dtype=np.uint8) + image = Image.fromarray(tensor[0]) + + # tagger特徴量の計算 + pfg_feature = self.infer(image) + + # pfgの計算 + pfg_cond = (weight @ pfg_feature + bias) * pfg_scale + pfg_cond = pfg_cond.reshape(1, -1, attention_dim) + + # text_embs + cond = positive[0][0] + uncond = negative[0][0] + + # cond側 + pfg_cond = pfg_cond.to(cond.device, dtype=cond.dtype) + pfg_cond = pfg_cond.repeat(cond.shape[0], 1, 1) + + # uncond側はゼロベクトルでパディング + cond = torch.cat([cond, pfg_cond], dim=1) + pfg_uncond_zero = torch.zeros(uncond.shape[0], pfg_cond.shape[1], uncond.shape[2]).to(uncond.device, dtype=uncond.dtype) + uncond = torch.cat([uncond, pfg_uncond_zero], dim=1) + + return ([[cond, positive[0][1]]], [[uncond, negative[0][1]]]) + +NODE_CLASS_MAPPINGS = { + "PFG": PFG +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "PFG": "Load PFG node", +} \ No newline at end of file diff --git a/pfg_model.py b/pfg_model.py new file mode 100644 index 0000000..c2de708 --- /dev/null +++ b/pfg_model.py @@ -0,0 +1,146 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F + +class StochDepth(nn.Module): + """Batchwise Dropout used in EfficientNet, optionally sans rescaling.""" + + def __init__(self, drop_rate, scale_by_keep=False): + super().__init__() + self.drop_rate = drop_rate + self.scale_by_keep = scale_by_keep + + def forward(self, x): + if not self.training: + return x + + batch_size = x.shape[0] + r = torch.rand([batch_size, 1, 1], dtype=x.dtype, device=x.device) + keep_prob = 1.0 - self.drop_rate + binary_tensor = torch.floor(keep_prob + r) + if self.scale_by_keep: + x = x / keep_prob + return x * binary_tensor + +class PosEmbed(nn.Module): + def __init__(self, input_size): + super().__init__() + + self.pos_embed = nn.Parameter( + torch.empty(input_size, dtype=torch.float32) + ) + torch.nn.init.trunc_normal_(self.pos_embed,mean=0.0, std=0.02) + + def forward(self, x): + return x + self.pos_embed.unsqueeze(0) + +class MLPBlock(nn.Module): + def __init__(self, input_dim, mlp_dim, stochdepth_rate): + super().__init__() + self.stochdepth_rate = stochdepth_rate + self.fc1 = nn.Linear(input_dim, mlp_dim) + self.fc2 = nn.Linear(mlp_dim, input_dim) + if stochdepth_rate > 0.0: + self.stochdepth = StochDepth(stochdepth_rate, scale_by_keep=True) + else: + self.stochdepth = None + + def forward(self, x): + out = F.gelu(self.fc1(x)) + if self.stochdepth: + out = self.stochdepth(out) + out = self.fc2(out) + return out + +class SkipInitChannelwise(nn.Module): + def __init__(self, channels, init_val=1e-6): + super().__init__() + self.channels = channels + self.init_val = init_val + self.skip = nn.Parameter(torch.ones(channels) * init_val) + + def forward(self, x): + return x * self.skip + +class ViTBlock(nn.Module): + def __init__(self, input_dim, heads, key_dim, mlp_dim, layerscale_init, stochdepth_rate): + super().__init__() + self.norm1 = nn.LayerNorm(input_dim,eps=1e-3) + self.attn = nn.MultiheadAttention(input_dim, heads, batch_first=True) + self.skip1 = SkipInitChannelwise(key_dim, init_val=layerscale_init) + self.stochdepth1 = StochDepth(stochdepth_rate, scale_by_keep=True) if stochdepth_rate > 0.0 else None + self.norm2 = nn.LayerNorm(input_dim,eps=1e-3) + self.mlp = MLPBlock(input_dim, mlp_dim, stochdepth_rate) + self.skip2 = SkipInitChannelwise(key_dim, init_val=layerscale_init) + self.stochdepth2 = StochDepth(stochdepth_rate, scale_by_keep=True) if stochdepth_rate > 0.0 else None + + def forward(self, x): + out = self.norm1(x) + out = self.attn(out, out, out)[0] + out = self.skip1(out) + if self.stochdepth1: + out = self.stochdepth1(out) + x = out + x + + out = self.norm2(x) + out = self.mlp(out) + out = self.skip2(out) + if self.stochdepth2: + out = self.stochdepth2(out) + + out = out + x + return out + +class ViT(nn.Module): + def __init__(self, in_channels=3, img_size=320, out_classes=2000, definition_name="B16"): + super().__init__() + self.definitions = { + "B16": { + "num_blocks": 12, + "patch_size": 16, + "key_dim": 768, + "mlp_dim": 3072, + "heads": 12, + "stochdepth_rate": 0.05, + }, + # Other definitions removed for simplicity + } + + definition = self.definitions[definition_name] + self.blocks = nn.ModuleList() + num_blocks = definition["num_blocks"] + patch_size = definition["patch_size"] + key_dim = definition["key_dim"] + mlp_dim = definition["mlp_dim"] + heads = definition["heads"] + stochdepth_rate = definition["stochdepth_rate"] + layerscale_init = 0.1 # Replacing CaiT_LayerScale_init(num_blocks) + + self.conv = nn.Conv2d(in_channels, key_dim, kernel_size=patch_size, stride=patch_size) + self.pos_embed = PosEmbed(((img_size // patch_size) ** 2,key_dim)) + + for i in range(num_blocks): + self.blocks.append( + ViTBlock(key_dim, heads, key_dim, mlp_dim, layerscale_init, stochdepth_rate) + ) + + self.norm = nn.LayerNorm(key_dim,eps=1e-3) + self.avgpool = nn.AdaptiveAvgPool1d(1) + self.fc = nn.Linear(key_dim, out_classes) + self.act = nn.Sigmoid() + + def forward(self, x): + x = (x - 127.5) / 127.5 + x = self.conv(x) + b, c, h, w = x.shape + x = x.view(b, c, h*w).permute(0, 2, 1) # (B, H*W, C) + x = self.pos_embed(x) + for block in self.blocks: + x = block(x) + x = self.norm(x) + x = self.avgpool(x.transpose(1, 2)).squeeze(-1) # (B, C) + + # pfg uses output of last pooling layer. + #x = self.fc(x) + #x = self.act(x) + return x \ No newline at end of file diff --git a/pfg_utils.py b/pfg_utils.py new file mode 100644 index 0000000..8cb1d6a --- /dev/null +++ b/pfg_utils.py @@ -0,0 +1,33 @@ +#このコードはhttps://github.com/kohya-ss/sd-scripts/blob/main/finetune/tag_images_by_wd14_tagger.pyを参考にしていますというかパクっています。 + +from huggingface_hub import hf_hub_download +import cv2 +import numpy as np +import os +IMAGE_SIZE = 448 + +TAGGER_REPO = "furusu/wd-v1-4-tagger-pytorch" +TAGGER_FILE = "wd-v1-4-vit-tagger-v2.ckpt" + +def download(path): + if not os.path.exists(os.path.join(path, TAGGER_FILE)): + hf_hub_download(TAGGER_REPO, TAGGER_FILE, cache_dir=path, force_download=True, force_filename=TAGGER_FILE) + +def preprocess_image(image): + image = image.convert("RGB") + image = np.array(image) + image = image[:, :, ::-1] # RGB->BGR + + # pad to square + size = max(image.shape[0:2]) + pad_x = size - image.shape[1] + pad_y = size - image.shape[0] + pad_l = pad_x // 2 + pad_t = pad_y // 2 + image = np.pad(image, ((pad_t, pad_y - pad_t), (pad_l, pad_x - pad_l), (0, 0)), mode="constant", constant_values=255) + + interp = cv2.INTER_AREA if size > IMAGE_SIZE else cv2.INTER_LANCZOS4 + image = cv2.resize(image, (IMAGE_SIZE, IMAGE_SIZE), interpolation=interp) + + image = image.astype(np.float32) + return image \ No newline at end of file