ProtoSeam: Lifting Classifier Training with Latent Gaussian Mixture Models
Authors: Robert Lampel, Timon Klein, Sebastian Sager
Organizations: Department of Mathematics, Otto von Guericke University (Magdeburg, Germany) · Max Planck Institute for Dynamics of Complex Technical Systems (Magdeburg, Germany)
We propose a lifted reformulation of supervised classification that improves the final accuracy of standard classifiers without changing the architecture at inference time. A network N=N2∘N1 is split at a single semantic interface and one learnable prototype per class is inserted there. Training combines a quadratic consensus penalty that pulls N1(x) toward the prototype of its class with a classification loss of N2 evaluated on samples drawn around the prototypes, whereat no gradient crosses the interface. At inference the prototypes are discarded and the unmodified network N2∘N1 is used. Across CIFAR-10, CIFAR-100, and TinyImageNet with ResNet and vision transformer backbones, lifted training improves test accuracy by up to five percentage points over variants without lifting under a shared tuning protocol. Moreover, we provide theoretical justification of those results.
Figures & tables
Figure 1: Output of the trained first network part N1 (left) and the sampled input for the second part N2 (right) for the MNIST ( Deng, 2012 ) dataset. The stars denote the position of the lifted variables, i.e., the class means of the sampled normal distribution. Here, N1 consists of two convolutional layers, mapping to R2 . N2 consists of one fully connected layer from R2 to R10 .
Figure 2: The lifting procedure. The original network (top) is split into two parts and one lifted variable is introduced per class (bottom). The part N1 is trained against si by the consensus term (orange), while N2 is trained on samples zi drawn around si , so no gradient crosses the interface (seam). The class-conditional covariance (teal) is estimated once per epoch from the embeddings N1(x) of the whole class, regularized by ( 3 ), and then used as the scale of the sampling step ( 6 ). At inference lifted variables and sampling step are discarded and the original path N2∘N1 is restored.
Figure 3: ViT-S on CIFAR-100, using the empirically computed covariance. On the left hand side we compare the test accuracy for different values of ρmax and a fixed embedding dimension of k=32 . The right plot shows the test accuracy for different embedding dimensions, this time for a fixed penalty of ρmax=16 . For both ablations we observe a monotonic climb which first increases steeply and then stabilizes. There is no meaningful difference between the test and validation accuracies. We use ρmax=16 and k=32 in the following, which we found to work best across all datasets.
Figure 4: ViT-S on TinyImageNet: validation accuracy over the epochs for baseline, unlifted, and lifted variant. We show the mean over 3 seeds ( 42,43,44 ), sampled every 5 epochs. The shaded bands show the full seed range (min to max) at each sampled epoch.
Dataset
Model
Variant
seed 42
seed 43
seed 44
Mean
Spread
CIFAR-10
ResNet-8
Baseline
94.88
95.09
94.77
94.91
0.32
Unlifted
95.09
95.38
95.00
95.16
0.38
ProtoSeam
95.27
95.54
95.57
95.46
0.30
ViT-S
Baseline
89.45
89.50
89.34
89.43
0.16
Unlifted
90.55
89.94
89.66
90.05
0.89
ProtoSeam
90.64
90.34
90.17
90.38
0.47
Table 1: Multi-seed confirmation (seeds 42/43/44): test accuracy (%) for baseline, unlifted, and lifted across both architectures and all three datasets. Bold = lifted’s mean, which beats both comparators at every individual seed on every combination.
Appendix figures & tables3 assets
Supplementary material from the paper’s appendix.
Appendix
Dataset
Model
Variant
η
λ
CIFAR10
ResNet8
Baseline
0.02
10−3
ProtoSeam
0.05
5×10−4
ViT
Baseline
0.02
10−4
ProtoSeam
0.01
5×10−4
CIFAR100
ResNet8
Baseline
0.05
2×10−3
ProtoSeam
0.02
2×10−3
Appendix
Table 2: Final, confirmed best hyperparameters for every (dataset, model, variant) combination determined by the grid search.
Figure 5: Left: the identity head g0 , the sampled optima ga⋆(σ) for σ∈{0.1,0.3,0.7} , and the small-noise limit g−1/2 . All curves pass through (−1,−1) and (1,1) ; the dotted lines are the tangents there, which are the only thing that distinguishes the family. Right: Monte-Carlo estimates of a⋆(σ) and of ∣ga⋆′(1)∣ against the closed forms ( 30 ) and ( 31 ). Markers are means over 25 seeds and error bars are one standard deviation; they are smaller than the markers at most noise levels.
Figure 6: Loss EY[(ga(Y+δ)−Y)2] under a bounded deterministic reconnection error δ , restricted to the neighborhood of the prototypes for which the local model is intended, at σ=0.2 . Grey curves are randomly drawn prototype-interpolating values of a , which illustrate that interpolation alone constrains nothing. Right: the same curves on log–log axes, showing the Θ(δ2) behavior of the identity head against the Θ(δ4) behavior of g−1/2 predicted by ( 40 ).