- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
本文是 google-research 仓库中 interpretability_benchmark 目录的完整技术指南。该目录实现了ROAR(RemOve And Retrain,移除并重训练)基准,用于评估深度神经网络中各类可解释性方法(特征重要性估计器)的近似准确度。读完本文,你将掌握 ROAR 的评测原理、从 TFRecord 数据集到显著性热力图生成再到 ResNet-50 重训练评估的完整实验流程,以及全部命令行参数与源码级实现细节。
ROAR 是什么:tl;dr 核心思想
深度神经网络的可解释性方法(interpretability methods)通常以"特征重要性估计"的形式回答一个问题:输入图像中每个像素对模型预测的贡献有多大。这些估计以像素级别的排序(ranking)呈现,例如集成梯度(Integrated Gradients)、敏感度热力图(Sensitivity Heatmaps)、引导反向传播(Guided Backprop)等。
在医疗、自动驾驶、信用评分等敏感领域,可解释性估计必须同时满足两个要求:
- 对人类有意义(meaningful to a human);
- 高度准确(highly accurate)——因为对模型行为的错误解释可能给人类福祉带来难以承受的代价。
ROAR 专注于第 2 点,即度量特征重要性估计器的近似准确度(approximate accuracy)。其核心思想非常直接:按照某个估计器给出的重要度排序,移除(RemOve)被判定为最重要的那一部分输入特征,然后对修改后的数据集**重新训练(Retrain)**模型,观察模型精度的变化。最准确的估计器,应当识别出"移除后对模型性能损害最大"的输入。也就是说,一个估计器是否优秀,取决于它能否精准定位那些真正支撑模型预测的像素——移除它们造成的性能损失应当远大于其他估计器。
从源码结构看,该目录的完整评测流水线分为三个环节,对应三个核心模块:
| 环节 | 脚本 | 作用 |
|---|---|---|
| 特征重要性估计生成 | saliency_data_gen/dataset_generator.py | 为 TFRecord 数据集中每张图像生成显著性热力图,产出新的 TFRecord |
| 数据修改与预处理 | data_input.py | 按估计器排序将最/不重要像素替换为均值,供训练/评估使用 |
| 模型重训练与评估 | train_resnet.py | 在修改后的数据集上重训练 ResNet-50 并输出 top-1/top-5 准确率 |
实验设计:为什么"移除 + 重训练"能度量准确度
理解 ROAR 之前,需要先厘清它区别于"可视化验证"的评测逻辑。通常的显著性方法验证是定性观察热力图是否"对齐"物体轮廓;而 ROAR 给出的是定量基准:
- 每个估计器产出一张显著性排序图;
- 按阈值移除排序最高的
threshold%像素(将其替换为全局均值,而非置零); - 在修改后的数据集上从零重训练模型;
- 记录重训练模型的 top-1/top-5 准确率作为该估计器的得分。
移除像素后模型性能下降越多,说明被移除的像素对预测越关键,该估计器的定位越准确。这一设计规避了一个常见陷阱:仅用"删掉像素后原模型输出是否改变"来评估会受模型平滑度、梯度饱和等干扰,而重训练要求模型必须真正依赖这些像素才能恢复精度。
需要强调的是,ROAR 采用的控制变量组设计(详见下文transformation参数):除了按显著性排序移除像素的modified_image之外,还包含未修改原图(raw_image)和随机移除像素(random_baseline)两种对照,用以排除"图像本身被破坏导致精度下降"这一混淆因素。
第一步:准备数据集并转换为 TFRecord
论文(A Benchmark for Interpretability Methods in Deep Neural Networks)在三个公开图像分类数据集上评估了模型解释的准确性:
- ImageNet(1000 类)
- Birdsnap(500 类)
- Food101(101 类)
要复现实验结果,首先下载目标数据集,并将其转换为 TFRecord 格式。README 建议参考 TensorFlow models 仓库中build_image_data.py一类的转换脚本,将原始图片转换为包含image/encoded(编码图像)与image/class/label(类别标签)字段的 TFRecord shard。
后续所有步骤(显著性图生成、训练、评估)都直接读取这些 TFRecord,因此数据格式的一致性至关重要。从 saliency_data_gen/data_helper.py 的parser实现可以看到,生成显著性图阶段读取的字段正是:
features={ 'image/encoded': tf.FixedLenFeature([], tf.string, default_value=''), 'image/class/label': (tf.FixedLenFeature([], tf.int64)), }其中 ImageNet 的标签在读取时会减去 1,使类别落在[0, 1000)区间(data_helper.py),与N_CLASSES = {'imagenet': 1000, 'food_101': 101, 'birdsnap': 500}的分类数设置对应(dataset_generator.py)。
第二步:生成特征重要性估计(显著性热力图)
saliency_data_gen/dataset_generator.py 为 TFRecord 数据集中每一张图像生成特征重要性估计。特征重要性估计本质上是"每个输入像素对模型预测贡献"的排序,这些属于训练后(post-training)可解释性方法,因此必须提供一个训练好的模型 checkpoint才能生成估计数据集。
支持的显著性方法
脚本基于 saliency 库实现,覆盖的方法包括三种基础方法及其 SmoothGrad、平方(squared)、SmoothGrad²、VarGrad 变体:
方法枚举值(--saliency_method) | 含义 | 底层调用 |
|---|---|---|
IG | 集成梯度 Integrated Gradients | IntegratedGradients.GetMask |
IG_SG | IG + SmoothGrad(非平方) | GetSmoothedMask(magnitude=False) |
IG_SG_2 | IG + SmoothGrad²(平方) | GetSmoothedMask(magnitude=True) |
SH | 敏感度热力图 Gradient Saliency | GradientSaliency.GetMask |
SH_SG | SH + SmoothGrad | GetSmoothedMask(magnitude=False) |
SH_SG_2 | SH + SmoothGrad² | GetSmoothedMask(magnitude=True) |
GB | 引导反向传播 Guided Backprop | GuidedBackprop.GetMask |
GB_SG | GB + SmoothGrad | GetSmoothedMask(magnitude=False) |
GB_SG_2 | GB + SmoothGrad² | GetSmoothedMask(magnitude=True) |
SOBEL | Sobel 边缘算子(非梯度基线) | scipy.ndimage.sobel |
方法分发逻辑在 saliency_helper.py 的generate_saliency_image中实现:get_saliency_image负责在 TensorFlow 计算图上实例化三类 saliency 对象(saliency_helper.py),随后按方法名调用对应的GetMask/GetSmoothedMask。SOBEL比较特殊——它不依赖梯度,直接对预处理图像做ndimage.sobel(img_out, axis=0),作为无监督的边缘基线参与对比(dataset_generator.py)。
运行参数
python -m interpretability_benchmark.saliency_data_gen.dataset_generator \ --data_path=/path/to/tfrecords/ \ --ckpt_path=/path/to/trained/model.ckpt \ --output_dir=/tmp/saliency/ \ --dataset_name=imagenet \ --saliency_method=IG_SG \ --split=validation \ --test_small_sample=False各参数说明(定义于 dataset_generator.py):
--master:TensorFlow master 名称,默认空字符串;--output_dir:TFRecord 输出目录,默认/tmp/saliency/;--data_path:输入 TFRecord 数据集路径;--ckpt_path:训练好的模型 checkpoint 路径;--split:枚举training/validation,指定为训练集还是验证集生成显著性图;--dataset_name:枚举food_101/imagenet/birdsnap,决定类别数N_CLASSES;--saliency_method:上述 10 种方法之一,默认SH_SG;--test_small_sample:布尔值,默认True。置真时仅用合成空图生成 2 个样本验证工作流。
内部处理流程
ProcessSaliencyMaps.produce_saliency_map(dataset_generator.py)展示了完整链路:
- 用
DataIterator读取 TFRecord shard,解析出原始图像与标签; - 按 ImageNet 统计值做标准化:
MEAN_RGB = [0.485*255, 0.456*255, 0.406*255],STDDEV_RGB = [0.229*255, 0.224*255, 0.225*255](dataset_generator.py); - 构建
resnet_50(num_classes, data_format='channels_last')前向图,用saver.restore加载 checkpoint; - 取
logits[0][neuron_selector](预 softmax 激活)作为归因目标,neuron_selector由模型预测类别prediction_out[0]填充; - 按
saliency_method生成热力图,展平后与原始图像、标签一起写入新的 TFRecord(image_to_tfexample,data_helper.py)。
输出的每个样本包含三个字段:raw_image(原始图像)、以方法命名的显著性字段(如ig_smooth、gradient_smooth、gb_smooth等,映射关系见 data_helper.py 的saliency_dict)以及label。非测试模式下,输出目录结构为output_dir/dataset_name/resnet_50/saliency_method/(dataset_generator.py),且每个输入 shard 对应一个输出 shard。
第三步:在修改后的数据集上重训练 ResNet-50
train_resnet.py 在显著性 TFRecord 数据集上重训练 ResNet-50。注意:这里训练的模型与生成显著性图时使用的 checkpoint 是同一个网络结构(ResNet-50),但参数完全不同——ROAR 的核心就在于每次都在"被移除重要像素"的数据上重新训练,以此度量估计器定位关键特征的能力。
三种数据变换模式(--transformation)
transformation参数决定模型训练在哪种数据上(train_resnet.py):
raw_image:直接在未修改的原图上训练——作为上界对照;modified_image:按显著性估计器排序移除最/不重要像素后训练——待评测的"治疗组";random_baseline:随机移除像素后训练——排除"图像损坏"这一混淆变量的对照。
modified_image与random_baseline模式下,数据管线在 data_input.py 中实现像素替换:
compute_feature_ranking(data_input.py):先按use_squared_value决定是否对热力图取平方,再通过rescale_input将热力图缩放到[0,1](实现见 preprocessing_helper.py,加1e-5epsilon 防除零);随后调用percentage_ranking用tf.nn.top_k选出数值最高的threshold%像素;percentage_ranking(data_input.py):对显著性值做top_k后经tf.scatter_nd还原为掩码,keep_information=True时保留这些像素(其余替换为均值),False时移除这些像素;random_ranking(data_input.py):用tf.nn.dropout以threshold/100的保留概率随机选择像素,语义上等价于随机损坏,作为对照基线。
实现细节上,替换值使用每个通道的全局均值而非 0(global_mean_constant),并在top_k前给热力图加上0.00001的 epsilon,用于区分"热力图原本为 0 的像素"与"被tf.scatter_nd置 0 的像素"(data_input.py)。
ROAR 关键参数
| 参数 | 默认值 | 说明 |
|---|---|---|
--transformation | raw_image | raw_image/random_baseline/modified_image三种模式 |
--saliency_method | ig_smooth_2 | 用于估计重要像素的方法,枚举与data_helper.saliency_dict一致(如gradient_image、ig_smooth、gb_smooth_2、sobel等 13 个值) |
--keep_information | False | True时保留(preserve)最重要像素,False时移除(remove)最重要像素 |
--squared_value | True | 排名基于热力图平方值还是原始像素值 |
--threshold | 80. | 被修改的输入特征比例(百分比),即移除/保留排序前多少比例的像素 |
saliency_dict的映射关系(train_resnet.py)保证train_resnet.py中使用的字段名(如ig_smooth_2)与dataset_generator.py生成 TFRecord 时的字段名一一对应,例如ig_smooth↔IG_SG、gradient_smooth_2↔SH_SG_2、gb_image↔GB、sobel↔SOBEL。
各数据集训练配置
训练脚本内置了三套数据集超参数(train_resnet.py),供实验复现参考:
| 配置项 | ImageNet | Food101 | Birdsnap |
|---|---|---|---|
train_batch_size | 4096 | 256 | 256 |
num_train_images | 1,281,167 | 75,750 | 47,386 |
num_eval_images | 50,000 | 25,250 | 2,443 |
num_label_classes | 1000 | 101 | 500 |
num_train_steps | 32,000 | 20,000 | 20,000 |
base_learning_rate | 0.1 | 0.7 | 1.0 |
eval_batch_size | 1024 | 256 | 224 |
此外,模型统一采用分段的阶梯学习率调度lr_schedule = [(1.0, 5), (0.1, 30), (0.01, 60), (0.001, 80)],即每个"倍率、起始 epoch"元组,配合 Nesterov Momentum(momentum=0.9)优化器,并叠加 L2 权重衰减(weight_decay=1e-4,不含 batch normalization 参数)与 0.1 的标签平滑(train_resnet.py)。数据格式统一为channels_last。
训练 / 评估运行方式
# 训练(默认模式) python -m interpretability_benchmark.train_resnet \ --mode=train \ --dataset_name=birdsnap \ --transformation=modified_image \ --saliency_method=ig_smooth_2 \ --threshold=80. \ --base_dir=/path/to/saliency/tfrecords/ \ --output_dir=/tmp/roar_experiments/ # 评估:监听 checkpoint 并计算 top-1 / top-5 准确率 python -m interpretability_benchmark.train_resnet \ --mode=eval \ --dataset_name=birdsnap \ --transformation=modified_image \ --saliency_method=ig_smooth_2 \ --threshold=80. \ --base_dir=/path/to/saliency/tfrecords/ \ --output_dir=/tmp/roar_experiments/关键机制说明:
--base_dir指向显著性 TFRecord 所在目录,脚本会自动拼接数据路径为base_dir/dataset_name/2018-12-10/resnet_50/saliency_method/split*(train_resnet.py),其中split在train模式下为training、eval模式下为validation;model_dir按实验配置自动组织为output_dir/dataset_name/transformation/threshold/base_learning_rate/weight_decay/squared_or_not/keep_or_remove/[saliency_method],保证不同配置的实验结果互不覆盖(train_resnet.py);- 评估模式使用
tf2.training.checkpoints_iterator持续监听新 checkpoint,自动计算top_1_accuracy与top_5_accuracy(train_resnet.py); - 数据管线中,
train模式会先 shuffle、repeat,再做parallel_interleave(cycle_length=num_cores)多文件并行读取与 64 路并行解析、预取(data_input.py)。
环境搭建与快速自测
依赖列表见 requirements.txt:absl-py>=0.6.0、tensorflow>=1.11.0(可替换为tensorflow-gpu获得 GPU 支持)、numpy>=1.15.2、scipy>=1.0.0、scikit-image,另需 pip 安装saliency库。代码基于tensorflow.compat.v1编写,并在脚本入口处调用tf.disable_v2_behavior()(dataset_generator.py),适用于 TensorFlow 1.x 或兼容 v1 API 的环境。
仓库提供的 run.sh 演示了完整的环境初始化与自测流程:
set -e set -x virtualenv -p python3 env source env/bin/activate pip install saliency pip3 install -r interpretability_benchmark/requirements.txt output_dir="/tmp/" python -m interpretability_benchmark.train_resnet_test --dest_dir=output_dir即:创建 python3 虚拟环境 → 安装saliency与 requirements 依赖 → 运行 train_resnet_test.py(对应train_resnet.py的单元测试)验证训练管线。两个核心脚本也都内置了--test_small_sample=True的合成数据自测开关:生成阶段仅处理 2 张空图样本,训练阶段使用 batch_size=2、10 步的小规模配置(train_resnet.py),便于快速验证代码工作流正确性。
结果解读与扩展
完成三个模式(raw_image、modified_image、random_baseline)在多个threshold与多个saliency_method下的实验后,通过对比不同估计器在相同阈值下的 top-1/top-5 准确率即可得到结论:同一阈值下,移除某估计器认为最重要像素后模型精度越低,该估计器越准确。而random_baseline提供了随机移除的基准线,raw_image则给出无任何修改的精度上界,二者共同构成评估的参照系。论文与 README 表明,实验通常在多个移除比例(如 10%~90%)下重复,以刻画"准确率-移除比例"曲线的整体形态,而非单点比较。
代码库欢迎以 pull request 形式新增更多待评测的可解释性方法(例如在 saliency_helper.py 中接入新的 saliency 实现,并在 data_helper.py 与 train_resnet.py 两处saliency_dict中同步注册字段名),或改进现有代码。基准背后的方法论细节与完整实验可参阅论文《A Benchmark for Interpretability Methods in Deep Neural Networks》(NeurIPS 2019)。
引用
若在研究中使用了 ROAR 代码,可引用:
@incollection{NIPS2019_9167, title = {A Benchmark for Interpretability Methods in Deep Neural Networks}, author = {Hooker, Sara and Erhan, Dumitru and Kindermans, Pieter-Jan and Kim, Been}, booktitle = {Advances in Neural Information Processing Systems 32}, editor = {H. Wallach and H. Larochelle and A. Beygelzimer and F. d\textquotesingle Alch\'{e}-Buc and E. Fox and R. Garnett}, pages = {9737--9748}, year = {2019}, publisher = {Curran Associates, Inc.}, url = {http://papers.nips.cc/paper/9167-a-benchmark-for-interpretability-methods-in-deep-neural-networks.pdf} }小结
ROAR 提供了一套可复现、可对照的定量评测框架:先由 dataset_generator.py 基于预训练 ResNet-50 生成 10 种显著性估计的 TFRecord,再由 train_resnet.py 在按估计器排序修改(或随机修改、或不修改)的数据上重训练并度量 top-1/top-5 准确率。整个流程环环相扣:saliency_dict保证字段名在生成与训练两端一致,transformation+threshold+keep_information组合出完整的消融矩阵,内置的三套数据集超参数与test_small_sample自测开关让实验既能复现论文结论,也能快速验证新估计器的表现。
- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
Google Research
相关推荐
kotaemon质量评估:检索准确性度量标准
kotaemon质量评估:检索准确性度量标准 引言:为什么检索质量评估如此重要? 在RAG(Retrieval Augmented Generation,检索增
人工智能大模型RAG向量数据库后端Hyprnote性能基准测试:转录速度与准确率评估
Hyprnote性能基准测试:转录速度与准确率评估 引言 在现代会议场景中,实时语音转录已成为提升工作效率的关键技术。Hyprnote作为一款本地优先的AI记事
AI 应用人工智能语音本地部署桌面应用音频如何使用fastai Captum实现深度学习模型可解释性与特征重要性分析:完整指南
如何使用fastai Captum实现深度学习模型可解释性与特征重要性分析:完整指南 fastai是一个强大的深度学习库,它通过Captum集成提供了直观的模型
人工智能深度学习
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考