File size: 1,817 Bytes
961cf0c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 | import torch
import torch.nn.functional as F
# ==========================================================
# BASE CAM
# ==========================================================
class BaseCAM:
def __init__(self, model, target_layer):
self.model = model
self.target_layer = target_layer
self.activations = None
target_layer.register_forward_hook(self._forward_hook)
def _forward_hook(self, module, inp, out):
self.activations = out
def _normalize(self, cam):
cam = F.relu(cam)
return (cam - cam.min()) / (cam.max() - cam.min() + 1e-8)
# ==========================================================
# GradCAM++
# ==========================================================
class GradCAMPlusPlus(BaseCAM):
def generate(self, image, metadata, target_class):
image.requires_grad_(True)
output = self.model(image, metadata)
score = output[:, target_class]
grads = torch.autograd.grad(
score,
self.activations,
retain_graph=True,
create_graph=True
)[0]
grads2 = grads ** 2
grads3 = grads ** 3
denominator = (
2 * grads2 +
torch.sum(self.activations * grads3,
dim=(2,3), keepdim=True) + 1e-8
)
alpha = grads2 / denominator
weights = torch.sum(alpha * F.relu(grads),
dim=(2,3), keepdim=True)
cam = torch.sum(weights * self.activations,
dim=1, keepdim=True)
cam = self._normalize(cam)
cam = F.interpolate(
cam,
size=image.shape[-2:],
mode="bilinear",
align_corners=False
)
return cam.squeeze().detach().cpu().numpy() |