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

如何定位影响神经网络输出指定索引的输入张量索引?

解决方案

一、手动推导(适合简单层/操作)

对于矩阵乘法、线性变换这类有明确数学规则的操作,直接通过计算逻辑就能推导出输入输出的映射关系。比如你提到的12×12矩阵相乘,输出的(0,0)元素是第一个矩阵第0行和第二个矩阵第0列的点积,所以第一个矩阵的(0,0)到(0,11)、第二个矩阵的(0,0)到(11,0)都会影响该输出位置。这种方法高效直观,适合规则明确的基础操作。

二、自动追踪工具(适合复杂神经网络)

如果模型包含多层非线性操作(比如卷积、激活函数、循环层),手动推导会非常繁琐,这时可以用梯度追踪或专门的可解释性工具:

1. PyTorch梯度反向传播法

利用反向传播的梯度特性,给目标输出位置设置梯度为1,其余为0,反向计算输入的梯度——梯度非零的输入位置就是对该输出有贡献的元素。

示例代码:

import torch
import torch.nn as nn

# 示例模型:自定义神经网络(这里以矩阵乘法为例)
class CustomModel(nn.Module):
    def forward(self, x1, x2):
        return torch.matmul(x1, x2)

model = CustomModel()
# 初始化输入张量,开启梯度追踪
x1 = torch.randn(12, 12, requires_grad=True)
x2 = torch.randn(12, 12, requires_grad=True)

output = model(x1, x2)
# 构造梯度掩码:只保留输出(0,0)位置的梯度
grad_mask = torch.zeros_like(output)
grad_mask[0, 0] = 1.0

# 反向传播计算输入梯度
output.backward(gradient=grad_mask)

# 提取有贡献的输入索引
x1_contrib = torch.nonzero(x1.grad).tolist()
x2_contrib = torch.nonzero(x2.grad).tolist()

print("Input1影响输出(0,0)的索引:", x1_contrib)
print("Input2影响输出(0,0)的索引:", x2_contrib)

这个方法适用于任意PyTorch模型,反向传播会自动处理链式法则下的多层映射,准确找出所有影响目标输出的输入元素。

2. 可解释性库辅助

比如Captum这类专注于模型可解释性的库,提供了InputXGradient等封装好的方法,能更便捷地计算输入对指定输出的贡献,还支持批量输入、多输出追踪等场景,无需手动构造梯度掩码。

三、可视化方案

1. 热力图(二维输入)

对于矩阵、图像这类二维输入,可以把输入的梯度值绘制成热力图,直观展示哪些区域对目标输出影响更大。

示例代码:

import matplotlib.pyplot as plt

# 可视化Input1的贡献梯度
plt.figure(figsize=(6,6))
plt.imshow(x1.grad.numpy(), cmap='hot')
plt.title("Input1对输出(0,0)的贡献热力图")
plt.colorbar()
plt.show()

# 可视化Input2的贡献梯度
plt.figure(figsize=(6,6))
plt.imshow(x2.grad.numpy(), cmap='hot')
plt.title("Input2对输出(0,0)的贡献热力图")
plt.colorbar()
plt.show()

热力图中颜色越鲜艳的位置,对应输入元素对目标输出的影响程度越高。

2. 索引列表展示

对于非二维输入(比如一维向量、高维张量),可以直接提取非零梯度的索引,整理成列表或表格形式展示,清晰列出所有有贡献的输入位置。

内容的提问来源于stack exchange,提问作者Inyoung Kim 김인영

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 10:55:21