Skip to main content
Back to timeline
arXivSource publication:

Polytopal Neural Networks constrain latent representations to learned polytopes, retaining 96.9% of unconstrained accuracy on ImageNet-1k

Synopsis

The authors propose Polytopal Neural Networks (PNNs), which constrain each layer's latent representation during training to a convex combination of learned archetypes so that the explanatory coordinates are the same coordinates the next layer receives; across image classification and reconstruction benchmarks, PNNs approach unconstrained networks with small accuracy loss, outperform VQ-VAE and Dirichlet VAE in unsupervised reconstruction, and offer a VQ training route without EMA updates, commitment loss, or straight-through estimator.

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

Interpretation

PNN inserts a simplex-constrained projection after selected layers, making each representation a convex combination of learned archetypes, with the constraint participating directly in the forward computation rather than serving as post-hoc explanation. Prior deep archetypal approaches typically place archetypes at a single bottleneck or enforce archetypal structure only through a penalty; PNN enforces AA constraints exactly at multiple layers and makes the explanatory coordinates identical to the next layer's input coordinates. The paper provides Lemma 1 and Theorem 1 showing that under positive definiteness the projection is a continuous piecewise affine map representable by a finite ReLU network; experiments span MNIST, FashionMNIST, CIFAR-10, SVHN, EuroSAT, MedMNIST, and ImageNet.

Scalable training combines a compact corpus learned jointly with the network and updated from mini-batches, plus amortized simplex inference. The corpus decouples archetypes from the full dataset, so gradient memory scales with batch and corpus size rather than sample count; the amortizer is trained only on the projection residual and never receives the task loss, so it cannot deform the encoder, head, corpus, or archetypes. Mini-batched training achieves accuracy comparable to full-data updates; replacing the deployed amortizer with exact SMO projection leaves all 150 configurations on the diagonal in test MSE.

Replacing the simplex with one-hot coordinates reduces PNN to a VQ model whose codebook is grounded in the corpus and trained without EMA updates, a commitment loss, or a straight-through estimator. Standard VQ-VAE relies on hard nearest-neighbor assignment, EMA codebook updates, and auxiliary optimization heuristics; PNN-VQ optimizes directly via the softmax reparameterization, with centroids given by the mean of corpus points assigned to each centroid. PNN-VQ achieves comparable and in several cases lower test MSE than VQ-VAE with a commitment loss; in the MLP ablation, PNN-VQ degrades markedly at low archetype counts and recovers as the count increases.

On pretrained large backbones, token-level PNN projects each spatial token onto a shared polytope, and a linear classifier acts on the token mean of the projections, so every token makes an exact additive contribution to the prediction. Unlike prototype networks that score patches by similarity to learned prototypes, PNN represents each token by its own coordinates in the polytope spanned by archetypes, requiring no prototype-specific losses. ConvNeXt-T retains 96.9% of unconstrained accuracy on ImageNet-100 and 96.9% on ImageNet-1k; a PNN head on a frozen compact ViT reaches 94.46% on CIFAR-10 versus 94.47% for the unconstrained classifier.

Perspective

The result targets researchers and engineers who need inspectable representation decompositions obtained during training, and applies to the tested settings of image classification, image reconstruction, and tabular classification, as well as token representations of pretrained vision backbones. As the paper recommends, users should systematically assess the effect of the archetype count K on performance and, when interpretability is prioritized, select the smallest K with sufficient performance; the current work uses the same K for each PNN layer, leaving per-layer values to future exploration. For a linear downstream classifier, archetype contributions can be traced directly to class scores; for nonlinear downstream networks, the coordinates describe the intermediate representation and their effect on the output must be assessed through that subsequent computation.

Open questions a careful reader would still watch: strategies for choosing the archetype count K and corpus size across datasets are not yet systematized; efficient search for per-layer K remains open; in nonlinear downstream networks, the causal effect of coordinates on outputs must be assessed through subsequent computation; the paper reports that reconstruction cost shrinks as K grows but does not give a unified trade-off criterion across datasets; and token-level PNN on ImageNet uses the amortizer without SMO refinement, so its difference from the refined version merits further observation.

Sources