Skip to main content
QUICK REVIEW

[论文解读] Discovering and Explaining the Representation Bottleneck of DNNs

Huiqi Deng, Qihan Ren|arXiv (Cornell University)|Nov 11, 2021
Explainable Artificial Intelligence (XAI)参考文献 36被引用 16
一句话总结

本文识别出深度神经网络(DNNs)中的一种表征瓶颈,即模型在学习输入变量之间的低阶与高阶交互时表现突出,但在处理中等复杂度交互时表现困难。作者通过多阶交互效用工具,从理论上和实证上证明了这种偏差,并将其与DNN与人类感知之间的认知差距联系起来,同时提出损失函数以调节交互复杂度的学习,揭示了不同交互阶次下表征能力的显著差异。

ABSTRACT

This paper explores the bottleneck of feature representations of deep neural networks (DNNs), from the perspective of the complexity of interactions between input variables encoded in DNNs. To this end, we focus on the multi-order interaction between input variables, where the order represents the complexity of interactions. We discover that a DNN is more likely to encode both too simple interactions and too complex interactions, but usually fails to learn interactions of intermediate complexity. Such a phenomenon is widely shared by different DNNs for different tasks. This phenomenon indicates a cognition gap between DNNs and human beings, and we call it a representation bottleneck. We theoretically prove the underlying reason for the representation bottleneck. Furthermore, we propose a loss to encourage/penalize the learning of interactions of specific complexities, and analyze the representation capacities of interactions of different complexities.

研究动机与目标

  • 调查DNN在表征特定类型特征交互时是否存在系统性偏差。
  • 确定DNN在图像分类任务中是否以与人类认知相似的方式编码视觉概念。
  • 识别并解释与交互复杂度相关的DNN表征能力的根本局限性。
  • 开发可鼓励或惩罚特定复杂度阶次交互学习的损失函数。

提出的方法

  • 作者使用多阶交互效用 $ I^{(m)}(i,j) $,即在所有大小为 $ m $ 的上下文下,变量 $ i $ 与 $ j $ 之间交互效用的平均值,来量化交互复杂度。
  • 他们将DNN输出分解为局部效用与多阶交互效用之和:$ \text{output} = \sum_{m=0}^{n-2} \sum_{i,j} w^{(m)} I^{(m)}(i,j) + \sum_i \text{local utility} + \text{bias} $。
  • 理论分析证明,第 $ m $ 阶交互的学习强度与 $ \frac{n-m-1}{n(n-1)} / \sqrt{\binom{n-2}{m}} $ 成正比,从而解释了观察到的瓶颈现象。
  • 他们提出一种损失函数,通过基于 $ I^{(m)}(i,j) $ 修改梯度更新,以惩罚或鼓励特定交互阶次的学习。
  • 在ResNet-18上使用Tiny-ImageNet进行实证验证,测量不同上下文大小和训练阶段下的交互效用。
  • 通过将梯度投影到主方向并确认不同上下文大小下均值接近零,验证了理论推导中的零均值假设。

实验结果

研究问题

  • RQ1DNN是否在学习中等复杂度交互时表现出系统性偏差?
  • RQ2交互复杂度(以阶次 $ m $ 衡量)如何影响DNN的表征能力?
  • RQ3DNN学习到的交互在多大程度上与人类视觉认知一致?
  • RQ4能否通过针对性的损失函数控制交互复杂度?
  • RQ5DNN中观察到的表征瓶颈的理论基础是什么?

主要发现

  • 存在表征瓶颈:DNN在低阶和高阶交互效用上表现出较高的绝对值,但中阶交互的值显著偏低。
  • 理论分析证实,第 $ m $ 阶交互的学习强度在交互阶次谱的中段急剧下降,解释了该瓶颈现象。
  • 在ResNet-18上的实证结果表明,即使在训练40个周期后,中阶交互在特征表征中仍持续被低估。
  • 理论模型中的零均值假设得到验证:梯度在主方向上的投影在不同上下文大小下均值接近零,支持了模型的可靠性。
  • 尽管交互学习存在差异,但具有不同交互复杂度偏好性的DNN在分类准确率上表现相似,表明瓶颈并未损害模型性能。
  • 所提出的损失函数成功调节了特定交互阶次的学习,证实了该方法的可控性与实际应用价值。

更好的研究,从现在开始

从阅读论文到最终审阅,大幅缩短您的研究时间。

无需绑定信用卡

本解读由 AI 生成,并经人工编辑审核。