def gradCAM(x): # Run model and predict logits = model(x) pred = logits.max(-1)[-1] # Returns index of max value (0 or 1) # Fetch activations at final conv layer last_conv = model.model_layers[:5] activations = last_conv(x) # Compute gradients with respect to model's prediction model.zero_grad() logits[0,pred].backward(retain_graph=True) # Compute average gradient per output channel of last conv layer pooled_grads = model.model_layers[3].weight.grad.mean((1,2,3)) # Multiply each output channel with its corresponding average gradient for i in range(activations.shape[1]): activations[:,i,:,:] *= pooled_grads[i] # Compute heatmap as average over all weighted output channels heatmap = torch.mean(activations, dim=1)[0].cpu().detach() return heatmap