Papers
Topics
Authors
Recent
Search
2000 character limit reached

SAFE-KD: Risk-Controlled Early Exit for Vision Models

Updated 10 February 2026
  • The framework SAFE-KD is a universal early-exit system that combines hierarchical knowledge distillation with conformal risk control to guarantee statistically bounded selective risk.
  • It attaches intermediate classifier exits to any vision backbone (CNN or ViT) and employs decoupled knowledge distillation alongside consistency regularization for calibrated risk control.
  • Empirical results show up to a 45% reduction in expected inference depth while maintaining or surpassing full-inference accuracy, ensuring efficiency and robustness.

SAFE-KD offers a universal, risk-controlled early-exit framework for modern vision backbones, combining hierarchical knowledge distillation with conformal risk control (CRC) to achieve statistically guaranteed bounds on selective misclassification risk for early-exit architectures. It enables substantial reductions in inference cost via early stopping for "easy" samples, while maintaining user-specified upper bounds on misclassification risk at each exit, calibrated on finite data. SAFE-KD is model-agnostic and deploys on a variety of convolutional (CNN) and transformer-based (ViT) image models (Khazem, 3 Feb 2026).

1. Architecture and Components

SAFE-KD is structured as a lightweight "wrapper" atop any standard vision backbone, supporting both CNNs and Vision Transformers. Its architecture comprises:

  • Base Backbone: Any pretrained or trainable vision model f(⋅)f(\cdot) (e.g., ResNet, ConvNeXt, ViT, Swin).
  • Intermediate Exit Heads: At KK select depths, SAFE-KD attaches classifiers producing logits zj(x)∈RCz_j(x)\in\mathbb{R}^C (j=1,…,Kj=1,\ldots,K), class probabilities pj(c∣x)=softmax(zj(x))cp_j(c|x)=\mathrm{softmax}(z_j(x))_c, and a confidence score (typically, Maximum Softmax Probability MSPj(x)=max⁡cpj(c∣x)\text{MSP}_j(x)=\max_c p_j(c|x)). For CNNs, exits use global average pooling and an optional MLP before a fully connected (FC) layer; for ViTs, exits use CLS or mean token pooling, optional LayerNorm, then FC.
  • Teacher Network: An Exponential Moving Average (EMA) of the full model serves as the teacher for knowledge transfer.

This configuration allows SAFE-KD to operate agnostically across architectures, minimally increasing inference overhead.

2. Decoupled Knowledge Distillation and Consistency

Training leverages hierarchical Decoupled Knowledge Distillation (DKD), coupled with deep-to-shallow consistency regularization:

  • DKD Loss: For each exit jj, knowledge is distilled from the teacher logits zT(x)z_T(x) to zj(x)z_j(x) using a split-KL loss:

LDKD(zj,zT,y)=λtar KL(pyT∥pj,yS)+λnon KL(p\yT∥pj,\yS)\mathcal{L}_{\mathrm{DKD}}(z_j,z_T,y) = \lambda_{\text{tar}}~\mathrm{KL}(p^T_y \parallel p^S_{j,y}) + \lambda_{\text{non}}~\mathrm{KL}(p^T_{\backslash y} \parallel p^S_{j,\backslash y})

where KK0 and KK1 are the (teacher, student) probabilities on the target class KK2 (ground-truth label), and KK3 and KK4 are normalized distributions over all non-target classes.

  • Consistency Regularization: To align intermediate exits with the final head, SAFE-KD adds

KK5

with weighting KK6, regularizing posterior agreement between each intermediate exit and the ultimate exit (KK7).

  • Total Loss: For weights KK8 summing to KK9, full training minimizes:

zj(x)∈RCz_j(x)\in\mathbb{R}^C0

where zj(x)∈RCz_j(x)\in\mathbb{R}^C1 is a scaling factor for DKD.

This hierarchical approach increases calibration, depth-to-exit consistency, and maintains high accuracy at all exits.

3. Conformal Risk Control for Early-Exit Thresholds

At inference, SAFE-KD employs Conformal Risk Control (CRC) to set data-driven confidence thresholds at each exit, guaranteeing a user-specified selective risk:

  • Nonconformity Score: zj(x)∈RCz_j(x)\in\mathbb{R}^C2 at exit zj(x)∈RCz_j(x)\in\mathbb{R}^C3.
  • Acceptance Set: zj(x)∈RCz_j(x)\in\mathbb{R}^C4.
  • Selective Misclassification Risk: zj(x)∈RCz_j(x)\in\mathbb{R}^C5.
  • Threshold Calibration: Using a held-out calibration set zj(x)∈RCz_j(x)\in\mathbb{R}^C6, thresholds zj(x)∈RCz_j(x)\in\mathbb{R}^C7 are chosen so the conformal upper bound:

zj(x)∈RCz_j(x)\in\mathbb{R}^C8

does not exceed the desired risk level zj(x)∈RCz_j(x)\in\mathbb{R}^C9. Here j=1,…,Kj=1,\ldots,K0.

CRC, under the exchangeability assumption, ensures

j=1,…,Kj=1,\ldots,K1

for each exit, providing finite-sample statistical guarantees.

4. Safe Inference Policy and Practical Deployment

At test time, early exit is governed by the following procedure:

  • Proceed through exits j=1,…,Kj=1,\ldots,K2, checking at each if j=1,…,Kj=1,\ldots,K3.
  • The first such j=1,…,Kj=1,\ldots,K4 is used for prediction. If none, inference proceeds to the final exit j=1,…,Kj=1,\ldots,K5.
  • For every exit j=1,…,Kj=1,\ldots,K6, the empirical misclassification risk among samples exiting there is guaranteed not to exceed j=1,…,Kj=1,\ldots,K7 (up to sampling correction).

