You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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
---&gt; 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!")
---&gt; 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代码还有几个需要调整的地方,保证逻辑正确:

  1. 修复输入张量的警告
    把创建输入张量的代码:

    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)
    
  2. 修正钩子重复注册问题
    每次调用generate_heatmap时重复注册钩子会导致多次触发,需要保存钩子句柄并在注册前移除旧钩子。

  3. 完善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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.24 03:20:11