news 2026/9/21 1:40:57

ROAR 可解释性基准评测指南:用 RemOve And Retrain 度量深度神经网络特征重要性估计的准确度

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
ROAR 可解释性基准评测指南:用 RemOve And Retrain 度量深度神经网络特征重要性估计的准确度
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/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)等。

在医疗、自动驾驶、信用评分等敏感领域,可解释性估计必须同时满足两个要求:

  1. 对人类有意义(meaningful to a human);
  2. 高度准确(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 GradientsIntegratedGradients.GetMask
IG_SGIG + SmoothGrad(非平方)GetSmoothedMask(magnitude=False)
IG_SG_2IG + SmoothGrad²(平方)GetSmoothedMask(magnitude=True)
SH敏感度热力图 Gradient SaliencyGradientSaliency.GetMask
SH_SGSH + SmoothGradGetSmoothedMask(magnitude=False)
SH_SG_2SH + SmoothGrad²GetSmoothedMask(magnitude=True)
GB引导反向传播 Guided BackpropGuidedBackprop.GetMask
GB_SGGB + SmoothGradGetSmoothedMask(magnitude=False)
GB_SG_2GB + SmoothGrad²GetSmoothedMask(magnitude=True)
SOBELSobel 边缘算子(非梯度基线)scipy.ndimage.sobel

方法分发逻辑在 saliency_helper.py 的generate_saliency_image中实现:get_saliency_image负责在 TensorFlow 计算图上实例化三类 saliency 对象(saliency_helper.py),随后按方法名调用对应的GetMask/GetSmoothedMaskSOBEL比较特殊——它不依赖梯度,直接对预处理图像做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)展示了完整链路:

  1. DataIterator读取 TFRecord shard,解析出原始图像与标签;
  2. 按 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);
  3. 构建resnet_50(num_classes, data_format='channels_last')前向图,用saver.restore加载 checkpoint;
  4. logits[0][neuron_selector](预 softmax 激活)作为归因目标,neuron_selector由模型预测类别prediction_out[0]填充;
  5. saliency_method生成热力图,展平后与原始图像、标签一起写入新的 TFRecord(image_to_tfexample,data_helper.py)。

输出的每个样本包含三个字段:raw_image(原始图像)、以方法命名的显著性字段(如ig_smoothgradient_smoothgb_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_imagerandom_baseline模式下,数据管线在 data_input.py 中实现像素替换:

  • compute_feature_ranking(data_input.py):先按use_squared_value决定是否对热力图取平方,再通过rescale_input将热力图缩放到[0,1](实现见 preprocessing_helper.py,加1e-5epsilon 防除零);随后调用percentage_rankingtf.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.dropoutthreshold/100的保留概率随机选择像素,语义上等价于随机损坏,作为对照基线。

实现细节上,替换值使用每个通道的全局均值而非 0(global_mean_constant),并在top_k前给热力图加上0.00001的 epsilon,用于区分"热力图原本为 0 的像素"与"被tf.scatter_nd置 0 的像素"(data_input.py)。

ROAR 关键参数

参数默认值说明
--transformationraw_imageraw_image/random_baseline/modified_image三种模式
--saliency_methodig_smooth_2用于估计重要像素的方法,枚举与data_helper.saliency_dict一致(如gradient_imageig_smoothgb_smooth_2sobel等 13 个值)
--keep_informationFalseTrue时保留(preserve)最重要像素,False时移除(remove)最重要像素
--squared_valueTrue排名基于热力图平方值还是原始像素值
--threshold80.被修改的输入特征比例(百分比),即移除/保留排序前多少比例的像素

saliency_dict的映射关系(train_resnet.py)保证train_resnet.py中使用的字段名(如ig_smooth_2)与dataset_generator.py生成 TFRecord 时的字段名一一对应,例如ig_smoothIG_SGgradient_smooth_2SH_SG_2gb_imageGBsobelSOBEL

各数据集训练配置

训练脚本内置了三套数据集超参数(train_resnet.py),供实验复现参考:

配置项ImageNetFood101Birdsnap
train_batch_size4096256256
num_train_images1,281,16775,75047,386
num_eval_images50,00025,2502,443
num_label_classes1000101500
num_train_steps32,00020,00020,000
base_learning_rate0.10.71.0
eval_batch_size1024256224

此外,模型统一采用分段的阶梯学习率调度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),其中splittrain模式下为trainingeval模式下为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_accuracytop_5_accuracy(train_resnet.py);
  • 数据管线中,train模式会先 shuffle、repeat,再做parallel_interleave(cycle_length=num_cores)多文件并行读取与 64 路并行解析、预取(data_input.py)。

环境搭建与快速自测

依赖列表见 requirements.txt:absl-py>=0.6.0tensorflow>=1.11.0(可替换为tensorflow-gpu获得 GPU 支持)、numpy>=1.15.2scipy>=1.0.0scikit-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_imagemodified_imagerandom_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

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

相关推荐

上一篇:Awesome-Dify-Workflow:Ollama模型集成方案
下一篇:如何用ansible-docker实现Docker镜像仓库登录与凭证管理?实用技巧分享

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/21 1:39:58

从技术要点到完整博客:素材驱动的写作方法论

简介:这是一份系统讲解OpenCV多传感器融合与位姿估计优化的技术文档,共483页,面向机器人、自动驾驶与视觉SLAM方向的中高级开发者,旨在解决时间同步、状态估计和传感器标定等工程落地难题。资源为单个PDF文件,大小12.7…

作者头像 李华
网站建设 2026/9/21 1:39:38

BrowserSkill截图指南:视口、元素、整页3种模式与参数速查表

BrowserSkill截图指南:视口、元素、整页3种模式与参数速查表 【免费下载链接】BrowserSkill Let AI agents use your real, logged-in browser without interrupting your work. CLI extension for browser automation across any shell-capable AI agent. 项目地…

作者头像 李华
网站建设 2026/9/21 1:38:16

Sails 框架 res.json() 完全指南:用法、源码实现与最佳实践

Sails 框架 res.json() 完全指南:用法、源码实现与最佳实践 【免费下载链接】sails Realtime MVC Framework for Node.js 项目地址: https://gitcode.com/gh_mirrors/sa/sails 导读 res.json() 是 Sails(基于 Node.js 的 Realtime MVC 框架&…

作者头像 李华
网站建设 2026/9/21 1:35:00

零基础选AI软件:先搞懂三件事,按场景匹配不踩坑

1. 先别急着下载:零基础选 AI 软件,先搞懂三件事这段时间我收到特别多类似的提问:大家都是零基础,看到网上铺天盖地的 AI 软件推荐,脑子里全是问号。某某 AI 能做 PPT,某某 AI 能写代码,某某 AI…

作者头像 李华