这个包是影响函数的即插即用 PyTorch 重新实现。Pang Wei Koh 和 Percy Liang (ICML2017) 在通过影响函数理解黑盒预测的论文中介绍了影响函数。
项目描述
PyTorch 的影响函数
这是 ICML2017 最佳论文中影响函数的 PyTorch 重新实现: Pang Wei Koh 和 Percy Liang 的通过影响函数理解黑盒预测。参考实现可以在这里找到:链接。
为什么要使用影响函数?
影响函数可帮助您根据数据集调试深度学习模型的结果。在测试单个测试图像时,您可以计算哪些训练图像在分类结果上的结果最大。因此,您可以轻松地在数据集中找到错误标记的图像,或者将数据集稍微压缩为对您的个人测试数据集最重要的最有影响力的图像。这可以提高预测准确性、减少训练时间并减少内存需求。有关详细信息,请参阅此处链接的原始论文。
影响函数当然也可以用于图像以外的数据,只要您有监督学习问题。
要求
- Python 3.6 或更高版本
- PyTorch 1.0 或更高版本
- NumPy 1.12 或更高版本
要运行测试,进一步的要求是:
- torchvision 0.3 或更高版本
- 太平船务
安装
您可以直接通过 pip 安装此软件包:
pip3 install --user pytorch-influence-functions
或者,您可以克隆存储库并将其作为包导入到您的PATH.
用法
计算训练数据集的各个样本对最终预测的影响是直截了当的。
让代码运行的最简单的方法是这样的:
import pytorch_influence_functions as ptif
# Supplied by the user:
model = get_my_model()
trainloader, testloader = get_my_dataloaders()
ptif.init_logging()
config = ptif.get_default_config()
influences, harmful, helpful = ptif.calc_img_wise(config, model, trainloader, testloader)
# do someting with influences/harmful/helpful
这里,config包含影响函数计算的默认值,当然可以更改。有关详细信息和示例,请查看此处。
背景和文件
在近似影响时,可以通过使用更多的迭代和/或更多的递归来调整输出的精度。
配置
config是一个字典,其中包含用于计算影响的参数。config您可以通过调用获取默认值ptif.get_default_config()。
我建议您根据自己的喜好更改以下参数。下面的列表分为影响计算的参数和影响其他一切的参数。
其他参数
save_pth:默认None,保存的文件夹和需要保存s_test的grad_z文件outdir: 写入结果 json 文件的文件夹名称log_filename:默认None,如果设置,除了stdout.
计算参数
seed:默认 = 42,numpy、random、pytorch 的随机种子gpu:默认 = -1,-1用于在 CPU 上计算,否则为 GPU idcalc_method:默认 = img_wise,在此处列出的两种计算方法之间进行选择。DataLoader所需数据集的对象train_loader和test_loader
test_sample_start_per_class: 默认 = False,每类索引从哪里开始计算影响函数。如果False,它将从 开始0。如果您想计算整个测试数据集的影响函数并手动将计算拆分到多个线程/机器/GPU 上,这将非常有用。然后,您可以从数据集中的各个点开始。test_sample_num:默认 = False,每个类的样本数从 开始test_sample_start_per_class计算影响函数。例如,如果您的数据集有 10 个类别,并且您将此值设置为1,则将为10 * 1测试样本计算影响函数,每个类别一个。如果False,计算所有图像的影响。
s_test
recursion_depth:默认值 = 5000,s_test计算的递归深度。更大的递归深度提高了精度。r:默认 = 1,s_test取平均值的计算次数。更大的 r 平均提高了精度。- 结合起来,原始论文建议
recursion_depth * r应该等于训练数据集的大小,因此上述值r = 10和recursion_depth = 5000对于训练数据集大小为 50000 个项目的 CIFAR-10 是有效的。 damp:默认 = 0.01,计算时的阻尼系数s_test。scale:默认 = 25,计算期间的比例因子s_test。
计算模式
这个包提供了两种计算模式来计算影响函数。第一种模式称为,在此期间,在计算单个图像的影响时,动态计算每个训练图像calc_img_wise的两个值s_test和。grad_z然后算法移动到下一个图像。调用第二种模式calc_all_grad_then_test并首先计算grad_z所有图像的值并将它们保存到磁盘。然后,它会计算所有s_test值并将其保存到磁盘。随后,该算法将通过从磁盘读取两个值并基于它们计算影响来计算所有图像的影响函数。这可能会占用大量磁盘空间(100 GB),但使用快速 SSD 可以显着加快计算速度,因为不会发生重复计算。之所以如此,是因为grad_z必须计算两次,一次用于第一次近似,s_test一次与s_test
向量结合以计算影响。然而,最重要的是,s_test仅取决于测试样本。虽然grad_z在计算过程中使用一个来估计 Hessian 的初始值s_test,但这并不重要。grad_z另一方面,只依赖于训练样本。因此,在该calc_img_wise模式下,我们丢弃所有grad_z
计算,即使我们可以将它们重用于所有后续s_test
计算,这可能是成千上万的。但是,如上所述,grad_z仅当它们可以更快地加载/保存在 RAM 中而不是动态计算它们时,保留 s 才有意义。
TL;DR:推荐的方法是使用calc_img_wise,除非你有一个疯狂的快速 SSD、大量可用存储空间,并且想要计算对整个数据集甚至超过 1000 个测试样本的预测结果的影响。
输出变量
可视化,输出可能如下所示:
左上角的测试图像是计算影响的测试图像。为了得到ship的正确测试结果,训练数据集中的有用图像是最有帮助的,而有害图像是最有害的。在这里,我们使用 CIFAR-10 作为数据集。该模型是 ResNet-110。图像上方的数字显示了计算出的实际影响值。
下图显示了相同但不同的模型 DenseNet-100/12。因此,我们可以看到不同的模型从不同的图像中学到的东西更多。
影响
是一个 dict/json,包含对每个测试数据样本的所有训练数据样本计算的影响。dict 结构看起来与此类似:
{
"0": {
"label": 3,
"num_in_dataset": 0,
"time_calc_influence_s": 129.6417362689972,
"influence": [
-0.00016939856868702918,
4.3426321099104825e-06,
-9.501376189291477e-05,
...
],
"harmful": [
31527,
5110,
47217,
...
],
"helpful": [
5287,
22736,
3598,
...
]
},
"1": {
"label": 8,
"num_in_dataset": 1,
"time_calc_influence_s": 121.8709237575531,
"influence": [
3.993639438704122e-06,
3.454859779594699e-06,
-3.5805194329441292e-06,
...
有害
有害是一个数字列表,这些数字是按有害程度排序的训练数据样本的 ID。如果计算多张测试图像的影响函数,则按照对处理后的测试样本预测结果的平均危害程度排序。
有帮助
Helpful 是一个数字列表,这些数字是按照帮助性排序的训练数据样本的 ID。如果针对多个测试图像计算影响函数,则帮助度按对处理后的测试样本的预测结果的平均帮助度排序。
路线图
v0.2
- 使变量名等数据集独立
- 从代码中删除所有数据集名称检查
- 禁用 shell 输出的能力,例如
display_progress从配置中 - 添加适当的结果绘图支持
- 添加数据加载器以仅对最具影响力的样本进行训练
- 添加一些结果的可视化
- 添加对原始论文的一些图表的重新创建以验证实现
v0.3
- 使配置成为一个类,以便它可以重新调整自己,例如当
r和recursion_depth值可以降低而不会产生太大影响时 - 检查杀戮数据增强!?
- 在
calc_influence_function.pyin 中load_s_test,load_grad_z不要对文件名进行硬编码
v0.4
- 集成 myPy 类型注释(静态类型检查)
- 使用多处理计算影响
- 使用
r"doc"像 pytorch 这样的文档字符串
项目详情
下载文件
下载适用于您平台的文件。如果您不确定要选择哪个,请了解有关安装包的更多信息。
源分布
pytorch_influence_functions -0.1.1.tar.gz 的哈希值
| 算法 | 哈希摘要 | |
|---|---|---|
| SHA256 | 6ec4730ad2e0f38e5deb34d60979df5439f4c60b92a0e6390bc50f444f749670 |
|
| MD5 | 4630403a99cbf18634f5b4caae1fbad6 |
|
| 布莱克2-256 | 683eb023b7c3fdbfbeced7399f65556b8cc2e07d5e37b915e939deade4ce7c9e |