From bcd892c0f730d40ff576d38da14d64822e1ca812 Mon Sep 17 00:00:00 2001 From: laksjdjf Date: Mon, 25 Mar 2024 20:33:12 +0900 Subject: [PATCH] grad cam --- scripts/wd-tagger/__init__.py | 4 +- scripts/wd-tagger/node.py | 73 +++++++++++++++++++++++++++++++++-- 2 files changed, 72 insertions(+), 5 deletions(-) diff --git a/scripts/wd-tagger/__init__.py b/scripts/wd-tagger/__init__.py index 63eca43..e523024 100644 --- a/scripts/wd-tagger/__init__.py +++ b/scripts/wd-tagger/__init__.py @@ -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"] diff --git a/scripts/wd-tagger/node.py b/scripts/wd-tagger/node.py index a161878..830eff7 100644 --- a/scripts/wd-tagger/node.py +++ b/scripts/wd-tagger/node.py @@ -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) \ No newline at end of file + 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, ) \ No newline at end of file