使用PyTorch计算Shap values时绘图遇索引越界错误求助
修复PyTorch中SHAP图像可视化的IndexError问题
问题描述
在PyTorch中使用SHAP分析MNIST模型时,运行图像可视化代码出现IndexError: index 1 is out of bounds for axis 0 with size 1错误。
运行代码
import numpy as np import torch from torch import nn, optim from torch.nn import functional as F from torchvision import datasets, transforms import shap batch_size = 128 num_epochs = 2 device = torch.device("cpu") class Net(nn.Module): def __init__(self): super().__init__() self.conv_layers = nn.Sequential( nn.Conv2d(1, 10, kernel_size=5), nn.MaxPool2d(2), nn.ReLU(), nn.Conv2d(10, 20, kernel_size=5), nn.Dropout(), nn.MaxPool2d(2), nn.ReLU(), ) self.fc_layers = nn.Sequential( nn.Linear(320, 50), nn.ReLU(), nn.Dropout(), nn.Linear(50, 10), nn.Softmax(dim=1), ) def forward(self, x): x = self.conv_layers(x) x = x.view(-1, 320) x = self.fc_layers(x) return x def train(model, device, train_loader, optimizer, epoch): model.train() for batch_idx, (data, target) in enumerate(train_loader): data, target = data.to(device), target.to(device) optimizer.zero_grad() output = model(data) loss = F.nll_loss(output.log(), target) loss.backward() optimizer.step() if batch_idx % 100 == 0: print( f"Train Epoch: {epoch} [{batch_idx * len(data)}/{len(train_loader.dataset)}" f" ({100.0 * batch_idx / len(train_loader):.0f}%)]" f"\tLoss: {loss.item():.6f}" ) def test(model, device, test_loader): model.eval() test_loss = 0 correct = 0 with torch.no_grad(): for data, target in test_loader: data, target = data.to(device), target.to(device) output = model(data) test_loss += F.nll_loss(output.log(), target).item() # sum up batch loss pred = output.max(1, keepdim=True)[ 1 ] # get the index of the max log-probability correct += pred.eq(target.view_as(pred)).sum().item() test_loss /= len(test_loader.dataset) print( f"\nTest set: Average loss: {test_loss:.4f}," f" Accuracy: {correct}/{len(test_loader.dataset)}" f" ({100.0 * correct / len(test_loader.dataset):.0f}%)\n" ) train_loader = torch.utils.data.DataLoader( datasets.MNIST( "mnist_data", train=True, download=True, transform=transforms.Compose([transforms.ToTensor()]), ), batch_size=batch_size, shuffle=True, ) test_loader = torch.utils.data.DataLoader( datasets.MNIST( "mnist_data", train=False, transform=transforms.Compose([transforms.ToTensor()]) ), batch_size=batch_size, shuffle=True, ) model = Net().to(device) optimizer = optim.SGD(model.parameters(), lr=0.01, momentum=0.5) for epoch in range(1, num_epochs + 1): train(model, device, train_loader, optimizer, epoch) test(model, device, test_loader) # since shuffle=True, this is a random sample of test data batch = next(iter(test_loader)) images, _ = batch #images = images.view(-1, 1, 28, 28) background = images[:100] test_images = images[100:110] e = shap.DeepExplainer(model, background) shap_values = e.shap_values(test_images) shap_numpy = [np.swapaxes(np.swapaxes(s, 1, -1), 1, 2) for s in shap_values] test_numpy = np.swapaxes(np.swapaxes(test_images.numpy(), 1, -1), 1, 2) # plot the feature attributions shap.image_plot(shap_numpy, -test_numpy)
报错信息
Traceback (most recent call last): Cell In[5], line 5 shap.image_plot(shap_numpy, -test_numpy) File ~\anaconda3\lib\site-packages\shap\plots_image.py:154 in image if len(shap_values[0][row].shape) == 2: IndexError: index 1 is out of bounds for axis 0 with size 1
错误原因及修复方案
1. 模型输出层适配SHAP计算
你的模型最后使用了Softmax(dim=1)输出概率值,但SHAP的DeepExplainer更适合处理未归一化的logits或LogSoftmax输出,概率值会导致维度计算异常。
修复步骤:
- 修改模型全连接层,将
Softmax替换为LogSoftmax:self.fc_layers = nn.Sequential( nn.Linear(320, 50), nn.ReLU(), nn.Dropout(), nn.Linear(50, 10), nn.LogSoftmax(dim=1), # 替换原Softmax层 ) - 调整训练和测试的损失计算,去掉
output.log()(因为LogSoftmax已输出log概率):# 训练函数中 loss = F.nll_loss(output, target) # 测试函数中 test_loss += F.nll_loss(output, target).item()
2. 图像维度简化
MNIST是单通道图像,转换为numpy数组后可移除多余的通道维度,避免可视化时的维度冲突:
# 替换原维度转换代码 shap_numpy = [np.squeeze(np.swapaxes(s, 1, -1)) for s in shap_values] test_numpy = np.squeeze(np.swapaxes(test_images.numpy(), 1, -1))
3. 强制模型进入评估模式
在SHAP计算前,显式设置模型为评估模式,避免Dropout等层干扰计算结果:
model.eval() # 添加该行代码 e = shap.DeepExplainer(model, background) shap_values = e.shap_values(test_images)
修复后关键代码片段
# 修改后的模型 class Net(nn.Module): def __init__(self): super().__init__() self.conv_layers = nn.Sequential( nn.Conv2d(1, 10, kernel_size=5), nn.MaxPool2d(2), nn.ReLU(), nn.Conv2d(10, 20, kernel_size=5), nn.Dropout(), nn.MaxPool2d(2), nn.ReLU(), ) self.fc_layers = nn.Sequential( nn.Linear(320, 50), nn.ReLU(), nn.Dropout(), nn.Linear(50, 10), nn.LogSoftmax(dim=1), ) def forward(self, x): x = self.conv_layers(x) x = x.view(-1, 320) x = self.fc_layers(x) return x # 修改后的训练函数损失计算 def train(model, device, train_loader, optimizer, epoch): model.train() for batch_idx, (data, target) in enumerate(train_loader): data, target = data.to(device), target.to(device) optimizer.zero_grad() output = model(data) loss = F.nll_loss(output, target) loss.backward() optimizer.step() # 其余代码不变 # 修改后的测试函数损失计算 def test(model, device, test_loader): model.eval() test_loss = 0 correct = 0 with torch.no_grad(): for data, target in test_loader: data, target = data.to(device), target.to(device) output = model(data) test_loss += F.nll_loss(output, target).item() # 其余代码不变 # 修改后的SHAP可视化部分 model.eval() batch = next(iter(test_loader)) images, _ = batch background = images[:100] test_images = images[100:110] e = shap.DeepExplainer(model, background) shap_values = e.shap_values(test_images) shap_numpy = [np.squeeze(np.swapaxes(s, 1, -1)) for s in shap_values] test_numpy = np.squeeze(np.swapaxes(test_images.numpy(), 1, -1)) shap.image_plot(shap_numpy, -test_numpy)
内容的提问来源于stack exchange,提问作者Bea González
相关产品推荐
相关产品推荐

