Skip to main content
Back to timeline
arXivSource publication:

One Transformer block, reused at every depth: a depth coordinate steers an expert bank, and B/16 matches DeiT III with about 70% fewer stored parameters

Lead

A vision transformer's entire depth stack is compressed into one recurrent Transformer block, where a normalized depth coordinate softly merges a small expert bank into one FFN per recurrent depth, letting B/16 match DeiT III at comparable inference FLOPs with about 70% fewer stored parameters.

Source-provided article image: One Block, Multiple Depths: Recurrent Vision Transformers with Depth-Programmed Experts
arxiv.org

Story

A vision encoder's depth-wise stack can be replaced by a single Transformer block applied recurrently, with depth-specific computation supplied by different FFN weights at each recurrent depth. Vision transformers were previously built as fixed-depth stacks with independently parameterized blocks; weight-sharing schemes either targeted compression or retained several distinct blocks. Under supervised ImageNet-1k training, reViT-B/16 reaches 83.0% top-1 versus DeiT III's 82.8% at near-matched inference FLOPs, while storing 23.6M parameters against 86.6M.

The FFN at each recurrent depth is a convex combination of a small expert bank, and the mixing coefficients depend only on that depth's normalized coordinate, so the same coefficients are shared by every image and token. Earlier mixture-of-experts methods routed at the token level, sending each token through a different expert subset without merging expert parameters before execution. With the same recurrent backbone, task loss, and base recipe, seven MoE mechanisms were compared; at a nominal budget of one dense FFN, the depth-programmed merge outperformed the tested token-dispatch and output-mixture alternatives.

One trained checkpoint can run at multiple inference depths by resampling the same normalized coordinate interval. Changing the depth of a fixed-depth model previously required retraining or a separate checkpoint. Under distillation, the depth-conditioned variant improved on ADE20k, ImageNet linear probe, and NYUv2 as depth increased, while a feature-only control stayed nearly flat on ImageNet and ADE20k.

What to watch

Readers who need fixed-depth deployment can precompute the merged FFNs and write them into a conventional dense graph, removing online routing and merging at the cost of materializing one FFN per depth. Readers who want to change inference depth can resample the same checkpoint on the normalized coordinate interval without retraining. Readers who want more stored capacity without raising per-step dense FFN computation can enlarge the expert bank; for S/16, going from 4 to 8 experts improves ImageNet top-1 by 2.0 points.

The latency measurements use a direct PyTorch implementation that reconstructs merged weights with generic FP32 reductions, without cross-input caching or a custom kernel, so they characterize the current implementation rather than an optimized deployment. The dynamic form reduces total runtime memory by about half at batch 1 but incurs a latency overhead, which rises to about 39% at batch 64 where total memory becomes similar. In the distillation comparison, reViT matches only the teacher's final-layer features while Raptor uses intermediate features, so Table 2 is a reference comparison rather than a controlled ablation. Expert-deletion analysis shows the largest alignment losses come from deleting the earliest-used experts, so deletion sensitivity alone cannot separate an expert's role from its position in the recurrence.

Sources