This commit is contained in:
laksjdjf
2023-07-11 20:35:48 +09:00
committed by GitHub
parent da9a0b7ffc
commit 16aea0ce0b
4 changed files with 269 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
chinchin
+89
View File
@@ -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",
}
+146
View File
@@ -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
+33
View File
@@ -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