grad cam
This commit is contained in:
@@ -1,14 +1,16 @@
|
||||
from .node import LoadTagger, PredictTag
|
||||
from .node import LoadTagger, PredictTag, GradCam
|
||||
from ... import SYMBOL, NODE_SURFIX
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
f"LoadTagger{NODE_SURFIX}": LoadTagger,
|
||||
f"PredictTag{NODE_SURFIX}": PredictTag,
|
||||
f"GradCam{NODE_SURFIX}": GradCam,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
f"LoadTagger{NODE_SURFIX}": f"Load Tagger {SYMBOL}",
|
||||
f"PredictTag{NODE_SURFIX}": f"Predict Tag {SYMBOL}",
|
||||
f"GradCam{NODE_SURFIX}": f"Grad Cam {SYMBOL}",
|
||||
}
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
|
||||
@@ -2,6 +2,8 @@ from .preprocess import preprocess
|
||||
import timm
|
||||
import pandas as pd
|
||||
import torch
|
||||
import matplotlib.pyplot as plt
|
||||
|
||||
from ... import ROOT_NAME
|
||||
|
||||
CATEGORY_NAME = ROOT_NAME + "wd-tagger"
|
||||
@@ -31,6 +33,7 @@ class LoadTagger:
|
||||
|
||||
CATEGORY = CATEGORY_NAME
|
||||
|
||||
@torch.inference_mode(False)
|
||||
def load_tagger(self, tagger, dtype):
|
||||
|
||||
if self.loaded_model_name != tagger:
|
||||
@@ -56,15 +59,17 @@ class PredictTag:
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("BATCH_STRING", "STRING")
|
||||
RETURN_TYPES = ("BATCH_STRING", "STRING", "WD-TAGGER-FEATURES")
|
||||
FUNCTION = "predict_tag"
|
||||
CATEGORY = CATEGORY_NAME
|
||||
|
||||
@torch.inference_mode(False)
|
||||
def predict_tag(self, tagger, labels, image, rating, character_thereshold, general_thereshold):
|
||||
dtype = tagger.parameters().__next__().dtype
|
||||
image = preprocess(image).to("cuda", dtype=dtype)
|
||||
preprocessed_image = preprocess(image).to("cuda", dtype=dtype)
|
||||
with torch.no_grad():
|
||||
logits = tagger(image)
|
||||
feature = tagger.forward_features(preprocessed_image)
|
||||
logits = tagger.forward_head(feature)
|
||||
probs = logits.sigmoid()
|
||||
probs = probs.cpu().numpy()
|
||||
|
||||
@@ -83,4 +88,64 @@ class PredictTag:
|
||||
prompts.append(prompt)
|
||||
|
||||
string = "\n".join([f"prompt:{i}\n{prompt}" for i, prompt in enumerate(prompts)])
|
||||
return (prompts, string)
|
||||
id_to_tag = labels['name'].to_dict()
|
||||
tag_to_id = {v:k for k,v in id_to_tag.items()}
|
||||
|
||||
features = {
|
||||
"feature": feature,
|
||||
"image": ((preprocessed_image + 1) / 2).flip(1).permute(0, 2, 3, 1).float().cpu(), # なにこれは・・・
|
||||
"tag_to_id": tag_to_id,
|
||||
}
|
||||
|
||||
return (prompts, string, features)
|
||||
|
||||
class GradCam:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"tagger": ("WD_TAGGER",),
|
||||
"features": ("WD-TAGGER-FEATURES",),
|
||||
"target_tag": ("STRING",{"default": ""}),
|
||||
"heat_map_alpha": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", )
|
||||
FUNCTION = "grad_cam"
|
||||
CATEGORY = CATEGORY_NAME
|
||||
|
||||
@torch.inference_mode(False)
|
||||
def grad_cam(self, tagger, features, target_tag, heat_map_alpha):
|
||||
|
||||
image = features["image"]
|
||||
size = (image.shape[1], image.shape[2])
|
||||
target_id = features["tag_to_id"][target_tag.strip().replace(" ", "_")]
|
||||
|
||||
features = features["feature"].detach().clone().requires_grad_(True)
|
||||
|
||||
gradients = []
|
||||
for i in range(len(features)):
|
||||
feature = features[i].unsqueeze(0)
|
||||
outputs = tagger.forward_head(feature).sigmoid()
|
||||
output = outputs[:,target_id]
|
||||
gradients.append(torch.autograd.grad(output, feature, retain_graph=True)[0])
|
||||
tagger.zero_grad()
|
||||
features.grad = None
|
||||
gradients = torch.cat(gradients)
|
||||
|
||||
weight = torch.mean(gradients, dim=1, keepdim=True)
|
||||
heat_map = torch.sum(weight * features, dim=2).relu().reshape(-1, 1, 28, 28)
|
||||
heat_map = heat_map / heat_map.max()
|
||||
|
||||
heat_map = heat_map.permute(0, 2, 3, 1)
|
||||
heat_map = heat_map.reshape(-1, 1).detach().float().cpu().numpy()
|
||||
|
||||
c_map = plt.get_cmap("jet")
|
||||
heat_map = c_map(heat_map).reshape(-1, 28, 28, 4)[:,:,:,:3]
|
||||
heat_map = torch.from_numpy(heat_map)
|
||||
heat_map = heat_map.permute(0, 3, 1, 2)
|
||||
heat_map = torch.nn.functional.interpolate(heat_map, size=size, mode="bilinear", align_corners=False)
|
||||
heat_map = heat_map.permute(0, 2, 3, 1)
|
||||
|
||||
return (image * (1 - heat_map_alpha) + heat_map * heat_map_alpha, )
|
||||
Reference in New Issue
Block a user