This allows the system designer to select j=1,…,Kj=1,\ldots,K8 according to operational requirements, trading off computational savings for tightly controlled selective risk. All calibration is based on a held-out set.

5. Empirical Evaluation and Results

SAFE-KD has been empirically validated across six architectures (ResNet-50, MobileNetV3-S, EfficientNet-B0, ConvNeXt-T, ViT-S, Swin-T) and multiple image datasets (CIFAR-10/100, STL-10, Pets, Flowers102, Aircraft), delivering:

  • Compute-Accuracy Trade-offs: At j=1,…,Kj=1,\ldots,K9 risk, SAFE-KD achieves pj(c∣x)=softmax(zj(x))cp_j(c|x)=\mathrm{softmax}(z_j(x))_c0--pj(c∣x)=softmax(zj(x))cp_j(c|x)=\mathrm{softmax}(z_j(x))_c1 lower expected depth while matching or surpassing full-inference accuracy. Baseline methods (fixed MSP or entropy thresholds) violate the risk constraint, with observed risks pj(c∣x)=softmax(zj(x))cp_j(c|x)=\mathrm{softmax}(z_j(x))_c2--pj(c∣x)=softmax(zj(x))cp_j(c|x)=\mathrm{softmax}(z_j(x))_c3.
  • Calibration: SAFE-KD reduces negative log-likelihood (NLL) and expected calibration error (ECE) at all exits.
  • Risk Guarantee: Across sweeps in pj(c∣x)=softmax(zj(x))cp_j(c|x)=\mathrm{softmax}(z_j(x))_c4, the observed per-exit risk tracks the theoretical bound pj(c∣x)=softmax(zj(x))cp_j(c|x)=\mathrm{softmax}(z_j(x))_c5, confirming pj(c∣x)=softmax(zj(x))cp_j(c|x)=\mathrm{softmax}(z_j(x))_c6 tightness.
  • Robustness: On CIFAR-10-C corrupted data (severity 3), SAFE-KD attains lower mean corruption error (mCE) at both shallowest and deepest exits compared to comparable multi-exit and DKD-based models, e.g., mCE at exit 1: SAFE-KD pj(c∣x)=softmax(zj(x))cp_j(c|x)=\mathrm{softmax}(z_j(x))_c7, DKD pj(c∣x)=softmax(zj(x))cp_j(c|x)=\mathrm{softmax}(z_j(x))_c8, MultiExit pj(c∣x)=softmax(zj(x))cp_j(c|x)=\mathrm{softmax}(z_j(x))_c9; at final exit: SAFE-KD MSPj(x)=max⁡cpj(c∣x)\text{MSP}_j(x)=\max_c p_j(c|x)0.
  • Ablation Findings: Removing DKD degrades accuracy by MSPj(x)=max⁡cpj(c∣x)\text{MSP}_j(x)=\max_c p_j(c|x)1 and forces safer (more conservative) thresholds, increasing average depth. Removing consistency (MSPj(x)=max⁡cpj(c∣x)\text{MSP}_j(x)=\max_c p_j(c|x)2) triggers higher exit-variance, though risk guarantees persist.
  • Example Table (for CIFAR-100, ResNet-50, MSPj(x)=max⁡cpj(c∣x)\text{MSP}_j(x)=\max_c p_j(c|x)3):
Method Accuracy Exp. Depth Observed Risk
Fixed MSP 81.5% 0.72 6.8% (Unsafe)
Entropy gate 80.9% 0.65 7.5% (Unsafe)
SAFE-KD (CRC) 82.3% 0.59 4.8% (Safe)

SAFE-KD consistently defines the empirical Pareto frontier for the target risk constraint across tasks.

6. Calibration, Robustness, and Risk Guarantees

SAFE-KD's deployment of CRC uniquely enables it to deliver finite-sample, statistically-tight risk control not attainable with heuristic thresholds. Reliability diagrams confirm alignment of observed risk to the target MSPj(x)=max⁡cpj(c∣x)\text{MSP}_j(x)=\max_c p_j(c|x)4 across exit depths. Selective risk curves for a sweep of MSPj(x)=max⁡cpj(c∣x)\text{MSP}_j(x)=\max_c p_j(c|x)5 show empirical MSPj(x)=max⁡cpj(c∣x)\text{MSP}_j(x)=\max_c p_j(c|x)6 at or just under MSPj(x)=max⁡cpj(c∣x)\text{MSP}_j(x)=\max_c p_j(c|x)7.

For corrupted or hard samples, the framework naturally "defers" to deeper exits, preserving guaranteed selective risk at cost of additional computation. This property, together with out-of-the-box calibration from DKD and consistency, distinguishes SAFE-KD from prior early-exit and distillation methods without such formal risk control.

7. Summary and Broader Impact

SAFE-KD constitutes a general-purpose, modular extension to vision models requiring minimal architectural modification and no retraining of the backbone. Its integration of CRC, DKD, and deep-to-shallow consistency provides user-tunable, quantifiable risk guarantees for early exiting, fine-grained calibration, and enhanced robustness under dataset shift or corruption. The framework supports frequent regression testing and online adaptation to evolving operational requirements, making it especially suitable for resource-constrained or safety-critical deployments (Khazem, 3 Feb 2026).

Definition Search Book Streamline Icon: https://streamlinehq.com
References (1)

Topic to Video (Beta)

No one has generated a video about this topic yet.

Whiteboard

No one has generated a whiteboard explanation for this topic yet.

Follow Topic

Get notified by email when new papers are published related to SAFE-KD.