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

迁移学习修改ResNet18为二分类模型后计算SHAP值时出现RuntimeError的问题求助

迁移学习修改ResNet18为二分类模型后计算SHAP值时出现RuntimeError的问题求助

我把Resnet18模型改成了预测猫和狗两个类别(原本是1000类),然后想用SHAP库来可视化特征的边际贡献。

预训练的Resnet18在计算SHAP值时完全没问题,但我修改后的模型却抛出了这个错误:

RuntimeError: One of the differentiated Tensors does not require grad

修改后的模型本身是正常工作的,分类猫狗的准确率也不错,问题只在使用SHAP的时候才出现。

训练时我只用了猫狗数据集,但为了方便对比两个模型(感觉问题出在模型结构而非数据),计算SHAP值时我用了shap.datasets.imagenet50()。

下面是我构建模型的代码:

import glob
import warnings
import matplotlib.pyplot as plt
import numpy as np
import torch
from torch import nn
from torchvision.datasets import ImageFolder
from torchvision.transforms import Compose, Normalize, Resize, ToTensor

device = "cuda:0" if torch.cuda.is_available() else "cpu"

!wget https://download.microsoft.com/download/3/E/1/3E1C3F21-ECDB-4869-8368-6DEBA77B919F/kagglecatsanddogs_5340.zip
!unzip -qq kagglecatsanddogs_5340.zip
!rm -rf PetImages/Cat/666.jpg PetImages/Dog/11702.jpg readme\[1\].txt CDLA-Permissive-2.0.pdf

dataset = ImageFolder(
    "./PetImages",
    transform=Compose(
        [
            Resize((224, 224)),
            ToTensor(),
            Normalize((0.5, 0.5, 0.5), (1, 1, 1)),
        ]
    ),
)

train_set, val_set = torch.utils.data.random_split(
    dataset, [int(0.8 * len(dataset)), len(dataset) - int(0.8 * len(dataset))]
)

train_dataloader = torch.utils.data.DataLoader(train_set, batch_size=256, shuffle=True)
val_dataloader = torch.utils.data.DataLoader(val_set, batch_size=256, shuffle=False)

import pytorch_lightning as pl
from torchmetrics.functional import accuracy
from torchvision.models import resnet18
from torchmetrics.functional.classification import multiclass_accuracy

model = resnet18(pretrained=True)

class CatsVSDogsResnet(pl.LightningModule):
    def __init__(self, pretrained: bool = False) -> None:
        super().__init__()
        self.pretrained = pretrained

        if pretrained:
            # <YOUR CODE HERE>
            self.model = resnet18(pretrained=True)
            self.model.fc = nn.Identity()
            self.classifier = nn.Linear(512, 2)
            self.optimizer = torch.optim.Adam(self.classifier.parameters())
        else:
            # <YOUR CODE HERE>
            self.model = resnet18(pretrained=False)
            self.optimizer = torch.optim.Adam(self.model.parameters())

        self.loss = nn.CrossEntropyLoss()

    def forward(self, x) -> torch.Tensor:
        if self.pretrained:
            # <YOUR CODE HERE>
            with torch.no_grad():
                features = self.model(x)
            preds = self.classifier(features)
        else:
            # <YOUR CODE HERE>
            preds = self.model(x)
        return preds

    def configure_optimizers(self):
        return self.optimizer

    def training_step(self, train_batch, batch_idx) -> torch.Tensor:
        images, target = train_batch
        preds = self.forward(images)
        loss = self.loss(preds, target)
        self.log("train_loss", loss, prog_bar=True)
        return loss

    def validation_step(self, val_batch, batch_idx) -> None:
        images, target = val_batch
        preds = self.forward(images)
        loss = self.loss(preds, target)
        acc =  multiclass_accuracy(torch.argmax(preds, dim=-1).long(), target.long(), num_classes=2)
        self.log("val_loss", loss, prog_bar=True)
        self.log("accuracy", acc, prog_bar=True)

cats_vs_dogs_pretrained = CatsVSDogsResnet(pretrained=True)
trainer = pl.Trainer(accelerator="gpu", max_epochs=1)
trainer.fit(cats_vs_dogs_pretrained, train_dataloader, val_dataloader)

然后是SHAP相关的代码,错误就出在最后一行:

import json
import shap

mean = [0.485, 0.456, 0.406]
std = [0.229, 0.224, 0.225]

def normalize(image):
    if image.max() > 1:
        image /= 255
    image = (image - mean) / std
    # in addition, roll axes so that they suit pytorch
    return torch.tensor(image.swapaxes(-1, 1).swapaxes(2, 3)).float()

#model = resnet18(pretrained=True).eval()
model = cats_vs_dogs_pretrained.eval()

X, y = shap.datasets.imagenet50()
X /= 255
to_explain = X[[1, 41]]

# load the ImageNet class names
url = "https://s3.amazonaws.com/deep-learning-models/image-models/imagenet_class_index.json"
fname = shap.datasets.cache(url)
with open(fname) as f:
    class_names = json.load(f)

e = shap.GradientExplainer((model, model.model.layer1[0].conv2), normalize(X))
# while using resnet18 change to model.layer1[0].conv2

shap_values, indexes = e.shap_values(normalize(to_explain), ranked_outputs=2, nsamples=2)
#↑↑↑ Here the error emerges.

有没有大佬能帮我看看问题出在哪?怎么解决这个SHAP计算时的RuntimeError?

备注:内容来源于stack exchange,提问作者Katya Pleshakova

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.22 16:03:06