joenathan's picture
Update docstrings for the renamed notebook
95cdfc0 verified
Raw
History Blame Contribute Delete
3.92 kB
"""Shared model/Grad-CAM code, extracted from
mit-group8-explainable-defect-detection.ipynb (sections 3 and 6) so the
Space app and the training script can both import it without duplication."""
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision.transforms as T
import matplotlib
CLASS_NAMES = [
"crazing",
"inclusion",
"patches",
"pitted_surface",
"rolled-in_scale",
"scratches",
]
IMAGE_SIZE = 128
MEAN = (0.5, 0.5, 0.5)
STD = (0.25, 0.25, 0.25)
eval_transform = T.Compose([
T.Resize((IMAGE_SIZE, IMAGE_SIZE)),
T.ToTensor(),
T.Normalize(MEAN, STD),
])
_mean_t = torch.tensor(MEAN).view(3, 1, 1)
_std_t = torch.tensor(STD).view(3, 1, 1)
def to_image(tensor):
return (tensor.detach().cpu() * _std_t + _mean_t).clamp(0, 1).permute(1, 2, 0).numpy()
class ConvBlock(nn.Module):
def __init__(self, in_channels, out_channels, stride):
super().__init__()
self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride,
padding=1, bias=False)
self.norm = nn.BatchNorm2d(out_channels)
self.activation = nn.ReLU()
def forward(self, x):
return self.activation(self.norm(self.conv(x)))
class DefectCNN(nn.Module):
def __init__(self, number_of_classes=6):
super().__init__()
self.features = nn.Sequential(
ConvBlock(3, 24, stride=2),
ConvBlock(24, 48, stride=2),
ConvBlock(48, 96, stride=2),
ConvBlock(96, 128, stride=1),
)
self.pool = nn.AdaptiveAvgPool2d(1)
self.classifier = nn.Sequential(
nn.Dropout(0.20),
nn.Linear(128, number_of_classes),
)
def forward(self, x):
features = self.features(x)
return self.classifier(self.pool(features).flatten(1))
class GradCAM:
def __init__(self, model, target_layer):
self.model = model
self.target_layer = target_layer
self.activations = None
self.gradients = None
self.forward_handle = target_layer.register_forward_hook(self._capture_activations)
def _capture_activations(self, module, inputs, output):
self.activations = output
if output.requires_grad:
output.register_hook(self._capture_gradients)
def _capture_gradients(self, gradient):
self.gradients = gradient
def __call__(self, images, target_classes):
self.model.eval()
self.model.zero_grad(set_to_none=True)
with torch.enable_grad():
logits = self.model(images)
target_classes = target_classes.to(logits.device, dtype=torch.long)
selected_logits = logits.gather(1, target_classes[:, None]).sum()
selected_logits.backward()
weights = self.gradients.mean(dim=(2, 3), keepdim=True)
cam = torch.relu((weights * self.activations).sum(dim=1, keepdim=True))
cam = F.interpolate(cam, size=images.shape[-2:], mode="bilinear", align_corners=False)
flat = cam.flatten(start_dim=1)
minimum = flat.min(dim=1).values[:, None, None, None]
maximum = flat.max(dim=1).values[:, None, None, None]
cam = (cam - minimum) / (maximum - minimum + 1e-8)
return cam[:, 0].detach(), logits.detach()
def close(self):
self.forward_handle.remove()
def overlay_heatmap(image_tensor, heatmap, heatmap_weight=0.45):
image = to_image(image_tensor)
color = matplotlib.colormaps["inferno"](heatmap.detach().cpu().numpy())[..., :3]
return np.clip((1 - heatmap_weight) * image + heatmap_weight * color, 0, 1)
def load_model(checkpoint_path, device="cpu"):
model = DefectCNN(number_of_classes=len(CLASS_NAMES)).to(device)
state_dict = torch.load(checkpoint_path, map_location=device)
model.load_state_dict(state_dict)
model.eval()
return model