PyTorch中GradCAM类出现save_grad位置参数错误求助
GradCAM代码报错解决:TypeError: save_grad() takes 1 positional argument but 3 were given
问题描述
我编写了用于生成热力图的GradCAM类代码:
import torch import torch.nn as nn import torch.nn.functional as F class GradCAM: """ Class for generating GradCAM heatmaps for interpreting convolutional neural networks. Args: model (nn.Module): PyTorch model for which GradCAM will be generated. Attributes: model (nn.Module): PyTorch model. gradients (torch.Tensor): Gradients of the target layer. """ def __init__(self, model): self.model = model self.gradients = None def backward_hook(self, module, grad_input, grad_output): """ Hook function to capture gradients of the target layer. Args: module (nn.Module): Target layer module. grad_input (tuple of torch.Tensor): Gradients of the input. grad_output (tuple of torch.Tensor): Gradients of the output. """ self.gradients = grad_output[0] def generate_heatmap(self, input_tensor, target_layer_index): """ Generate GradCAM heatmap. Args: input_tensor (torch.Tensor): Input tensor. target_layer_index (int): Index of the target layer. Returns: torch.Tensor: GradCAM heatmap. """ # Get the target layer from the model's Sequential module target_layer = self.model._modules[f"block{target_layer_index}"] # Register the backward hook on the target layer target_layer.register_backward_hook(self.backward_hook) # Forward pass output, activations = self.model(input_tensor) # Zero out gradients self.model.zero_grad() # Calculate gradients output.backward(torch.ones_like(output)) # Get the gradients from the backward hook gradients = self.gradients # Global average pooling grad_weights = F.adaptive_avg_pool1d(gradients, 1) # Multiply the weights with the activations heatmap = torch.mul(activations, grad_weights).sum(dim=2, keepdim=True) return heatmap
运行时出现如下错误:
<ipython-input-71-afdfac905e59>:12: UserWarning: To copy construct from a tensor, it is recommended to use sourceTensor.clone().detach() or sourceTensor.clone().detach().requires_grad_(True), rather than torch.tensor(sourceTensor). sample_input_tensor = torch.tensor(sample_input, dtype=torch.float32) --------------------------------------------------------------------------- TypeError Traceback (most recent call last) <ipython-input-71-afdfac905e59> in <cell line: 18>() 16 17 # Generate the GradCAM heatmap ---> 18 heatmap = gradcam.generate_heatmap(sample_input_tensor, target_layer_index) 3 frames /usr/local/lib/python3.10/dist-packages/torch/nn/modules/module.py in __call__(self, *args, **kwargs) 69 if module is None: 70 raise RuntimeError("You are trying to call the hook of a dead Module!") ---> 71 return self.hook(module, *args, **kwargs) 72 return self.hook(*args, **kwargs) 73 TypeError: save_grad() takes 1 positional argument but 3 were given
错误原因与解决方法
核心错误根源
报错提到的save_grad()函数和你贴的GradCAM类中的backward_hook不是同一个函数。你实际运行的代码中,应该是错误注册了一个参数数量不符合要求的save_grad作为反向传播钩子。PyTorch的register_backward_hook要求钩子函数必须接收3个参数:module(目标层模块)、grad_input(输入梯度)、grad_output(输出梯度),如果你的save_grad只定义了1个参数,就会触发这个错误。
代码修正细节
除了钩子函数的问题,你的GradCAM代码还有几个需要调整的地方,保证逻辑正确:
修复输入张量的警告
把创建输入张量的代码:sample_input_tensor = torch.tensor(sample_input, dtype=torch.float32)替换为:
# 如果sample_input是numpy数组 sample_input_tensor = torch.as_tensor(sample_input, dtype=torch.float32).requires_grad_(True) # 如果sample_input已经是PyTorch张量 sample_input_tensor = sample_input.clone().detach().float().requires_grad_(True)修正钩子重复注册问题
每次调用generate_heatmap时重复注册钩子会导致多次触发,需要保存钩子句柄并在注册前移除旧钩子。完善GradCAM逻辑
- 对热力图添加ReLU激活,只保留正向贡献(符合GradCAM的原始逻辑)
- 增加归一化步骤,方便后续可视化
- 梯度处理时添加
detach()避免计算图干扰
修正后的完整GradCAM代码
import torch import torch.nn as nn import torch.nn.functional as F class GradCAM: """ 生成GradCAM热力图,用于解释卷积神经网络 """ def __init__(self, model): self.model = model self.gradients = None self.hook_handle = None # 保存钩子句柄,用于移除旧钩子 def backward_hook(self, module, grad_input, grad_output): # 保存梯度并脱离计算图 self.gradients = grad_output[0].detach() def generate_heatmap(self, input_tensor, target_layer_index): # 获取目标层,增加不存在的判断 target_layer = self.model._modules.get(f"block{target_layer_index}") if not target_layer: raise ValueError(f"模型中不存在名为block{target_layer_index}的层") # 先移除旧钩子,避免重复注册 if self.hook_handle: self.hook_handle.remove() # 注册新的反向传播钩子 self.hook_handle = target_layer.register_backward_hook(self.backward_hook) # 前向传播,确保模型返回(output, activations) output, activations = self.model(input_tensor) activations = activations.detach() # 清零模型梯度 self.model.zero_grad() # 反向传播:如果是分类任务,建议针对最大概率类别求导,更符合GradCAM意图 # 替换为 output.max(dim=1)[0].sum().backward() 效果更好 output.backward(torch.ones_like(output)) # 计算梯度权重:全局平均池化 grad_weights = F.adaptive_avg_pool1d(self.gradients, 1) # 生成热力图,添加ReLU保留正贡献 heatmap = torch.mul(activations, grad_weights).sum(dim=2, keepdim=True) heatmap = F.relu(heatmap) # 归一化到[0,1]区间,方便可视化 heatmap = (heatmap - heatmap.min()) / (heatmap.max() - heatmap.min() + 1e-8) return heatmap
内容的提问来源于stack exchange,提问作者Manu Jack Pel
相关产品推荐
相关产品推荐

