Putting Registers to Work: Task Registers for Token Pruning in Vision Transformers
Hongsen Cao, Mona Jaber, Shanxin Yuan, Ahmed Sayed
Queen Mary University of London
cs.CV, cs.AI
Submitted: 2026-08-11
Updated: 2026-08-12
Comments: 24 pages, 9 figures. Includes supplementary material
License: http://creativecommons.org/licenses/by/4.0/
Importance score: 100/100
The gist: This paper investigates whether token-pruning policies transfer across different vision tasks (image classification, semantic segmentation, and object detection) when using pretrained Vision
Terminology
Summary
This paper investigates whether token-pruning policies transfer across different vision tasks (image classification, semantic segmentation, and object detection) when using pretrained Vision Transformers. The authors conduct controlled probes that freeze the no-pruning checkpoint and apply parameter-free reduction criteria at one eligible layer at a time without retraining. The probes reveal three key differences: segmentation and detection rank the criteria differently, classification is especially sensitive to attention-based pruning in the earliest layers, and the dense tasks prefer opposite recovery endpoints.
These findings motivate Task-Adaptive Pruning (TAP), which introduces one task register per task and activates only the current one.
The evolving register state ranks tokens, distributes an exact removal budget over depth, and sets the recovery scale for dense features.
At a final keep rate of ρ = 0.5, the jointly adapted model TAP-J reaches 47.0 mIoU at 1.30× encoder throughput on ADE20K and 53.7 box AP at 1.32× encoder throughput on COCO while remaining competitive on ImageNet-1K.
The paper notes that Self-attention in a Vision Transformer (ViT) scales quadratically with the number of tokens
and that Token pruning lowers this cost by removing low-ranked patch tokens before later blocks.
Most pruning criteria are developed for image classification, yet pretrained ViT backbones are now routinely adapted to segmentation and detection.
The authors argue that Applying the same pruning policy across these pipelines assumes that token importance transfers across tasks,
and when this assumption fails, equal token budgets need not preserve equal amounts of task-relevant information.
The controlled experiments show that the effectiveness of a pruning criterion changes across network depths and tasks.
For classification, attention-based selection causes large drops in the earliest layers,
but a spatially distributed criterion (coverage, implemented with Farthest Point Sampling) is more stable.
For segmentation and detection, the relative advantage of attention and coverage changes with depth and differs between the two tasks.
Figure 1(c) reveals a negative rank correlation between segmentation and detection tasks,
with the criterion having the smallest detection drop ranking among the poorer choices for segmentation. The recovery probe shows that restoring this difference improves the segmentation task only, while using the stand-in alone improves the detection task.
The contributions are threefold: (1) a controlled framework that freezes each no-pruning pipeline and changes one pruning choice at one eligible layer, showing that segmentation and detection rank pruning criteria differently, spatial coverage is more stable than attention in the early classification layers, and the dense pipelines favor opposite recovery endpoints
; (2) task registers as active, task-specific controllers of sparse computation rather than task-agnostic stores for feature artifacts
; (3) TAP's task-conditioned pruning process that treats token removal and feature recovery as a unified operation.
The paper reviews token reduction criteria including EViT (retains tokens receiving high attention from the class token), DynamicViT (predicts token retention probabilities), Evo-ViT (maintains slow and fast token sets), ATS (samples tokens by attention scores), ToMe (bipartite matching for token merging), DiffRate (learns layerwise compression rates), and TPS (fuses information from pruned tokens). For task-aware pruning, BAT preserves attentive tokens while merging similar inattentive tokens,
Token Cropr learns task-relevant scores through auxiliary prediction heads,
and VLTP conditions token relevance on vision-language guidance.
The paper distinguishes TAP by assigning one register to each task and using the active register to jointly condition token selection, layerwise budget allocation, and recovery scaling.
A Vision Transformer splits an image into patch tokens and may prepend a class token c. Let X(l) = [x1(l),..., xNl(l)] denote patch tokens entering layer l, where xi(l) ∈ Rd and Nl is the number of patch tokens. The model processes one task t ∈ T at a time and prunes at M ordered layers, L = l1 <... < lM. At each pruning layer, tokens are scored and hard-selected before self-attention, so only the retained patch tokens are passed into the block.
For each task t ∈ T, TAP adds a learned initial register rt(0) ∈ Rd whose dimension matches that of the patch tokens and class token.
Standard transformer blocks update the register together with the image tokens. Neither the class token nor the register is pruned.
The evolving register state drives all pruning decisions.
At a pruning layer, TAP normalizes the pre-attention states and reuses the block's query and key projections to compute a pre-softmax register-to-patch score.
The score is computed as:
fi(l) = Σh=1H (Wq(l,h) r̄t(l))⊤ (Wk(l,h) x̄i(l)) / √dh
where Wq(l,h) and Wk(l,h) are the query and key projections for head h, and dh = d/H. For each head, scoring forms one register query and one key for each candidate patch by reusing the block's existing projections, so it introduces no additional projection parameters.
At inference, TAP ranks patches by decreasing fi(l) and resolves ties by ascending original patch index. Only the patches in S(l) enter full self-attention and the feed-forward sublayer together with the class token and task register.
A learned allocation must satisfy the requested global token budget exactly.
Let N1 denote the initial number of patch tokens. A target keep rate ρ ∈ (0, 1] sets the final integer keep count K and the removal budget R = N1 − K. At pruning layer l, B(l) denotes the unspent budget, and C(l) = min B(l), Nl − 1 is the removal capacity. The sigmoid of a shared linear readout (w, b) gives the fraction of B(l) proposed for that layer.
The algorithm produces an exact integer allocation, with the final layer consuming the remainder. The paper proves that the recursion removes exactly R = N1 − K patches and leaves exactly K ≥ 1 patches after the final pruning layer.
Hard top-k selection blocks gradients to the scoring and allocation readouts.
TAP trains with a cardinality-constrained straight-through mask
using Gumbel perturbations: zi(l) = fi(l) + γi, where γi are independent standard Gumbel samples. The soft keep probabilities satisfy a cardinality constraint solved by bisection. The straight-through mask is:
mi(l) = stopgrad(mi,hard(l) − mi,soft(l)) + mi,soft(l)
The forward value is exactly the retained sequence because mi(l) = 1 on the gathered indices.
The temperature is annealed during training, and Gumbel noise is omitted at inference. Because the hard forward pass already enforces the target budget, no rate loss is needed.
For each removed token i ∈ P(l), TAP selects the retained stand-in with the largest key-space cosine similarity
:
π(i) = arg maxj∈S(l) k̃i(l)⊤ k̃j(l)
The offset δi(l) = xi(l) − xπ(i)(l) is stored. At a dense read point L, the removed position is reconstructed as:
x̂i(L) = xπ*(i)(L) + αt(l) δi(l), αt(l) = σ(u⊤ rt(l) + v)
where π*(i) follows the pointer chain to a token that survives at the read point. Reconstruction combines the original offset δi with the final surviving endpoint and omits offsets from intermediate stand-ins.
The shared readout controls how much of the early offset is added to the later endpoint.
Every surviving token retains its original patch index.
Reconstructed features are consumed only by the task head and never reenter the backbone.
Classification reads the class token and skips recovery.
Detection prunes only at global-attention blocks,
with each subsequent window block grouping surviving tokens by their original window and applying variable-length attention to every nonempty group.
TAP-J adds one register and two readouts, totaling 3d + 2 = 2,306 pruning-specific parameters for ViT-B.
Recovery stores one pointer and one d-dimensional offset per removed token.
TAP-F stores one frozen backbone and adds approximately 1.2M low-rank parameters and one register per task.
The paper evaluates ImageNet-1K classification, ADE20K semantic segmentation, and COCO object detection and instance segmentation. All models use a ViT-B backbone with MAE pretraining. Classification uses a linear head, segmentation uses a SegFormer MLP head, and detection uses ViTDet with Cascade Mask R-CNN. Classification and segmentation prune at three evenly spaced layers. Detection prunes tokens in global attention blocks.
The probe freezes the corresponding no-pruning backbone and head, prunes each eligible layer in turn at a fixed token count without retraining.
The parameter-free criteria span attention, activation magnitude, spatial coverage, and pairwise feature affinities used for merging.
Across 11 criteria, the Spearman correlation between the two dense-task rankings is −0.62.
Coverage ranks first for segmentation but last for detection,
while key redundancy shows the opposite pattern.
At ρ = 0.5, TAP-J achieves:
-
Classification: 83.2 top-1 accuracy (Δ −0.4), 1.26× encoder speedup
-
Segmentation: 47.0 mIoU (Δ −0.2), 1.30× encoder speedup
-
Detection: 53.7 box AP (Δ −0.3), 1.32× encoder speedup
TAP-J has the smallest accuracy drop on segmentation and detection among the fully fine-tuned methods.
Compared with the static rule, TAP-J gains 0.1 point on classification, 1.3 mIoU on segmentation, and 0.5 box AP on detection.
TAP-F reduces task-specific storage by sharing the frozen base,
trailing no pruning by 1.0 top-1 points, 1.4 mIoU, and 1.5 box AP.
Table 4 tests the three register-driven choices. Replacing Equation (1) with incoming patch attention loses 0.9 mIoU and 0.5 box AP.
Using the initial register instead of its evolving state costs 0.4 mIoU and 0.3 box AP.
Zero-filling is worse than both recovery endpoints, and the preferred endpoint differs between segmentation and detection.
Sharing the initial register across task-specific adapters costs 0.2 top-1 points, 1.3 mIoU, and 1.1 box AP.
At ρ = 0.5, TAP-J reaches encoder speedups of 1.26×, 1.30×, and 1.32× with corresponding losses of 0.4 top-1 points, 0.2 mIoU, and 0.3 box AP.
The register query, budget readout, and top-k bookkeeping account for 0.7% of encoder runtime.
Standard deviations are at most 0.24, and peak inference memory is lower for every task and regime.
Figure 5(a) shows that TAP removes fewer tokens early and learns different mean schedules for segmentation and detection.
Figure 5(b) places segmentation near the full-offset endpoint and detection near the stand-in endpoint.
Allocation and recovery are computed per image from the register.
The evidence argues against treating token reduction as a single transferable rule.
TAP-J trails DiffRate by 0.1 point on classification because Classification reads a class token and skips dense recovery, leaving one of TAP's three mechanisms unused.
The paper suggests finer-grained task-adaptive computation
for future work, including image-conditioned budgets, multiple stand-ins, cross-layer transport, and matched pipelines.
Pruning behavior changes with task and depth, which motivates us to propose Task-Adaptive Pruning (TAP).
Rather than relying on a pruning rule tailored to one visual task, TAP puts an evolving task register to work as the controller of a shared sparse computation process under an exact global budget.
The design maintains a favorable balance between predictive performance and encoder throughput across tasks, making TAP particularly well suited to unified vision systems that serve multiple tasks through a shared backbone.
Improvements for AI systems
Improvements to AI Systems Based on This Paper:
-
Task-Conditioned Sparse Computation Controllers: Replace static, task-agnostic token pruning policies with learned
task registers
that dynamically control token selection, layerwise budget allocation, and feature recovery. This enables a single shared backbone to serve multiple vision tasks (classification, segmentation, detection) with task-specific sparse computation, improving efficiency without sacrificing task-relevant information. -
Unified Pruning and Recovery Mechanism: Implement a joint operation where token removal and dense-feature reconstruction are controlled by the same task register. This allows the system to adaptively choose between recovering pruned tokens via offset-based reconstruction (beneficial for segmentation) or using stand-in tokens (beneficial for detection), based on the active task.
-
Exact Budget Allocation with Learned Depth Scheduling: Use a learned, cardinality-constrained allocation that distributes a global token removal budget across network depth, rather than using fixed per-layer rates. This lets the system automatically learn task-specific schedules (e.g., prune fewer tokens early for segmentation/detection) while guaranteeing exact final token counts.
-
Depth- and Task-Aware Criterion Selection: Replace single pruning criteria with task- and depth-sensitive scoring. The system learns to prefer spatial coverage (Farthest Point Sampling) in early classification layers, while dynamically switching between attention-based and coverage-based selection for dense tasks based on layer depth and task identity.
-
Parameter-Efficient Multi-Task Adaptation: Share a frozen pretrained backbone across tasks while adding only one task register and two lightweight readouts (≈2,306 parameters for ViT-B) for pruning control. This enables unified vision systems to serve multiple tasks with minimal task-specific storage and near-zero overhead (0.7% of encoder runtime).
-
Dense-Feature Reconstruction with Task-Adaptive Endpoints: For segmentation and detection, reconstruct pruned token features at the readout by combining the original offset with a task-controlled scaling factor applied to the surviving endpoint. This improves dense prediction quality by preserving spatial detail without re-entering the backbone.
-
Stable Training via Cardinality-Constrained Straight-Through Estimation: Use Gumbel-perturbed soft masks with bisection-based cardinality constraints and straight-through gradient estimation. This avoids rate-loss tuning, ensures exact budget enforcement during training, and stabilizes learning across tasks.
What the Improved AI System Can Do:
-
Serve image classification, semantic segmentation, and object detection simultaneously through one shared ViT backbone, with task-specific sparse computation that adapts pruning behavior per task and per layer.
-
Achieve 1.26–1.32× encoder throughput at 50% token retention while maintaining near-baseline accuracy (e.g., 47.0 mIoU on ADE20K, 53.7 box AP on COCO, 83.2% top-1 on ImageNet-1K), outperforming static pruning rules by up to 1.3 mIoU and 0.5 box AP.
-
Dynamically allocate token removal budgets across depth per image and task, learning to preserve more tokens early for dense tasks and adapt recovery strategy (offset vs. stand-in) based on task needs.
-
Reconstruct dense feature maps for segmentation/detection heads with task-adaptive scaling, improving spatial fidelity without extra computational cost in the backbone.
-
Extend to new tasks by adding a single register and readout, enabling scalable multi-task deployment with minimal parameter overhead and no retraining of the shared backbone.
Abstract
Token-pruning policies are usually designed for a single recognition pipeline, but pretrained Vision Transformers are reused across tasks with different spatial demands. We ask which parts of a pruning policy transfer across image classification, semantic segmentation, and object detection. For each pipeline, controlled probes freeze the no-pruning checkpoint and apply a series of parameter-free reduction criteria at one eligible layer at a time without retraining. The probes reveal three differences: segmentation and detection rank the criteria differently, classification is especially sensitive to attention-based pruning in the earliest layers, and the dense tasks prefer opposite recovery endpoints. These findings motivate Task-Adaptive Pruning (TAP). Existing register tokens serve as task-agnostic storage for feature artifacts. TAP instead introduces one task register per task and activates only the current one. Its evolving state ranks tokens, distributes an exact removal budget over depth, and sets the recovery scale for dense features. At a final keep rate of rho=0.5, our jointly adapted model, TAP-J, reaches 47.0 mIoU at 1.30 times encoder throughput on ADE20K and 53.7 box AP at 1.32 times encoder throughput on COCO while remaining competitive on ImageNet-1K.
Sources
Related papers
- Loss Knows Best: Detecting Annotation Errors in Videos via Loss Trajectories
- AnchorWeave: World-Consistent Video Generation with Retrieved Local Spatial Memories
- Benchmarking the Robustness of Foundation Models for Mammography under Domain Shift
- MambaX-Net: Dual-Input Mamba-Enhanced Cross-Attention Network for Longitudinal MRI Segmentation
- TeleOCR: Navigating Document Parsing Across Digital and Camera-Captured Documents
- A Survey on Efficient Vision-Language-Action Models