技术摘要:快速近似估计线性解释器中的条件 Shapley 值
问题陈述
本文解决了在协变量具有依赖性时,为线性回归模型估计**条件 Shapley 值(conditional Shapley values)**所面临的计算瓶颈。虽然 Shapley 值提供了对特定预测的特征贡献的博弈论分解,但计算条件值需要估计所有可能的特征子集(联盟)的模型期望输出。
对于一个拥有 p 个预测变量的模型,存在 2p 个可能的联盟。现有的方法,例如 shapr R 包中的**顺序(Sequential)和迭代(Iterative)**方法,面临着显著的挑战:
- 顺序估计: 需要拟合 2p 个独立的回归模型,随着 p 的增加,这在计算上变得难以承受。
- 迭代估计: 通过采样联盟直到满足收敛标准。虽然在少量联盟即可收敛时速度较快,但在困难情况下,它仍可能需要采样大量的 2p 联盟比例才能达到收敛,导致运行时间(数小时)与顺序方法相当。
- 特征依赖性: 与边际 Shapley 值不同,条件值必须考虑协变量之间的依赖结构,因此需要复杂的条件期望计算。
作者旨在开发能够同时且快速地估计所有 2p 个子模型系数的方法,并利用其底层数学结构的稀疏性。
方法论
所提出的方法利用**约束高斯马尔可夫随机场(GMRF)**理论和稀疏矩阵代数,来联合而非顺序地估计所有子模型的回归系数。作者推导出了三种不同的算法:
1. 对角修正法(近似)
该方法通过修改全模型的精度矩阵(Q=XTX)来近似约束回归。
- 机制: 为了强制执行一个约束,即某个系数为零(即该特征被排除在子模型之外),需要在精度矩阵对应的对角元素上加上一个大的标量值 κ。
- 理论基础: 基于 Woodbury 矩阵恒等式,当 κ→∞ 时,添加 κATA(其中 A 是约束矩阵)可以近似条件协方差矩阵。
- 实现: 构建一个用于所有 2p 个模型的联合约束矩阵 A。该方法求解线性系统 (QF+κATA)−1m,其中 QF 是包含重复全模型精度矩阵的块对角矩阵。
- 调优: 需要一个调优参数 κ,建议设定为原始精度矩阵最大特征值的 105 倍。
2. 投影法(近似)
该方法通过投影参数空间来移除受约束的参数。
- 机制: 应用投影矩阵 Z=I−ATA 到精度矩阵。受约束的模型被视为一个内在的 GMRF。为了确保数值稳定性(正定性),在投影后的精度矩阵对角线上添加一个很小的标量 ϵ(即 ZQFZ+ϵI)。
- 理论基础: 当 ϵ→0+ 时,解收敛于精确的约束估计。
- 实现: 求解系统 Z(ZQFZ+ϵI)−1Zm。
- 调优: 需要一个很小的调优参数 ϵ(例如 10−5),该参数应相对于 Q 的最小特征值较小,但足以保证 Cholesky 分解的稳定性。
3. 精确变换法(精确)
该方法提供了一个没有近似误差或调优参数的精确解。
- 机制: 该方法不直接在全参数空间中工作并强制系数为零,而是将问题转化为一个仅包含非零参数的低维空间。
- 实现: 构建一个映射矩阵 E 来仅选择每个子模型中活跃的参数。系统求解方式为 (EQFET)−1Em。
- 优势: 消除了对大型惩罚项(κ)或小型正则化项(ϵ)的需求,从而避免了潜在的数值不稳定或近似偏差。
所有三种方法都利用稀疏矩阵代数(特别是 Cholesky 分解)来高效处理数以千计模型的联合估计。由于这些模型是相互独立的(具有块对角结构),且离角元素为零表示条件独立,因此联合精度矩阵是高度稀疏的。
主要贡献
- 三种新算法: 本文引入了两种近似方法(对角修正、投影)和一种精确方法(精确变换),用于线性回归中所有子模型的联合估计。
- 收敛证明: 作者提供了数学证明,证明了近似方法会随着各自的调优参数(κ→∞ 和 ϵ→0+)趋向于真实的条件估计。
- 计算效率: 通过利用稀疏性和联合估计,这些方法将计算时间从(典型
shapr 的)数小时缩短至数秒或数分钟,即使在枚举所有 2p 个联盟的情况下也是如此。
- 实证验证: 通过三个数据集(Adult Income、模拟高斯数据集和 WHO 预期寿命数据集)将新方法与
shapr 包(包括 Sequential 和 Iterative 版本)进行了广泛的数值研究对比。
结果
数值案例研究得出以下发现:
- 准确性: 新方法产生的 Shapley 值估计量与
shapr 中的 Sequential 方法(用作全枚举的基准真相)几乎完全一致。
- 在 Adult 数据集(p=11)中,新方法与 Sequential 方法之间的中值绝对误差(MedAE)在 10−9 数量级。
- 在 模拟数据集(p=21)中,MedAE 在 10−7 到 10−5 数量级。
- 在 WHO 数据集(p=16)中,MedAE 在 10−6 数量级。
- 速度:
- Adult 数据集: 新方法运行时间为 2.5 至 9 秒,而
shapr(Sequential 和 Iterative)需要 17 至 19 分钟。
- 模拟数据集: 新方法在估计所有 221≈210 万个模型时耗时 1.4 至 2.5 分钟。
shapr 的 Iterative 方法更快(6.48 秒),但它仅估计了 200 个联盟,而非全集。
- WHO 数据集: 新方法耗时 1.1 至 1.2 分钟。
shapr 的 Sequential 方法耗时 131 分钟(2.2 小时),Iterative 方法(并行)耗时 16 分钟(或串行耗时 36 分钟)。
- 可扩展性: 这些方法成功地在几分钟内估计了超过两百万个子模型。识别出的主要限制是内存约束(维度灾难),当 p 在标准硬件上超过约 21 时会出现此问题,但作者指出可以通过增加 RAM 或使用专门的稀疏矩阵包来缓解。
意义与主张
论文声称,所提出的方法在不牺牲准确性的情况下,显著提高了条件 Shapley 值估计的计算效率。
- 全枚举中的优越性: 当
shapr 的 Iterative 方法需要大量联盟才能收敛时(如在 Adult 和 WHO 数据集中所示),新方法在保持“相等或更好准确度”的同时,“显著更快”。
- 全面枚举: 不同于通过采样联盟的迭代方法,新方法提供了所有 2p 个子模型的全枚举,确保不会因为采样方差而丢失信息。
- 建议: 作者推荐在实际中使用精确变换法(Exact Transformation Method)。虽然三种方法的结果都非常接近,但精确法不需要调优参数,从而消除了超参数选择和潜在近似误差的来源。
- 未来潜力: 作者建议该方法论可以扩展到处理样条函数(例如来自
rms 包),并且并行化(例如通过 OpenMP)可以进一步缩短运行时间,尽管目前的实现依赖于串行执行,以便与并行化的 shapr 包进行公平比较。
论文总结道,这些方法使得线性模型中混合协变量类型的条件 Shapley 值的精确计算在以往被认为计算成本过高的场景下变得可行。