Course Content
Transfer Learning and Pretraining
3 sections · 7 lessons
Transfer Learning in Practice — ResNet, EfficientNet, and MobileNet
A team builds a plant-disease classifier for smallholder farmers. Thirty-eight disease classes, photographs taken on cheap Android phones in a field. They reason that a harder problem deserves a bigger model, so they fine-tune ResNet-152 and reach 96.1% top-1 accuracy. Excellent result.
Then they ship it. The exported model is 230 MB — larger than most of the apps on the target handsets. Inference on a mid-range phone takes 1.9 seconds per photo, during which the UI freezes. Within a week the app's average rating is 2.1 stars and the most common review is that it "hangs".
They rebuild it on MobileNetV3-Large. Accuracy drops to 94.8% — 1.3 points worse. The model is now 21 MB and runs in 45 milliseconds. That is a 42× speedup and an 11× size reduction, bought with 1.3 points of accuracy that no farmer would ever notice.
The lesson is not that MobileNet is better than ResNet. It is that "which backbone" is a deployment question disguised as an accuracy question, and getting it wrong is expensive in ways that never show up in a validation score.
The three backbone families and what each is actually for
ResNet: the reliable default
ResNet's contribution was the residual connection. In a plain deep network each layer computes H(x) directly; in a ResNet each block computes F(x)+x, so the layer only has to learn the difference from the identity. If a block has nothing useful to add, driving F(x) towards zero is easy, and the block becomes a pass-through rather than a source of noise.
The gradient consequence is what made 50- and 100-layer networks trainable. Differentiating y=F(x)+x gives ∂x∂y=∂x∂F+1. That +1 means gradients always have a clean path backwards, whatever the multiplicative terms do. Without it, a product of 50 small Jacobians drives the gradient to zero and the early layers never learn.
For transfer learning specifically, ResNet has an underrated property: the block structure is clean and named. model.layer1 through model.layer4 give you four natural groups for progressive unfreezing and discriminative learning rates, with no guesswork about where to cut. Architectures with more irregular structures make that harder.
EfficientNet: get more accuracy per FLOP
Before EfficientNet, scaling a network up meant picking one dimension arbitrarily — more layers, or wider layers, or bigger input images. EfficientNet's observation was that these three must scale together, in a fixed ratio, or you waste capacity.
Compound scaling introduces a single coefficient ϕ and sets depth d=αϕ, width w=βϕ, resolution r=γϕ, subject to α⋅β2⋅γ2≈2. The constraint exists because FLOPs scale linearly with depth but quadratically with both width and resolution, so that product keeps total compute roughly doubling per unit of ϕ.
The payoff is stark. EfficientNet-B0 reaches about 77.1% ImageNet top-1 with 5.3M parameters and 0.39 GFLOPs. ResNet-50 reaches 76.1% with 25.6M parameters and 4.1 GFLOPs. That is 1 point more accuracy for 4.8× fewer parameters and 10.5× fewer FLOPs.
The caveat that catches people: EfficientNet's FLOP efficiency does not translate proportionally into wall-clock speed. Depthwise separable convolutions have low arithmetic intensity — few operations per byte of memory traffic — so GPUs, which are starved for bandwidth rather than arithmetic, run them far below peak. EfficientNet-B0 has 10.5× fewer FLOPs than ResNet-50 but is typically only 2–3× faster on a GPU. EfficientNetV2 was designed specifically to fix this, replacing depthwise convolutions with ordinary ones in the early stages.
MobileNet: built for the phone in someone's pocket
MobileNet's core move is the depthwise separable convolution, which splits a standard convolution into two cheaper steps: a depthwise convolution that filters each input channel independently, then a 1×1 pointwise convolution that mixes channels.
The arithmetic is worth doing once. A standard 3×3 convolution from 128 input channels to 256 output channels over a 28×28 feature map costs:
The separable version costs the depthwise part plus the pointwise part:
That is 8.7× cheaper. The general ratio is Cout1+k21, which for a 3×3 kernel and any reasonably wide layer approaches 1/9.
MobileNetV3 adds two more things worth knowing about. Squeeze-and-excitation blocks let each channel be reweighted according to global context, at almost no cost. And the hard-swish activation, x⋅6ReLU6(x+3), approximates swish using only operations that quantise cleanly to 8-bit integers — which matters enormously when the model runs on a phone's NPU.
Side by side
| Model | Params | GFLOPs | ImageNet top-1 | Feature dim | Typical GPU fine-tune | Best for |
|---|---|---|---|---|---|---|
| ResNet-18 | 11.7M | 1.8 | 69.8% | 512 | Fast | Prototyping, tiny datasets |
| ResNet-50 | 25.6M | 4.1 | 76.1% | 2048 | Moderate | The default. Start here. |
| ResNet-101 | 44.5M | 7.8 | 77.4% | 2048 | Slow | Large datasets, accuracy-critical |
| EfficientNet-B0 | 5.3M | 0.39 | 77.1% | 1280 | Moderate | Accuracy per parameter |
| EfficientNet-B3 | 12.2M | 1.8 | 81.6% | 1536 | Slow | Best accuracy/size trade-off |
| EfficientNetV2-S | 21.5M | 8.4 | 83.9% | 1280 | Moderate | Modern default when GPU speed matters |
| MobileNetV3-Small | 2.5M | 0.06 | 67.4% | 576 | Very fast | Extreme edge constraints |
| MobileNetV3-Large | 5.4M | 0.22 | 75.2% | 960 | Very fast | Mobile and embedded deployment |
| ViT-B/16 | 86M | 17.6 | 81.1% | 768 | Very slow | Large datasets, strong pretraining |
Pick the backbone from the deployment constraint first and the accuracy target second. A model that cannot run where it must run has an accuracy of zero.
A decision shortcut that holds up in practice: if inference happens on a server, start with ResNet-50 and move to EfficientNetV2-S if you need more. If inference happens on a phone or an embedded board, start with MobileNetV3-Large. If you have fewer than a thousand images per class, prefer the smaller model in either family — you do not have the data to justify the larger one.
The table is not the whole menu in 2026. ConvNeXt (in both torchvision and timm) is a modern convolutional network that is a common default where ResNet-50 used to be, and MobileNetV4 (in timm) is the newer mobile family. For a frozen feature extractor, backbones pretrained without labels or on image–text pairs — DINOv2 and DINOv3, SigLIP 2 — usually beat ImageNet-supervised ones, and timm ships all of them. The workflow in this lesson is the same whichever you pick.
Loading pretrained weights the current way
Why pretrained=True is on its way out
Older tutorials write models.resnet50(pretrained=True). That API has been deprecated since torchvision 0.13 — it still works but prints a warning — and the replacement exists for two reasons that are worth understanding rather than working around.
First, it is ambiguous. There is no longer one set of ImageNet weights per architecture. torchvision ships IMAGENET1K_V1 for ResNet-50 at 76.13% top-1, and IMAGENET1K_V2 — the same architecture retrained with a modern recipe — at 80.86%. A boolean flag cannot express which you want, and for full fine-tuning, silently getting the weaker one can cost you points for free. The reverse can also happen: the V2 recipe's heavy regularisation (label smoothing, mixup, cutmix) makes features that transfer worse to a frozen linear head — Kornblith et al. (2019) found this effect for label smoothing, and a quick linear probe on Flowers-102 gives about 85% with V1 features against under 80% with V2. If you are only training a head, try both.
Second, it discards the preprocessing. Every set of weights was trained with specific resize dimensions, crop size, interpolation mode and normalisation constants. Get those wrong and you feed the network inputs from a distribution it has never seen. The weights enum carries that metadata with it.
1import torch, torch.nn as nn2from torchvision import models3from torchvision.models import ResNet50_Weights45weights = ResNet50_Weights.IMAGENET1K_V26model = models.resnet50(weights=weights)78# The preprocessing that these exact weights were trained with.9preprocess = weights.transforms()10print(preprocess)11# ImageClassification(12# crop_size=[224], resize_size=[232],13# mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225],14# interpolation=InterpolationMode.BILINEAR15# )1617print(weights.meta["num_params"]) # 2555703218print(weights.meta["categories"][:3]) # ['tench', 'goldfish', 'great white shark']Note resize_size=232, not 256. The V2 recipe used a different resize, and a mismatch here costs roughly half a point on its own. Never hard-code preprocessing constants when the weights object will tell you.
To start from random weights instead — which you want for a from-scratch control run — pass weights=None.
Replacing the head, per architecture
The final layer lives in a different place in each family, and this is a routine source of confusion.
1num_classes = 3823# ResNet: a single Linear called `fc`4resnet = models.resnet50(weights=ResNet50_Weights.IMAGENET1K_V2)5resnet.fc = nn.Linear(resnet.fc.in_features, num_classes) # 2048 -> 3867# EfficientNet: classifier is Sequential(Dropout, Linear)8effnet = models.efficientnet_b0(weights="IMAGENET1K_V1")9effnet.classifier[1] = nn.Linear(effnet.classifier[1].in_features, num_classes)1011# MobileNetV3: classifier is Sequential(Linear, Hardswish, Dropout, Linear)12mnet = models.mobilenet_v3_large(weights="IMAGENET1K_V2")13mnet.classifier[3] = nn.Linear(mnet.classifier[3].in_features, num_classes)Always read in_features off the existing layer rather than typing the number. Hard-coding 2048 works until someone switches the backbone to ResNet-18, whose feature dimension is 512, and then you get a shape error at best or — if the dimensions happen to line up — a silently worse model.
Using timm for everything else
PyTorch Image Models (timm) carries well over a thousand pretrained backbones, including many that torchvision does not ship, and gives them a uniform interface.
1import timm23# num_classes= builds the correct head for you, whatever the architecture.4model = timm.create_model("efficientnet_b3", pretrained=True, num_classes=38)56# Feature extraction mode: no head at all, returns pooled features.7backbone = timm.create_model("resnet50", pretrained=True, num_classes=0)8feat_dim = backbone.num_features # 2048910# The right preprocessing for these specific weights.11cfg = timm.data.resolve_data_config({}, model=model)12transform = timm.data.create_transform(**cfg, is_training=False)1314# Multi-scale features, for detection or segmentation heads.15pyramid = timm.create_model("resnet50", pretrained=True,16 features_only=True, out_indices=(1, 2, 3, 4))17print(pyramid.feature_info.channels()) # [256, 512, 1024, 2048]num_classes=0 is the cleanest way to get a headless backbone — no surgery, no nn.Identity() hacks, and num_features tells you the output dimension so you never hard-code it.
Inspecting what you loaded
Two checks before you train anything, both of which have caught real bugs.
1total = sum(p.numel() for p in model.parameters())2trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)3print(f"{trainable:,} / {total:,} trainable "4 f"({100 * trainable / total:.2f}%)")56with torch.no_grad():7 out = model(torch.randn(2, 3, 224, 224))8print(out.shape) # must be torch.Size([2, 38])If you intended feature extraction and the trainable percentage reads 100%, your freeze loop ran before the head replacement or not at all. If the output shape's second dimension is 1000, you forgot to replace the head entirely and are about to train a plant classifier to predict ImageNet categories.
Transformer and vision-language backbones
Convolutional backbones are not the only option, and two alternatives are worth knowing about.
Vision Transformers split an image into fixed patches — 16×16 pixels for ViT-B/16, giving 196 patches from a 224×224 image — embed each patch as a token, and run self-attention over them. Because attention is global from the first layer, a ViT has no built-in assumption that nearby pixels are related. That assumption, which convolutions get for free, is genuinely useful when data is scarce, which is why ViTs trained from scratch on ImageNet-1k (1.3 million images) underperform comparable CNNs, and overtake them only when pretrained on tens to hundreds of millions of images. For transfer learning the practical rule is: use a ViT when its pretraining was large (ImageNet-21k or bigger), and freeze more of it than you would freeze of a ResNet.
CLIP was trained on hundreds of millions of image–caption pairs to place matching images and texts near each other in a shared embedding space. Two consequences follow. Its image encoder is an unusually strong frozen feature extractor, often beating an ImageNet-pretrained backbone by several points under linear probing, because caption supervision covers a far broader concept space than 1,000 ImageNet labels. And it classifies with no training at all: embed the text "a photo of a leaf with powdery mildew" for each class, embed the image, and pick the nearest. Zero-shot accuracy will not match a fine-tuned model, but it gives you a working baseline on day one with zero labelled examples — which is exactly what you want while annotation is still in progress.
The feature extraction pipeline, end to end
1import torch, torch.nn as nn2from torch.utils.data import DataLoader3from torchvision import datasets, transforms, models4from torchvision.models import ResNet50_Weights56device = "cuda" if torch.cuda.is_available() else "cpu"7weights = ResNet50_Weights.IMAGENET1K_V28norm = transforms.Normalize(mean=[0.485, 0.456, 0.406],9 std=[0.229, 0.224, 0.225])1011train_tf = transforms.Compose([12 transforms.RandomResizedCrop(224, scale=(0.7, 1.0)),13 transforms.RandomHorizontalFlip(),14 transforms.ToTensor(), norm,15])16eval_tf = transforms.Compose([17 transforms.Resize(232), transforms.CenterCrop(224),18 transforms.ToTensor(), norm,19])2021train_ds = datasets.ImageFolder("data/train", train_tf)22val_ds = datasets.ImageFolder("data/val", eval_tf)23train_dl = DataLoader(train_ds, batch_size=32, shuffle=True,24 num_workers=4, pin_memory=True)25val_dl = DataLoader(val_ds, batch_size=64, num_workers=4)2627model = models.resnet50(weights=weights)28for p in model.parameters(): # freeze FIRST29 p.requires_grad = False30model.fc = nn.Linear(model.fc.in_features, len(train_ds.classes)) # then replace31model = model.to(device)Note that the evaluation transform uses Resize(232) then CenterCrop(224), matching the V2 recipe exactly, while the training transform uses RandomResizedCrop — random cropping is itself the augmentation, so no separate resize is needed.
1def run_epoch(model, loader, criterion, optimiser=None):2 train = optimiser is not None3 model.train(train)4 total_loss, correct, n = 0.0, 0, 05 for x, y in loader:6 x, y = x.to(device, non_blocking=True), y.to(device)7 with torch.set_grad_enabled(train):8 out = model(x)9 loss = criterion(out, y)10 if train:11 optimiser.zero_grad(set_to_none=True)12 loss.backward()13 optimiser.step()14 total_loss += loss.item() * y.size(0)15 correct += (out.argmax(1) == y).sum().item()16 n += y.size(0)17 return total_loss / n, correct / n1819criterion = nn.CrossEntropyLoss(label_smoothing=0.1)20optimiser = torch.optim.AdamW(21 [p for p in model.parameters() if p.requires_grad], lr=1e-3, weight_decay=1e-4)22scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimiser, T_max=15)2324best = 0.025for epoch in range(15):26 tr_loss, tr_acc = run_epoch(model, train_dl, criterion, optimiser)27 va_loss, va_acc = run_epoch(model, val_dl, criterion)28 scheduler.step()29 print(f"epoch {epoch:2d} train {tr_acc:.3f} val {va_acc:.3f} "30 f"gap {tr_acc - va_acc:+.3f}")31 if va_acc > best:32 best = va_acc33 torch.save(model.state_dict(), "feature_extraction.pt")Printing the train–validation gap every epoch, not just the accuracies, is the single most useful habit in this loop. A gap under 5 points means you have headroom to unfreeze; a gap over 15 means you are memorising and should regularise instead.
Fine-tuning: unfreeze in stages, then use discriminative rates
Take the frozen-baseline checkpoint and go further. The rule that prevents disaster: the head must be trained before any backbone weight is allowed to move. A random head produces large, meaningless gradients, and at a from-scratch learning rate those gradients destroy pretrained features within the first fifty steps.
1model.load_state_dict(torch.load("feature_extraction.pt"))23def set_trainable(model, groups):4 for p in model.parameters():5 p.requires_grad = False6 for g in groups:7 for p in g.parameters():8 p.requires_grad = True910# Phase 2: add the last residual block, 10x lower learning rate.11set_trainable(model, [model.fc, model.layer4])12opt = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad],13 lr=1e-4, weight_decay=1e-4)14for epoch in range(8):15 run_epoch(model, train_dl, criterion, opt)1617# Phase 3: everything trainable, with per-depth learning rates.18for p in model.parameters():19 p.requires_grad = True2021base_lr, factor = 1e-4, 2.622groups = [model.fc, model.layer4, model.layer3, model.layer2, model.layer1]23opt = torch.optim.AdamW(24 [{"params": g.parameters(), "lr": base_lr / factor ** d}25 for d, g in enumerate(groups)],26 weight_decay=1e-4)Those rates work out as 1.00e-4 for the head, 3.85e-5 for layer4, 1.48e-5 for layer3, 5.69e-6 for layer2 and 2.19e-6 for layer1. The bottom of the network moves 46× more slowly than the top — edge detectors get nudged, the classifier gets rebuilt.
Validate between every phase. If phase 3 does not beat phase 2, keep the phase 2 checkpoint. Your dataset has told you it cannot support that many free parameters, and the correct response is to spend effort on data rather than on the optimiser.
Augmentation sized to your dataset
Augmentation strength should scale inversely with dataset size. Heavy augmentation on a large dataset wastes capacity fighting distortions that add nothing; light augmentation on a tiny dataset leaves the model free to memorise.
| Images per class | Geometric | Photometric | Advanced | Typical dropout |
|---|---|---|---|---|
| Under 50 | RandomResizedCrop(0.5–1.0), flip, rotate ±20° | Strong colour jitter (0.4) | Mixup + CutMix + RandomErasing | 0.5 |
| 50–500 | RandomResizedCrop(0.7–1.0), flip, rotate ±15° | Moderate jitter (0.2) | RandAugment, light Mixup | 0.3–0.4 |
| 500–5,000 | RandomResizedCrop(0.8–1.0), flip | Light jitter (0.1) | RandomErasing | 0.2 |
| Over 5,000 | RandomResizedCrop(0.8–1.0), flip | None needed | Optional | 0.1 |
1from torchvision.transforms import RandAugment, RandomErasing23small_dataset_tf = transforms.Compose([4 transforms.RandomResizedCrop(224, scale=(0.5, 1.0)),5 transforms.RandomHorizontalFlip(),6 transforms.RandomRotation(20),7 RandAugment(num_ops=2, magnitude=9),8 transforms.ToTensor(), norm,9 RandomErasing(p=0.25, scale=(0.02, 0.2)),10])The failure mode here is label-destroying augmentation, and it is common. Horizontal flips are free on plant leaves and wrong on chest X-rays, where mirroring relocates the heart. Rotation beyond about 20° is fine on satellite imagery and harmful on handwritten digits, where a rotated 6 becomes a 9. Aggressive colour jitter destroys the signal in histopathology, where stain colour is diagnostic. Before enabling any transform, ask whether a human expert would still assign the same label to the transformed image. If not, you are training the model to be wrong.
Every augmentation is a claim that the label is invariant to that change. If the claim is false, the augmentation is not regularisation — it is label noise you injected deliberately.
What this means when you build something
The plant-disease team's real mistake was ordering the decisions wrongly. They chose accuracy first and discovered the deployment constraint after shipping. Reverse that order and the whole project gets easier.
- Write down the hard constraints before you open an editor. Maximum model size, maximum latency, target hardware, whether inference is batched. These eliminate most of the backbone table immediately. A 42× latency difference is not a tuning detail you fix later.
- Establish the frozen baseline on the constrained backbone. Not on the biggest one. If MobileNetV3-Large is what you can ship, that is what you measure, because a ResNet-152 number you cannot deploy tells you nothing useful.
- Let the weights object supply the preprocessing. Call
weights.transforms()or timm'sresolve_data_configrather than typing constants copied from a blog post. A mismatched resize or the wrong normalisation costs real accuracy and produces no error message. - Unfreeze in stages, validating between each. Head only, then the last block at a 10× lower rate, then discriminative rates across everything. Stop at the first stage that fails to improve.
- Match augmentation strength to dataset size, and check every transform preserves the label. This is where domain knowledge earns its keep, and where a generic recipe copied from a cats-and-dogs tutorial will quietly cost you points.
- Keep a from-scratch control if your domain is unusual. Medical, scientific and sensor data all deserve the check. If random initialisation matches pretrained initialisation, ImageNet features are not helping and you should look for a better-matched source.
The team eventually shipped MobileNetV3-Large at 94.8%, then recovered a point by distilling from their ResNet-152 into it — the large model as a teacher rather than as the product. That is the shape of a mature solution: the big model earns its keep in training, and the small model is what the user's phone actually runs.