Choosing a backbone
Picking a vision backbone is a decision about your deployment target, your data size and your labels, and the leaderboard is the least useful input to it.
- 14 min read
- 3 reading levels
- Updated
Read these first
On this page 7
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Choosing a backbone means picking the pretrained vision model you will build on. The right answer depends on where it has to run.
Think about getting a fridge across town. An auto-rickshaw, a car, a tempo and a lorry are all available.
The lorry is the most capable vehicle by every measure. It is also the wrong choice if the lane to your building is three feet wide.
Model choice works the same way. The most accurate network is often the wrong one, because of where it has to fit.
What a backbone is
A backbone is the main body of a vision model: everything except the final layer that names the classes.
You take that body, already trained on millions of pictures, and attach your own small head. The body already knows about edges, textures, shapes and objects. You are only teaching the last part.
[ pretrained backbone ] -> [ your small head ] -> your classes
already knows you train this
about pictures on a few hundred examplesThis is why a small project can work at all. Training a body from scratch needs enormous data. Borrowing one needs very little.
The four questions that decide it
Where does it run? A phone, a browser, a server with a graphics card, or a tiny board on a factory line. This constrains the answer more than anything else.
How much time do you have per picture? A live camera needs an answer in a few tens of milliseconds. A nightly batch job can take a second.
How many labelled examples do you have? A few hundred means freeze the body and train only the head. Tens of thousands means you can afford to retrain the body too.
How different are your pictures from ordinary photos? X-rays, satellite images and microscope slides are far from holiday snaps, and the usual pretrained bodies help less.
The mistake nearly everyone makes
People pick the model at the top of the accuracy table. Then they discover it is too slow, or too big, or needs more labels than they have.
Two numbers usually decide it instead. How long does it take per picture on your actual machine. How much accuracy do you lose by going one size smaller.
Very often the second number is under two percent, and the first is a factor of five. That trade is worth taking almost every time.
The trap in the numbers
A model that does less arithmetic is not always faster. This catches everyone once.
The measurements further down show one model doing five times less arithmetic than another. On the same computer, it takes longer. Counting operations is not the same as timing them.
Time the model on the machine it will actually run on. Nothing else settles it.
Where you have already seen this
- A phone camera app choosing a small model so the preview stays smooth.
- A photo website using a large accurate model, because it runs overnight.
- A shop's billing camera using a mid-size model, because it must answer in real time.
Remember this
- Choose for the deployment target first, and accuracy second.
- Fewer operations does not mean faster; measure on the real hardware.
- With few labels, freeze the backbone and train only the head.
What to learn next
- Transfer learning in PyTorch — attaching your head and training it properly.
- ONNX — exporting the model you chose, and re-measuring after export.
- Model serving — putting it behind an interface people can call.
Developer — Code and libraries.
Setup
pip install "torch==2.5.1" "torchvision==0.20.1" "timm==1.0.26"Run against PyTorch 2.5.1, torchvision 0.20.1 and timm 1.0.26 on CPU, in about thirty seconds. No weights are downloaded: accuracy and operation counts come from metadata that ships with torchvision.
Get the real numbers before you choose
import statistics
import time
import torch
from torchvision.models import get_model, get_model_weights
CANDIDATES = ["mobilenet_v3_small", "efficientnet_b0", "resnet18", "resnet50",
"efficientnet_v2_s", "convnext_tiny", "swin_t", "vit_b_16"]
print("what torchvision publishes about each set of weights (no download needed)")
print(f"{'model':20s} {'weights':16s} {'params':>12s} {'GMACs':>7s} {'top-1':>7s}")
for name in CANDIDATES:
w = get_model_weights(name).DEFAULT
m = w.meta
print(f"{name:20s} {str(w).split('.')[-1]:16s} {m['num_params']:12,} "
f"{m['_ops']:7.2f} {m['_metrics']['ImageNet-1K']['acc@1']:7.2f}")
torch.set_num_threads(4)
print(f"\nmedian CPU latency, batch of 1, 224x224, {torch.get_num_threads()} threads")
print("these numbers are specific to this machine and this build: measure your own")
print(f"{'model':20s} {'ms/image':>9s} {'GMACs':>7s} {'GMACs per ms':>13s}")
for name in CANDIDATES:
model = get_model(name, weights=None).eval()
x = torch.zeros(1, 3, 224, 224)
times = []
with torch.no_grad():
for _ in range(5):
model(x) # warm-up, so we time steady state
for _ in range(20):
t0 = time.perf_counter()
model(x)
times.append((time.perf_counter() - t0) * 1000)
ms = statistics.median(times)
ops = get_model_weights(name).DEFAULT.meta["_ops"]
print(f"{name:20s} {ms:9.1f} {ops:7.2f} {ops / ms:13.3f}")
print("\nevery model wants its own preprocessing, and guessing costs accuracy silently")
for name in ["resnet50", "efficientnet_b0", "convnext_tiny", "swin_t"]:
t = get_model_weights(name).DEFAULT.transforms()
print(f" {name:18s} resize {t.resize_size[0]:>3d} -> crop {t.crop_size[0]:>3d}, "
f"{str(t.interpolation).split('.')[-1]}")
print("\npulling intermediate feature maps out, two ways")
from torchvision.models.feature_extraction import create_feature_extractor
from torchvision.models import resnet18
fx = create_feature_extractor(resnet18(weights=None), {"layer2": "s8", "layer4": "s32"})
out = fx(torch.zeros(1, 3, 224, 224))
print(" torchvision:", {k: tuple(v.shape) for k, v in out.items()})
import timm
tm = timm.create_model("resnet18", pretrained=False, features_only=True)
shapes = [tuple(o.shape) for o in tm(torch.zeros(1, 3, 224, 224))]
print(" timm channels :", tm.feature_info.channels())
print(" timm strides :", tm.feature_info.reduction())
print(" timm shapes :", shapes)
print(" timm data cfg :", timm.data.resolve_model_data_config(tm))
print(f"\ntimm version {timm.__version__}, "
f"{len(timm.list_models(pretrained=True)):,} pretrained checkpoints available")
print(" matching 'convnext*':", len(timm.list_models("convnext*", pretrained=True)))what torchvision publishes about each set of weights (no download needed)
model weights params GMACs top-1
mobilenet_v3_small IMAGENET1K_V1 2,542,856 0.06 67.67
efficientnet_b0 IMAGENET1K_V1 5,288,548 0.39 77.69
resnet18 IMAGENET1K_V1 11,689,512 1.81 69.76
resnet50 IMAGENET1K_V2 25,557,032 4.09 80.86
efficientnet_v2_s IMAGENET1K_V1 21,458,488 8.37 84.23
convnext_tiny IMAGENET1K_V1 28,589,128 4.46 82.52
swin_t IMAGENET1K_V1 28,288,354 4.49 81.47
vit_b_16 IMAGENET1K_V1 86,567,656 17.56 81.07
median CPU latency, batch of 1, 224x224, 4 threads
these numbers are specific to this machine and this build: measure your own
model ms/image GMACs GMACs per ms
mobilenet_v3_small 16.8 0.06 0.003
efficientnet_b0 39.7 0.39 0.010
resnet18 35.7 1.81 0.051
resnet50 89.6 4.09 0.046
efficientnet_v2_s 112.5 8.37 0.074
convnext_tiny 85.9 4.46 0.052
swin_t 172.2 4.49 0.026
vit_b_16 218.7 17.56 0.080
every model wants its own preprocessing, and guessing costs accuracy silently
resnet50 resize 232 -> crop 224, BILINEAR
efficientnet_b0 resize 256 -> crop 224, BICUBIC
convnext_tiny resize 236 -> crop 224, BILINEAR
swin_t resize 232 -> crop 224, BICUBIC
pulling intermediate feature maps out, two ways
torchvision: {'s8': (1, 128, 28, 28), 's32': (1, 512, 7, 7)}
timm channels : [64, 64, 128, 256, 512]
timm strides : [2, 4, 8, 16, 32]
timm shapes : [(1, 64, 112, 112), (1, 64, 56, 56), (1, 128, 28, 28), (1, 256, 14, 14), (1, 512, 7, 7)]
timm data cfg : {'input_size': (3, 224, 224), 'interpolation': 'bicubic', 'mean': (0.485, 0.456, 0.406), 'std': (0.229, 0.224, 0.225), 'crop_pct': 0.95, 'crop_mode': 'center'}
timm version 1.0.26, 1,699 pretrained checkpoints available
matching 'convnext*': 95Latencies were measured on one desktop CPU with four threads. They move between machines, and between runs on the same machine: a repeat run here gave 47 ms for resnet18 rather than 36. The ordering and the ratios held both times. Everything else in this output is metadata and reproduces exactly.
Reading the output, which is the point of the lesson
efficientnet_b0 does one fifth of resnet18's arithmetic and takes longer. 0.39 GMACs against 1.81, and 39.7 ms against 35.7 ms. Depthwise and squeeze-excitation layers move a lot of memory per multiplication. If you had chosen on GMACs alone you would have chosen wrong.
swin_t and convnext_tiny have the same GMACs and differ twofold in time. 4.49 against 4.46 GMACs; 172.2 ms against 85.9 ms. Window partitioning, rolling and masking are data movement that the arithmetic count does not see.
The GMACs per ms column is a hardware-efficiency score. vit_b_16 scores highest at 0.080 and is still the slowest model in the table, because it has so much arithmetic to do. mobilenet_v3_small scores lowest at 0.003 and is the fastest. High efficiency and low latency are different things.
resnet50's DEFAULT weights are IMAGENET1K_V2 at 80.86%. The V1 weights on the same architecture score 76.13%. The number attached to an architecture name is a property of the checkpoint, not the design.
Preprocessing differs between models that look identical. All four crop to 224. They resize to 232, 256, 236 and 232 first, using two different interpolation modes. This is why weights.transforms() exists. Hand-writing Resize(256), CenterCrop(224) for every model quietly costs accuracy.
Both libraries give you the four stage outputs. torchvision's create_feature_extractor takes a dictionary of node names. timm's features_only=True returns all five levels, with strides and channel counts described. Any detection or segmentation neck you attach needs exactly this information.
timm ships 1,699 pretrained checkpoints, 95 of them ConvNeXt variants. torchvision ships one or two per architecture. When you need a specific pretraining recipe, an unusual resolution, or a self-supervised checkpoint, timm is where it will be.
A decision procedure that works
- Write down the latency budget and the hardware. Not "fast". A number, in milliseconds, on a named device.
- Start with
resnet50orconvnext_tinyas your reference point. Get an end-to-end pipeline working and measure the accuracy you actually need. - Move down the size ladder until the budget is met, checking accuracy at each step. Stop at the first model that fits.
- Only then consider exotic options. A self-supervised checkpoint such as DINOv2 is worth trying when labels are scarce. So is domain-specific pretraining, when your images are far from natural photographs.
- Re-measure after export. ONNX, TensorRT and CoreML change the ranking. A model that wins in PyTorch can lose after conversion.
Common mistakes
Choosing on ImageNet top-1. It correlates with transfer accuracy but the correlation is weakest exactly where people need help: small datasets and fine-grained tasks.
Benchmarking with the wrong batch size. Throughput at batch 64 is a data-centre metric. If you serve one request at a time, measure batch 1, which is what the script above does.
Forgetting the warm-up. The first few forward passes include allocation and kernel selection. Discard them, then take a median rather than a mean so one scheduling hiccup does not distort the result.
Fine-tuning everything with 300 images. Freeze the backbone, train the head, and only unfreeze the last stage if the validation curve asks for it. See transfer learning in PyTorch.
Ignoring the licence. Some checkpoints are research-only, and some are gated behind accepting terms. Check before you build a product on one.
Try it yourself
Add "resnet101" and "convnext_small" to CANDIDATES and rerun. Then change torch.set_num_threads(4) to 1 and watch the ranking shift, because the models parallelise differently. That shift is the reason a benchmark from someone else's blog post cannot answer your question.
What to learn next
- Transfer learning in PyTorch — attaching your head and training it properly.
- ONNX — exporting the model you chose, and re-measuring after export.
- Model serving — putting it behind an interface people can call.
Researcher — Mathematics and papers.
What the selection is actually over
Three axes are usually conflated into "which model".
- Architecture. ResNet, ConvNeXt, EfficientNet, Swin, ViT.
- Pretraining data and objective. ImageNet-1k supervised, ImageNet-21k supervised, LAION with a contrastive text objective, LVD-1689M with self-distillation.
- Training recipe. Schedule, augmentation, optimiser, regularisation.
The evidence from the last few years is that axes 2 and 3 frequently dominate axis 1 at fixed compute. ResNet strikes back moves ResNet-50 from 76.1% to 80.4% by changing only axis 3. ConvNeXt's ablation attributes 2.7 of its points to the same axis. torchvision's own V1 and V2 ResNet-50 weights differ by 4.7 points on identical architecture.
A comparison that varies architecture while holding the other two fixed is rare in the literature and is the only kind that supports an architectural conclusion.
Transfer, and where the correlation breaks
Kornblith, Shlens and Le (2019), arxiv.org/abs/1805.08974, report correlations of 0.99 and 0.96 between ImageNet top-1 and transfer accuracy for fixed-feature and fine-tuned settings across 12 datasets. Three caveats they raise are the operative ones for practitioners:
- Fine-grained datasets benefit far less from stronger ImageNet models.
- Regularisation that helps ImageNet, notably label smoothing, can degrade the linear separability of penultimate features.
- With enough target data, training from scratch closes much of the gap, consistent with He, Girshick and Dollár (2018), arxiv.org/abs/1811.08883, who reach 50.9 AP on COCO from random initialisation given a long enough schedule.
The regime where pretraining is decisive is small labelled target sets. That is also the regime where most applied projects live.
FLOPs, parameters and latency are three different quantities
The measurements above make the point empirically. The underlying model is the roofline: a kernel is bound either by arithmetic throughput or by memory bandwidth, and which one depends on arithmetic intensity, operations per byte moved.
- Dense 3x3 convolutions have high intensity and sit near the compute roof. FLOP reductions translate into time reductions.
- Depthwise convolutions have intensity lower by roughly $k^2$ and sit under the memory roof. FLOP reductions translate poorly.
- Attention with explicit reshaping, rolling and masking spends time on data movement that no FLOP counter reports.
Dollár, Singh and Girshick (2021), Fast and Accurate Model Scaling, propose scaling under an activation-count constraint rather than a FLOP constraint, on the empirical grounds that activations predict runtime better. That is the correct instinct: pick the proxy that correlates with your actual cost, and verify the correlation on your hardware.
Domain shift
ImageNet pretraining transfers well to natural photographs and unevenly elsewhere. For medical imaging, Raghu et al. (2019), Transfusion: Understanding Transfer Learning for Medical Imaging, find that large ImageNet architectures give little benefit over far smaller ones on retinal and chest X-ray tasks, and that most of the gain comes from weight scaling rather than learned features.
Practical implications when your domain is far from natural images:
- Prefer a smaller architecture; the capacity of a large ImageNet model is not doing what you think.
- Consider self-supervised pretraining on your own unlabelled data, which is usually abundant in exactly these domains. This is the strongest argument for the previous three lessons in this section.
- Check the input assumptions. Single-channel inputs, 16-bit depth and non-square aspect ratios all break the assumptions baked into a standard preprocessing pipeline.
A defensible default set
| Situation | Reasonable first choice |
|---|---|
| Server GPU, accuracy matters, plenty of labels | ConvNeXt-T or ConvNeXt-S, or a ViT-B with strong pretraining |
| Server CPU, moderate latency budget | ResNet-50 with the V2 weights |
| Mobile or embedded | MobileNetV3, or MobileNetV4 where the runtime supports it |
| Detection or segmentation | A backbone with published stride-4/8/16/32 features and an existing neck integration |
| Few labels, natural images | A frozen DINOv2 checkpoint with a linear head |
| Few labels, unusual domain | Small architecture, plus self-supervised pretraining on your own unlabelled data |
| Zero-shot or text-driven | A CLIP-family model; see CLIP |
Treat these as starting points to be measured against, not as answers.
Reporting a model choice honestly
State the checkpoint identifier, not the architecture name. State the preprocessing. State the hardware, batch size, thread count and precision for every latency number. State the size of the labelled set and whether the backbone was frozen.
Without those, a comparison between two backbones is not reproducible, and most published comparisons between backbones are not.
What to learn next
- Transfer learning in PyTorch — attaching your head and training it properly.
- ONNX — exporting the model you chose, and re-measuring after export.
- Model serving — putting it behind an interface people can call.