跳到主要内容
返回时间线
arXiv来源发表:

Polytopal Neural Network 将隐表示约束在学习到的多面体上,在 ImageNet-1k 上保留 96.9% 的无约束精度

核心概要

作者提出 Polytopal Neural Network(PNN),在训练中把每层隐表示约束为学习到的原型(archetype)的凸组合,使解释坐标同时就是下一层接收的坐标;在图像分类与重建基准上,PNN 以很小的精度损失逼近无约束网络,在无监督重建中优于 VQ-VAE 与 Dirichlet VAE,并给出无需 EMA、承诺损失或直通估计器的 VQ 训练路径。

Source-provided article image: The Polytopal Neural Network
Figure 1 ·

Figure 1: Two PathMNIST test images traced through two PNN layers ( K = 10 K=10 ). Each layer passes a convex combination of archetypes to the next layer. Numbers denote mixture weights; each archetype is illustrated by its two nearest training images.

arXiv

深度剖析

PNN 在选定的层后插入单纯形约束投影,使每个表示成为学习到的原型的凸组合,且该约束在训练中直接参与前向计算,而非事后解释。 此前的深度原型分析多只在单个瓶颈处放置原型,或仅以惩罚项施加原型结构;PNN 在多个层精确施加 AA 约束,并让解释坐标与下一层输入坐标一致。 论文给出 Lemma 1 与 Theorem 1,说明在正定条件下投影是连续分段仿射映射并可由有限 ReLU 网络精确表示;实验覆盖 MNIST、FashionMNIST、CIFAR-10、SVHN、EuroSAT、MedMNIST 与 ImageNet。

可扩展训练由两部分组成:与网络联合学习、按小批量流式更新的紧凑语料,以及摊销单纯形推断。 语料把原型与全数据集解耦,梯度内存随批量与语料规模而非样本数增长;摊销器只拟合投影残差,不接收任务损失,因此不会改变编码器、头部、语料或原型。 小批量训练与全数据更新取得相当精度;摊销器替换为精确 SMO 投影后,150 组配置的测试 MSE 落在对角线上。

把单纯形替换为 one-hot 坐标后,PNN 退化为一种 VQ 模型,其码本由语料定义,训练不需要 EMA 更新、承诺损失或直通估计器。 标准 VQ-VAE 依赖硬最近邻分配、EMA 码本更新与辅助优化启发式;PNN-VQ 用基于 softmax 重参数化的梯度直接优化,质心由分配到该质心的语料点均值给出。 PNN-VQ 与带承诺损失的 VQ-VAE 相比取得相当、在若干情形更低的测试 MSE;在 MLP 消融中,PNN-VQ 在低原型数时明显退化,随原型数增加而恢复。

在预训练大骨干上,token 级 PNN 把每个空间 token 投影到共享多面体,线性分类器作用于投影的 token 均值,因此每个 token 对预测有精确的加性贡献。 与按原型相似度给 patch 打分的原型网络不同,PNN 用 token 自身在多面体中的坐标分解,不需要原型专用损失。 ImageNet-100 上 ConvNeXt-T 保留 96.9% 的无约束精度,ImageNet-1k 上同样保留 96.9%;冻结紧凑 ViT 的 PNN 头在 CIFAR-10 上达到 94.46%,无约束分类器为 94.47%。

启示与展望

该结果面向需要在训练中直接获得可检查表示分解的研究者与工程师,适用于图像分类、图像重建与表格分类等已测试设置,也适用于预训练视觉骨干的 token 表示。按论文建议,使用时应系统评估原型数 K 对性能的影响,并在可解释性优先时选择性能足够的较小 K;论文当前对每个 PNN 层使用相同的 K,逐层取值留待后续探索。对于线性下游分类器,原型贡献可直接追溯到类别分数;对于非线性下游网络,坐标描述的是中间表示,其对输出的影响需通过后续计算评估。

读者仍需关注:原型数 K 与语料规模在不同数据上的选择策略尚未系统化;逐层不同 K 的高效搜索仍是开放问题;在非线性下游网络中,坐标到输出的因果影响需要通过后续计算评估;论文报告的重建代价随 K 增大而缩小,但未给出跨数据集的统一权衡准则;此外,token 级 PNN 在 ImageNet 上使用摊销器而不做 SMO 精化,其与精化版本的差异值得进一步观察。

来源