C.W.K.
Stream
Lesson 02 of 06 · published

Vision Models — ResNet, EfficientNet, ConvNeXt, ViT

~12 min · resnet, efficientnet, convnext, vit

Level 0Tensor Curious
0 XP0/62 lessons0/13 achievements
0/120 XP to next level120 XP to go0% complete

The pretrained vision model menu

torchvision ships pretrained weights for dozens of architectures. The four families you'll reach for most:

  • ResNet (resnet18/34/50/101/152) — the workhorse. Reliable, well-understood, fast on every backend. Default starting point for any new task.
  • EfficientNet (efficientnet_b0..b7, _v2_*) — better accuracy/parameter trade-off. Use when model size matters (mobile deployment, large-scale inference).
  • ConvNeXt (convnext_tiny/small/base/large) — modern ConvNet that closed the gap to ViTs. Surprisingly strong, faster than ViT for the same accuracy.
  • Vision Transformer (vit_b_16/32, vit_l_16, vit_h_14) — Transformer applied to images. State-of-the-art for many tasks but more data-hungry than CNNs and slightly slower per parameter.

Adapting them — head replacement is architecture-specific

Each architecture exposes its classifier head differently:

  • ResNet: model.fc = nn.Linear(model.fc.in_features, num_classes)
  • EfficientNet: model.classifier[1] = nn.Linear(model.classifier[1].in_features, num_classes)
  • ConvNeXt: model.classifier[2] = nn.Linear(model.classifier[2].in_features, num_classes)
  • ViT: model.heads.head = nn.Linear(model.heads.head.in_features, num_classes)

The pattern: print(model) first, find the final Linear, replace it. The convention is consistent enough that you can write a small helper function once and reuse it.

Code

Loading the four families·python
from torchvision import models
from torchvision.models import (
    ResNet50_Weights, EfficientNet_B0_Weights,
    ConvNeXt_Base_Weights, ViT_B_16_Weights,
)

resnet = models.resnet50(weights=ResNet50_Weights.IMAGENET1K_V2)
effnet = models.efficientnet_b0(weights=EfficientNet_B0_Weights.IMAGENET1K_V1)
convnext = models.convnext_base(weights=ConvNeXt_Base_Weights.IMAGENET1K_V1)
vit = models.vit_b_16(weights=ViT_B_16_Weights.IMAGENET1K_V1)

for m in [resnet, effnet, convnext, vit]:
    n = sum(p.numel() for p in m.parameters())
    print(f"{type(m).__name__:14s}  {n/1e6:6.1f}M params")
Architecture-aware head replacement·python
import torch.nn as nn
from torchvision import models
from torchvision.models import (
    ResNet50_Weights, EfficientNet_B0_Weights,
    ConvNeXt_Base_Weights, ViT_B_16_Weights,
)

def make_classifier(arch, num_classes):
    if arch == 'resnet50':
        m = models.resnet50(weights=ResNet50_Weights.IMAGENET1K_V2)
        m.fc = nn.Linear(m.fc.in_features, num_classes)
    elif arch == 'efficientnet_b0':
        m = models.efficientnet_b0(weights=EfficientNet_B0_Weights.IMAGENET1K_V1)
        m.classifier[1] = nn.Linear(m.classifier[1].in_features, num_classes)
    elif arch == 'convnext_base':
        m = models.convnext_base(weights=ConvNeXt_Base_Weights.IMAGENET1K_V1)
        m.classifier[2] = nn.Linear(m.classifier[2].in_features, num_classes)
    elif arch == 'vit_b_16':
        m = models.vit_b_16(weights=ViT_B_16_Weights.IMAGENET1K_V1)
        m.heads.head = nn.Linear(m.heads.head.in_features, num_classes)
    return m

model = make_classifier('convnext_base', num_classes=5)
timm — when torchvision doesn't have the architecture you need·python
# pip install timm
import timm

# timm has 1000+ pretrained vision models — ConvNeXt v2, MaxViT, EVA, BEiT, etc.
model = timm.create_model('convnext_base.fb_in22k_ft_in1k', pretrained=True, num_classes=5)

# timm also gives you the matching transforms
data_cfg = timm.data.resolve_data_config({}, model=model)
transform = timm.data.create_transform(**data_cfg)

# For research, timm is often more up-to-date than torchvision
# For production stability, torchvision is the safer pick.

External links

Exercise

Use the make_classifier function to instantiate all four architectures with num_classes=5. Time a single forward pass on a (8, 3, 224, 224) batch on your fastest device. Note the per-architecture latency — useful for picking what to deploy.

Progress

Progress is local-only — sign in to sync across devices.
Spotted a bug or have feedback on this page?Report an Issue

Comments 0

🔔 Reply notifications (sign in)
Sign inPlease sign in to comment.

No comments yet — be the first.