a
This commit is contained in:
@@ -0,0 +1 @@
|
||||
chinchin
|
||||
@@ -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
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user