[论文解读] Optimizing for Interpretability in Deep Neural Networks with Tree Regularization
本文通过引入树正则化,训练出既高度准确又可由人类模拟的深度神经网络,方法是促使模型的决策函数能被紧凑的、轴对齐的决策树良好近似。该方法采用L0稀疏性(通过sparsemax实现)的全局与区域树正则化,以在保持与原始模型高保真度的同时,使领域专家能够通过易于模拟的决策规则来解释预测结果。
Deep models have advanced prediction in many domains, but their lack of interpretability remains a key barrier to the adoption in many real world applications. There exists a large body of work aiming to help humans understand these black box functions to varying levels of granularity -- for example, through distillation, gradients, or adversarial examples. These methods however, all tackle interpretability as a separate process after training. In this work, we take a different approach and explicitly regularize deep models so that they are well-approximated by processes that humans can step-through in little time. Specifically, we train several families of deep neural networks to resemble compact, axis-aligned decision trees without significant compromises in accuracy. The resulting axis-aligned decision functions uniquely make tree regularized models easy for humans to interpret. Moreover, for situations in which a single, global tree is a poor estimator, we introduce a regional tree regularizer that encourages the deep model to resemble a compact, axis-aligned decision tree in predefined, human-interpretable contexts. Using intuitive toy examples as well as medical tasks for patients in critical care and with HIV, we demonstrate that this new family of tree regularizers yield models that are easier for humans to simulate than simpler L1 or L2 penalties without sacrificing predictive power.
研究动机与目标
- 为解决深度学习中可解释性的关键障碍,使模型具备人类可模拟性——即能够由领域专家手动逐步推导。
- 克服事后可解释性方法的局限,后者无法捕捉模型的全局逻辑或需要复杂推理。
- 开发一种训练时正则化方法,明确优化可模拟性,而非在模型训练后才应用可解释性。
- 使领域专家能够通过提供可解释的、树状结构的决策函数,对模型决策进行审计、验证和改进。
提出的方法
- 引入一种全局树正则化项,促使深度模型的决策函数能被单一、紧凑、轴对齐的决策树良好近似。
- 提出一种区域树正则化框架,将训练数据划分为R个可由人类理解的区域,每个区域拥有其自身的局部决策树。
- 使用sparsemax(L0范数的可微分近似)来强制区域选择的稀疏性,防止对简单决策边界的过度正则化。
- 采用代理模型来估计蒸馏树的平均路径长度(APL),作为可解释性的代理指标,并通过深度模型反向传播该指标。
- 端到端训练深度模型,损失函数由标准预测损失与基于APL和蒸馏树保真度的正则化项组合而成。
- 支持可定制的代理模型训练频率和区域优先级,将区域选择视为多臂赌博机问题,以降低计算成本。
实验结果
研究问题
- RQ1是否可以通过在训练过程中显式正则化,使深度神经网络在不牺牲性能的前提下,既高度准确又具备人类可模拟性?
- RQ2在深度模型上强制施加全局或区域树结构,是否能在不损害预测性能的前提下提升可解释性?
- RQ3基于L0的稀疏性(通过sparsemax实现)与L1或L2正则化相比,在保持有意义且非平凡的决策边界方面表现如何?
- RQ4区域树正则化是否能让领域专家在特定上下文、临床相关的子人群中理解模型行为?
主要发现
- 区域树正则化模型在深度模型与蒸馏树之间实现了89%的保真度,表明在大多数样本中决策逻辑高度一致。
- 基于L0的区域正则化(sparsemax)在实现低APL(平均路径长度)和高AUC极小值方面优于L1、L2和softmax近似,避免了平凡的决策函数。
- 在Sepsis数据集上,区域树正则化使每轮训练时间增加约39.9秒(相比L2的约2.4秒),但该开销可控且可通过重用代理模型实现可扩展性。
- 重症监护和HIV领域的医生能够快速理解、验证并提出对蒸馏决策树的改进建议,证明了其实际可解释性。
- 该方法在AUC上优于标准决策树,同时保持低APL,表明经过树正则化的深度模型既准确又可模拟。
- 发现无梯度优化方法(如Nelder-Mead和输入扰动)不稳定或计算成本过高,而基于代理的优化方法则更稳定高效。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。