12X速度提升:如何用Quantus批处理指标让Faithfulness指标计算快12倍
【免费下载链接】Quantus[JMLR 2023] Quantus is an eXplainable AI toolkit for responsible evaluation of neural network explanations项目地址: https://gitcode.com/gh_mirrors/qu/Quantus
Quantus 是一个可解释AI(XAI)责任评估工具包(JMLR 2023 论文配套开源项目),本文以Faithfulness 指标(faithfulness 评估)为例,讲解如何使用 Quantus 的批处理指标(batched metrics)实现让 Faithfulness 指标计算提速 12 倍,并介绍batch_size关键参数与quantus.evaluate()的大规模评估工作流。
为什么 Faithfulness 指标计算这么慢?
Faithfulness 指标回答的是"解释与模型行为有多一致":它会迭代地扰动输入(把解释认为重要的特征替换/掩码掉),再观察模型预测如何变化。以 Monotonicity 为例,它从基线出发逐步加回重要性最高的特征,每一步都需要一次完整的前向预测。
旧实现采用逐样本循环:对 batch 中每个样本单独调用模型预测,GPU 经常吃不饱,Python 循环本身成为瓶颈——样本越多,等待越久。
批处理指标:从"逐个算"到"向量化批量算"的 12X 提速
Quantus 官方更新说明明确写道:New batch implementation for 12X speedup of existing faithfulness metrics (!)——现有 Faithfulness 指标的计算速度提升 12 倍。🚀
核心改动有三个:
- 批处理扰动函数:如 quantus/functions/perturb_func.py 中的
batch_baseline_replacement_by_indices,一次完成整个 batch 所有样本的特征替换,替代原来的逐样本baseline_replacement_by_indices; - 批量模型推理:扰动后的整批输入一次性送入模型,GPU 矩阵运算被充分利用;
- 统一的
evaluate_batch接口:每个指标只需实现"对一批数据做评估",切分、预处理、聚合全部由基类托管。
快速上手:3 步跑通批处理评估
先安装(按需选择框架):
pip install "quantus[torch]"import quantus # 1) 实例化指标 metric = quantus.Monotonicity(features_in_step=1, display_progressbar=True) # 2) 直接传入整批数据,batch_size 控制内部切分粒度(默认 64) scores = metric( model=model, x_batch=x_batch, y_batch=y_batch, a_batch=a_batch_saliency, batch_size=64, )# 3) 大规模评估:多个指标 × 多个解释方法 results = quantus.evaluate( metrics={"monotonicity": quantus.Monotonicity()}, xai_methods={"Saliency": a_batch_saliency}, model=model, x_batch=x_batch, y_batch=y_batch, )关键参数与性能调优技巧 🛠️
batch_size(默认 64):控制指标内部切分粒度。批越大 GPU 利用率越高,但显存占用也越大;显存不足时调小即可。- 懒生成解释:不传
a_batch时,Quantus 按 batch 逐批调用explain_func生成解释,避免一次性生成整批解释导致 OOM(见batch_preprocess逻辑)。 return_aggregate/aggregate_func:把逐样本分数聚合成单值(默认np.mean),方便横向对比不同解释方法。display_progressbar:打开后批处理评估会显示 tqdm 进度条,方便观察长任务。
批处理指标是如何工作的(源码走读)
主入口是 quantus/metrics/base.py 中的Metric.__call__,流程为:
general_preprocess():统一通道布局、包装模型、对解释做归一化/取绝对值;generate_batches():按batch_size把数据切分成小批并逐批产出;batch_preprocess():必要时懒生成当前批的解释;evaluate_batch():每个具体指标实现此方法完成"整批扰动 + 整批预测",例如 quantus/metrics/faithfulness/monotonicity.py 中先对整个 batch 排序归因索引,再逐步替换并批量预测;- 收集
evaluation_scores,按需聚合后返回。
哪些 Faithfulness 指标享受提速?
Faithfulness 类别下 12 个指标全部支持批处理接口,位于quantus/metrics/faithfulness/目录:
| 指标 | 源文件 |
|---|---|
| Monotonicity(单调性) | monotonicity.py |
| Pixel Flipping(像素翻转) | pixel_flipping.py |
| Region Perturbation(区域扰动) | region_perturbation.py |
| Sensitivity-N | sensitivity_n.py |
| IROF | irof.py |
| ROAD | road.py |
| Infidelity | infidelity.py |
| Sufficiency | sufficiency.py |
| Selectivity | selectivity.py |
| Faithfulness Correlation / Estimate | faithfulness_correlation.py / faithfulness_estimate.py |
| Monotonicity Correlation | monotonicity_correlation.py |
总结
- 12X 提速= 批处理扰动函数 + 批量模型推理 + 统一的
evaluate_batch接口,三者缺一不可; - 用户开箱即用:传入整批数据即可,用
batch_size调节速度与显存的平衡; - 搭配
quantus.evaluate()做多指标、多解释方法的大规模基准测试,Faithfulness 评估从此不再是等待的艺术。⚡
【免费下载链接】Quantus[JMLR 2023] Quantus is an eXplainable AI toolkit for responsible evaluation of neural network explanations项目地址: https://gitcode.com/gh_mirrors/qu/Quantus
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考