布尔任务中ReLU-MLP的可认证可解释训练,保证真值表泛化
Certifiably Interpretable Training of ReLU-MLPs for Boolean Tasks with Guaranteed Truth-Table Generalization
- University of Toronto(多伦多大学)
- McMaster University(麦克马斯特大学)
- Vector Institute(向量研究所)
机构由 AI 辅助整理,请以论文原文为准。
AI总结:
提出MACCHIATO算法,通过迭代残差投影和电路编译,联合构建显式结构ReLU-MLP及布尔电路,实现可认证可解释性,并保证真值表泛化误差界。
AI中文摘要:
随着计算规模扩大、模型演进和训练算法进步,我们解释其所支撑的日益强大AI系统的能力正在削弱。为帮助保障可解释性,我们引入了一种专门的训练算法(MACCHIATO),该算法从部分真值表观测中联合构建(i)显式结构的$\operatorname{ReLU}$-MLP,以及(ii)一个基于带符号文字、具有$\{\operatorname{AND},\operatorname{OR},\operatorname{XOR}\}$门控的显式布尔电路,证明其子网络计算内容及其组合方式。直观上,我们迭代地将布尔函数的残差投影到低维$\{\operatorname{AND},\operatorname{OR},\operatorname{XOR}\}$电路类上,并将所得电路精确编译为$\operatorname{ReLU}$-MLP;我们结合了$\operatorname{ReLU}$-MLP电路编译、ESPRESSO逻辑最小化和基于影响力的变量选择。粗略地说,我们的可解释性证书辅以统计保证:在定理的影响力恢复条件下,如果$m$个逐阶段残差中的每一个至多依赖于$\log_2(B)$个比特,则我们算法的样本分裂变体在$T$个观测上训练后,返回一个宽度为$\mathcal{O}(mB)$的六层$\operatorname{ReLU}$-MLP(计入输入层),其真值表误差为$\mathcal{O}\bigl(\sqrt{m(B+\log(m/\delta))/T}\bigr)$。在合成随机junta任务上,我们的网络在若干数据稀疏或投影对齐机制中优于深度和隐藏宽度匹配的Adam训练MLP,而训练后的ReLU-MLP在其他机制中更强。此外,在我们显式的PyEDA真值表实现中,迭代过程在平坦环境维度ESPRESSO超过三小时计算预算的机制中完成。
英文摘要:
As compute scales, models evolve, and training algorithms advance, our ability to explain the increasingly powerful AI systems they enable is eroding. To help safeguard interpretability, we introduce a specialized training algorithm (MACCHIATO) that jointly constructs (i) an explicitly structured $\operatorname{ReLU}$-MLP from partial truth-table observations and (ii) an explicit Boolean circuit over signed literals with $\{\operatorname{AND},\operatorname{OR},\operatorname{XOR}\}$ gates certifying what its subnetworks compute and how they compose. Intuitively, we iteratively project the residuals of a Boolean function onto low-dimensional $\{\operatorname{AND},\operatorname{OR},\operatorname{XOR}\}$-circuit classes and exactly compile the resulting circuit into a $\operatorname{ReLU}$-MLP; we combine $\operatorname{ReLU}$-MLP circuit compilation, ESPRESSO logic minimization, and influence-based variable selection. Roughly speaking, our interpretability certificate is complemented by a statistical guarantee: under the theorem's influence-recovery conditions, if each of the $m$ stage-wise residuals depends on at most $\log_2(B)$ bits, a sample-splitting variant of our algorithm trained on $T$ observations returns a six-layer $\operatorname{ReLU}$-MLP (counting the input layer) of width $\mathcal{O}(mB)$ with truth-table error $\mathcal{O}\bigl(\sqrt{m(B+\log(m/δ))/T}\bigr)$. On synthetic random-junta tasks, our networks outperform depth- and hidden-width-matched Adam-trained MLPs in several data-sparse or projection-aligned regimes, while the trained ReLU-MLPs are stronger in others. Moreover, in our explicit PyEDA truth-table implementation, the iterative procedure completes in regimes where flat ambient-dimensional ESPRESSO exceeds the three-hour computational budget.