Grad-CAM Wafer Defect Classifier
Training a CNN backbone to classify spatial wafer defect topologies, then using gradient-weighted class activation mapping to show a process engineer why a wafer was flagged.
Automated Yield Engineering & Defect Classification
In short: the shape a defect makes on a wafer map tells you what went wrong upstream — a scratch, a ring, a random cluster all point to different causes. Engineers currently read that shape by eye. I trained a CNN to do it instead.
Wafer inspection maps show defects in recognizable shapes — scratch lines, center-rings, donuts, random clusters — and each shape is a diagnostic clue: photolithography misalignment, wet-bench contamination, plasma-etch uniformity drift. Reading those shapes by eye is slow and inconsistent, so this project trains a convolutional neural network to classify them automatically.
Model Architecture & ResNet-18 Backbone
The classifier is a ResNet-18 backbone with a Convolutional Block Attention Module (CBAM) added for sharper feature extraction. Defect classes are heavily imbalanced, so training uses focal loss to keep easy background patterns from drowning out the rarer, subtler defect shapes. Rotational and translational augmentation at test time keeps classification stable regardless of how the wafer map happens to be oriented.
Explainable AI with Grad-CAM
A model that’s simply “accurate” isn’t enough on a fab floor — an engineer needs to trust why it flagged a wafer, not just that it did. Gradient-weighted Class Activation Mapping (Grad-CAM) answers that: it shows which pixels actually drove the decision.
Grad-CAM computes the gradients of the target class score flowing into the final convolutional feature maps , generating coarse localization heatmaps that highlight the specific spatial regions driving the network’s classification decision:
By upsampling and overlaying these heatmaps onto the raw wafer inspection maps, we verified that the network learned physically meaningful spatial features rather than shortcut artifacts — attention concentrated along the wafer edge for scratch-pattern dies, radially outward for donut patterns, and tightly on the anomalous cluster itself for random-defect maps, instead of on sensor noise or die-boundary padding.
Implementation: Hooking the Target Convolutional Layer
Grad-CAM requires the raw activations and their gradients from a chosen convolutional layer, neither of which PyTorch retains by default mid-network. We registered a forward hook and a full backward hook directly on the ResNet-18 backbone’s final convolutional block — the last layer with meaningful spatial resolution before global average pooling collapses it:
class GradCAM:
def __init__(self, model, target_layer):
self.model = model
self.target_layer = target_layer
self.gradients = None
self.activations = None
self.target_layer.register_forward_hook(self.save_activations)
self.target_layer.register_full_backward_hook(self.save_gradients)
def save_activations(self, module, input, output):
self.activations = output
def save_gradients(self, module, input, output):
self.gradients = output[0]
The forward hook fires during the model’s normal forward pass and caches . The backward hook fires when gradients flow back through that same layer during .backward(), capturing without needing to modify the model’s forward() method or unroll the backbone.
From Gradients to a Localization Map
generate_heatmap runs a forward pass, isolates the score for the target class, and backpropagates only that score — everything downstream of the target layer is discarded, so the gradient signal reaching target_layer is specific to :
def generate_heatmap(self, input_tensor, class_idx):
self.model.eval()
output = self.model(input_tensor)
self.model.zero_grad()
score = output[0, class_idx]
score.backward()
weights = torch.mean(self.gradients, dim=(2, 3), keepdim=True)
cam = torch.sum(weights * self.activations, dim=1).squeeze()
cam = F.relu(cam)
cam = cam.detach().cpu().numpy()
cam = cv2.resize(cam, (224, 224))
if cam.max() > cam.min():
cam = (cam - cam.min()) / (cam.max() - cam.min() + 1e-8)
else:
cam = cam
return cam
Each line maps onto the equations above: averaging the gradients gives the per-channel weight , weighting and summing the activations gives , the ReLU keeps only features that push toward the target class, and the resize/normalize steps just get the result back to the wafer map’s resolution and into a displayable range.
Key Takeaways
Overlaying the heatmap caught a failure mode accuracy alone would have missed — attention needs to land on the actual defect, not a scanner artifact along the wafer edge. Using hooks instead of modifying the model also meant the same GradCAM wrapper works against any convolutional backbone unchanged.