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

使用PyTorch的Swin Transformer生成Grad-CAM热力图异常求助

针对Swin Transformer Grad-CAM热力图仅显示竖线的解决建议

核心问题根源

Swin Transformer依赖窗口注意力机制生成特征,输出特征的空间排布、维度逻辑和常规CNN差异很大,标准Grad-CAM逻辑直接套用会因为特征图的窗口化结构、维度处理不当导致异常。

具体修复方案

  • 换用合适的目标层
    别选norm1这类归一化层,优先挑Swin模块里输出完整空间特征的层,比如model.features[-1][-1].attn.proj(注意力投影层)或者model.features[-1][-1].mlp.fc2(MLP输出层),这些层的特征保留了正确的空间结构,适合Grad-CAM计算。
  • 修正特征图维度处理
    Swin的部分中间层特征是窗口化张量,检查Grad-CAM代码:
    1. 有没有误把窗口维度当成空间维度做压缩/展开,破坏了特征的空间排布。
    2. 反向传播时,确保梯度计算没有打乱特征图的空间结构,比如窗口维度的梯度未正确还原。
  • 规避窗口注意力的动态操作干扰
    把模型切到eval()模式,临时禁用可能的动态窗口调整(比如shifted window操作),避免特征图在推理时出现异常分割。
  • 调整Grad-CAM的加权逻辑
    标准Grad-CAM是对通道维度求平均加权,针对Swin的窗口特征,要么改成每个窗口内单独加权再合并,要么直接用Grad-CAM++这类对通道权重计算更精细的变体,适配窗口化特征。
  • 核对输入预处理
    确保输入图片的预处理和模型训练时完全一致,示例代码如下:
    from torchvision import transforms
    transform = transforms.Compose([
        transforms.Resize(256),
        transforms.CenterCrop(224),
        transforms.ToTensor(),
        transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
    ])
    
    尺寸不对、归一化参数错误会导致特征图无效,直接生成异常热力图。

验证用代码片段

修改目标层后的极简Grad-CAM核心逻辑:

import torch
import torch.nn.functional as F
from torchvision import models

# 初始化模型
model = models.swin_b(weights='IMAGENET1K_V1')
model.eval()

# 选择正确目标层
target_layers = [model.features[-1][-1].attn.proj]

# 注册钩子捕获特征和梯度
activations = []
grads = []
def forward_hook(module, input, output):
    activations.append(output)
def backward_hook(module, grad_input, grad_output):
    grads.append(grad_output[0])

for layer in target_layers:
    layer.register_forward_hook(forward_hook)
    layer.register_backward_hook(backward_hook)

# 处理输入并推理
img = transform(your_image).unsqueeze(0)
output = model(img)
pred_idx = output.argmax(dim=1)

# 反向传播计算梯度
model.zero_grad()
class_loss = output[0, pred_idx]
class_loss.backward()

# 计算并归一化热力图
activation = activations[0].squeeze()
grad = grads[0].squeeze()
weights = grad.mean(dim=[1,2])
cam = torch.einsum('c,chw->hw', weights, activation)
cam = F.relu(cam)
cam = (cam - cam.min()) / (cam.max() - cam.min())

内容的提问来源于stack exchange,提问作者Hazel

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.23 09:34:54