ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

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

ROAR 可解释性基准评测指南:用 RemOve And Retrain 度量深度神经网络特征重要性估计的准确度 人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载本文是 google-research 仓库中 interpretability_benchmark 目录的完整技术指南。该目录实现了ROARRemOve 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在三个公开图像分类数据集上评估了模型解释的准确性ImageNet1000 类Birdsnap500 类Food101101 类要复现实验结果首先下载目标数据集并将其转换为 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.GetMaskIG_SGIG SmoothGrad非平方GetSmoothedMask(magnitudeFalse)IG_SG_2IG SmoothGrad²平方GetSmoothedMask(magnitudeTrue)SH敏感度热力图 Gradient SaliencyGradientSaliency.GetMaskSH_SGSH SmoothGradGetSmoothedMask(magnitudeFalse)SH_SG_2SH SmoothGrad²GetSmoothedMask(magnitudeTrue)GB引导反向传播 Guided BackpropGuidedBackprop.GetMaskGB_SGGB SmoothGradGetSmoothedMask(magnitudeFalse)GB_SG_2GB SmoothGrad²GetSmoothedMask(magnitudeTrue)SOBELSobel 边缘算子非梯度基线scipy.ndimage.sobel方法分发逻辑在 saliency_helper.py 的generate_saliency_image中实现get_saliency_image负责在 TensorFlow 计算图上实例化三类 saliency 对象saliency_helper.py随后按方法名调用对应的GetMask/GetSmoothedMask。SOBEL比较特殊——它不依赖梯度直接对预处理图像做ndimage.sobel(img_out, axis0)作为无监督的边缘基线参与对比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_nameimagenet \ --saliency_methodIG_SG \ --splitvalidation \ --test_small_sampleFalse各参数说明定义于 dataset_generator.py--masterTensorFlow master 名称默认空字符串--output_dirTFRecord 输出目录默认/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_mapdataset_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_formatchannels_last)前向图用saver.restore加载 checkpoint取logits[0][neuron_selector]预 softmax 激活作为归因目标neuron_selector由模型预测类别prediction_out[0]填充按saliency_method生成热力图展平后与原始图像、标签一起写入新的 TFRecordimage_to_tfexampledata_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-50train_resnet.py 在显著性 TFRecord 数据集上重训练 ResNet-50。注意这里训练的模型与生成显著性图时使用的 checkpoint 是同一个网络结构ResNet-50但参数完全不同——ROAR 的核心就在于每次都在被移除重要像素的数据上重新训练以此度量估计器定位关键特征的能力。三种数据变换模式--transformationtransformation参数决定模型训练在哪种数据上train_resnet.pyraw_image直接在未修改的原图上训练——作为上界对照modified_image按显著性估计器排序移除最/不重要像素后训练——待评测的治疗组random_baseline随机移除像素后训练——排除图像损坏这一混淆变量的对照。modified_image与random_baseline模式下数据管线在 data_input.py 中实现像素替换compute_feature_rankingdata_input.py先按use_squared_value决定是否对热力图取平方再通过rescale_input将热力图缩放到[0,1]实现见 preprocessing_helper.py加1e-5epsilon 防除零随后调用percentage_ranking用tf.nn.top_k选出数值最高的threshold%像素percentage_rankingdata_input.py对显著性值做top_k后经tf.scatter_nd还原为掩码keep_informationTrue时保留这些像素其余替换为均值False时移除这些像素random_rankingdata_input.py用tf.nn.dropout以threshold/100的保留概率随机选择像素语义上等价于随机损坏作为对照基线。实现细节上替换值使用每个通道的全局均值而非 0global_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_image、ig_smooth、gb_smooth_2、sobel等 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_smooth↔IG_SG、gradient_smooth_2↔SH_SG_2、gb_image↔GB、sobel↔SOBEL。各数据集训练配置训练脚本内置了三套数据集超参数train_resnet.py供实验复现参考配置项ImageNetFood101Birdsnaptrain_batch_size4096256256num_train_images1,281,16775,75047,386num_eval_images50,00025,2502,443num_label_classes1000101500num_train_steps32,00020,00020,000base_learning_rate0.10.71.0eval_batch_size1024256224此外模型统一采用分段的阶梯学习率调度lr_schedule [(1.0, 5), (0.1, 30), (0.01, 60), (0.001, 80)]即每个倍率、起始 epoch元组配合 Nesterov Momentummomentum0.9优化器并叠加 L2 权重衰减weight_decay1e-4不含 batch normalization 参数与 0.1 的标签平滑train_resnet.py。数据格式统一为channels_last。训练 / 评估运行方式# 训练默认模式 python -m interpretability_benchmark.train_resnet \ --modetrain \ --dataset_namebirdsnap \ --transformationmodified_image \ --saliency_methodig_smooth_2 \ --threshold80. \ --base_dir/path/to/saliency/tfrecords/ \ --output_dir/tmp/roar_experiments/ # 评估监听 checkpoint 并计算 top-1 / top-5 准确率 python -m interpretability_benchmark.train_resnet \ --modeeval \ --dataset_namebirdsnap \ --transformationmodified_image \ --saliency_methodig_smooth_2 \ --threshold80. \ --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模式下为validationmodel_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_accuracytrain_resnet.py数据管线中train模式会先 shuffle、repeat再做parallel_interleave(cycle_lengthnum_cores)多文件并行读取与 64 路并行解析、预取data_input.py。环境搭建与快速自测依赖列表见 requirements.txtabsl-py0.6.0、tensorflow1.11.0可替换为tensorflow-gpu获得 GPU 支持、numpy1.15.2、scipy1.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_diroutput_dir即创建 python3 虚拟环境 → 安装saliency与 requirements 依赖 → 运行 train_resnet_test.py对应train_resnet.py的单元测试验证训练管线。两个核心脚本也都内置了--test_small_sampleTrue的合成数据自测开关生成阶段仅处理 2 张空图样本训练阶段使用 batch_size2、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保证字段名在生成与训练两端一致transformationthresholdkeep_information组合出完整的消融矩阵内置的三套数据集超参数与test_small_sample自测开关让实验既能复现论文结论也能快速验证新估计器的表现。赞分享人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载相关推荐kotaemon质量评估检索准确性度量标准kotaemon质量评估检索准确性度量标准 引言为什么检索质量评估如此重要 在RAGRetrieval Augmented Generation检索增人工智能大模型RAG向量数据库后端Hyprnote性能基准测试转录速度与准确率评估Hyprnote性能基准测试转录速度与准确率评估 引言 在现代会议场景中实时语音转录已成为提升工作效率的关键技术。Hyprnote作为一款本地优先的AI记事AI 应用人工智能语音本地部署桌面应用音频如何使用fastai Captum实现深度学习模型可解释性与特征重要性分析完整指南如何使用fastai Captum实现深度学习模型可解释性与特征重要性分析完整指南 fastai是一个强大的深度学习库它通过Captum集成提供了直观的模型人工智能深度学习上一篇Awesome-Dify-WorkflowOllama模型集成方案下一篇如何用ansible-docker实现Docker镜像仓库登录与凭证管理实用技巧分享创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进