Code
pytorch
import torch
import torch.nn as nn
from torchvision import models, transforms
from torch.utils.data import DataLoader
# Load pretrained ResNet
model = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)
# Freeze backbone parameters
for param in model.parameters():
param.requires_grad = False
# Replace the classifier head
num_features = model.fc.in_features
model.fc = nn.Linear(num_features, 10)
# Train only the head
optimizer = torch.optim.Adam(model.fc.parameters(), lr=1e-3)
criterion = nn.CrossEntropyLoss()
# Data augmentation for images
transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406],
std=[0.229, 0.224, 0.225]),
])
# Later, unfreeze and fine-tune at a lower LR
# for param in model.parameters():
# param.requires_grad = True
# optimizer = torch.optim.Adam(model.parameters(), lr=1e-5)