PyTorch 迁移学习实战:ResNet18 实现 20 类食物图像分类(完整可运行代码)

目录

一、项目前言

环境依赖

二、完整源码

三、代码分模块深度解析

[3.1 迁移学习核心:冻结主干网络](#3.1 迁移学习核心:冻结主干网络)

两种训练模式切换

[3.2 答疑:model = resnet_model.to(device) 为什么不用加括号?](#3.2 答疑:model = resnet_model.to(device) 为什么不用加括号?)

[3.3 数据增强与归一化说明](#3.3 数据增强与归一化说明)

[3.4 自定义 Dataset 数据集](#3.4 自定义 Dataset 数据集)

[3.5 训练 / 测试流程关键点](#3.5 训练 / 测试流程关键点)

四、数据集文件配置说明

五、拓展作业:单张图片推理预测(输入图片输出分类结果)

六、常见问题

七、总结


一、项目前言

传统从零搭建 CNN 训练图像分类,需要海量数据、长时间迭代,收敛速度慢。迁移学习可以直接复用 ImageNet 预训练好的 ResNet 残差网络,仅微调最后一层全连接层即可适配自定义数据集,大幅降低训练成本、提升精度。

本文基于ResNet18搭建 20 分类食物识别模型,完整包含:数据集自定义、数据增强、模型冻结、优化器 + 学习率衰减、训练 / 测试循环、最优精度保存逻辑,附带两种训练模式(冻结主干 / 全量训练),适合深度学习入门学习迁移学习。

环境依赖

bash

运行

复制代码
pip install torch torchvision pillow numpy

二、完整源码

python

运行

python 复制代码
import torch
import torchvision.models as models
from torch import nn
from torch.utils.data import Dataset, DataLoader
from torchvision import transforms
from PIL import Image
import numpy as np

# ====================== 1. 加载预训练ResNet18并冻结主干 ======================
# 加载ImageNet预训练权重的ResNet18
resnet_model = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)

# 冻结主干网络所有参数,不更新卷积层权重
for param in resnet_model.parameters():
    param.requires_grad = False

# 获取原模型最后一层全连接层输入特征维度
in_features = resnet_model.fc.in_features
# 替换全连接层:输出改为20,适配20类食物分类
resnet_model.fc = nn.Linear(in_features, 20)

# 收集仅需要更新的参数(只有最后一层全连接层)
params_to_update = []
for param in resnet_model.parameters():
    if param.requires_grad == True:
        params_to_update.append(param)

# ====================== 2. 数据增强与预处理 ======================
data_transforms = {
    'trainda':
        transforms.Compose([
            transforms.Resize([300, 300]),
            transforms.RandomRotation(45),        # 随机旋转-45~45°
            transforms.CenterCrop(224),           # 中心裁剪224×224(ResNet标准输入尺寸)
            transforms.RandomHorizontalFlip(p=0.5),# 随机水平翻转
            transforms.RandomVerticalFlip(p=0.5),  # 随机垂直翻转
            transforms.RandomGrayscale(p=0.1),     # 小概率转灰度图
            transforms.ToTensor(),
            # ImageNet标准归一化均值、方差
            transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
        ]),
    'valid':
        transforms.Compose([
            transforms.Resize([224, 224]),
            transforms.ToTensor(),
            transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
        ]),
}

# ====================== 3. 自定义数据集Dataset ======================
class food_dataset(Dataset):
    def __init__(self, file_path, transform=None):
        self.file_path = file_path
        self.imgs = []
        self.labels = []
        self.transform = transform
        # 读取txt标注文件:每行格式 图片路径 类别标签
        with open(self.file_path, 'r', encoding='utf-8') as f:
            samples = [x.strip().split(' ') for x in f.readlines()]
            for img_path, label in samples:
                self.imgs.append(img_path)
                self.labels.append(label)

    # 返回数据集总样本数量
    def __len__(self):
        return len(self.imgs)

    # 根据索引读取单张图片+标签
    def __getitem__(self, idx):
        image = Image.open(self.imgs[idx]).convert("RGB")
        # 执行数据增强/归一化
        if self.transform:
            image = self.transform(image)
        # 标签转int64张量,适配CrossEntropyLoss
        label = self.labels[idx]
        label = torch.from_numpy(np.array(label, dtype=np.int64))
        return image, label

# ====================== 4. 构建DataLoader数据加载器 ======================
training_data = food_dataset(file_path='./train.txt', transform=data_transforms['trainda'])
test_data = food_dataset(file_path='./test.txt', transform=data_transforms['valid'])

train_dataloader = DataLoader(training_data, batch_size=64, shuffle=True)
test_dataloader = DataLoader(test_data, batch_size=64, shuffle=True)

# ====================== 5. 设备自动适配(GPU/CUDA/MPS/CPU) ======================
device = "cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu"
print(f"Using {device} device")

# 模型移至GPU/CPU,无需括号原因下文详解
model = resnet_model.to(device)

# ====================== 6. 损失函数、优化器、学习率衰减 ======================
loss_fn = nn.CrossEntropyLoss()  # 多分类标准损失函数
# 仅更新解冻的全连接层参数
optimizer = torch.optim.Adam(params_to_update, lr=0.001)
# 每5轮epoch学习率×0.5,逐步降低学习率
scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.5)

# ====================== 7. 训练一轮函数 ======================
def train(dataloader, model, loss_fn, optimizer):
    model.train()  # 开启训练模式(启用dropout/bn更新)
    for X, y in dataloader:
        X, y = X.to(device), y.to(device)
        pred = model(X)  # 等价model.forward(X),推荐简写写法
        loss = loss_fn(pred, y)

        # 标准反向传播四步
        optimizer.zero_grad()  # 清空历史梯度
        loss.backward()        # 反向传播求梯度
        optimizer.step()       # 根据梯度更新权重

# ====================== 8. 测试/验证函数 ======================
best_acc = 0
acc_s = []   # 保存每轮精度
loss_s = []  # 保存每轮损失
def test(dataloader, model, loss_fn):
    global best_acc
    size = len(dataloader.dataset)
    num_batches = len(dataloader)
    model.eval()  # 评估模式,关闭dropout、冻结BN层
    test_loss, correct = 0, 0
    # 关闭梯度计算,节省显存/内存
    with torch.no_grad():
        for X, y in dataloader:
            X, y = X.to(device), y.to(device)
            pred = model(X)
            test_loss += loss_fn(pred, y).item()
            # argmax(1)取每行最大概率索引,即为预测类别
            correct += (pred.argmax(1) == y).type(torch.float).sum().item()
    test_loss /= num_batches
    correct /= size
    print(f"Test result: \n Accuracy: {(100*correct):.2f}%, Avg loss: {test_loss:.4f}")
    acc_s.append(correct)
    loss_s.append(test_loss)
    # 记录最优精度
    if correct > best_acc:
        best_acc = correct

# ====================== 9. 完整训练循环 ======================
epochs = 100
for t in range(epochs):
    print(f"Epoch {t+1}\n-------------------------------")
    train(train_dataloader, model, loss_fn, optimizer)
    scheduler.step()  # 每轮更新学习率
    test(test_dataloader, model, loss_fn)
print('最优训练准确率:', f"{best_acc*100:.2f}%")

三、代码分模块深度解析

3.1 迁移学习核心:冻结主干网络

python

运行

python 复制代码
resnet_model = models.resnet18(weights=models.ResNet18_Weights.DEFAULT)
# 冻结所有卷积层参数
for param in resnet_model.parameters():
    param.requires_grad = False
# 替换最后一层全连接层,适配20分类
in_features = resnet_model.fc.in_features
resnet_model.fc = nn.Linear(in_features, 20)
  1. weights=models.ResNet18_Weights.DEFAULT:加载 ImageNet 百万图像预训练权重,网络已经学会通用边缘、纹理、色彩特征;
  2. param.requires_grad = False:冻结参数,反向传播时不会更新卷积层权重,只训练最后自定义全连接层;
  3. ResNet18 默认输出 1000 类,替换fc层将输出改为 20,适配食物 20 分类任务。
两种训练模式切换
  1. 模式 1(代码默认):冻结主干,仅微调全连接层 适合数据集较小、硬件算力不足,训练快、不易过拟合;

  2. 模式 2:解冻全部参数,全量微调 注释冻结循环代码,优化器改为读取全部参数:

    python

    运行

    python 复制代码
    # 注释冻结代码
    # for param in resnet_model.parameters():
    #     param.requires_grad = False
    # 优化器传入全部参数
    optimizer = torch.optim.Adam(resnet_model.parameters(), lr=0.001)

    适合数据集量大、算力充足,整体精度上限更高。

3.2 答疑:model = resnet_model.to(device) 为什么不用加括号?

新手自定义 CNN 网络时写法:model = CNN().to(device)

  • CNN()实例化网络 ,创建新对象; 本文代码:resnet_model 已经提前实例化完成,不需要再次调用构造函数,直接调用.to(device)迁移设备即可。

python

运行

python 复制代码
# 分步拆解
# 1. 实例化预训练模型(已完成)
resnet_model = models.resnet18(...)
# 2. 直接迁移至GPU,无需再次实例化
model = resnet_model.to(device)

3.3 数据增强与归一化说明

训练集使用大量随机变换扩充样本,防止过拟合;验证集仅做基础缩放,不添加随机操作:

  1. 旋转、翻转、灰度化:模拟真实场景拍摄角度、光线变化;
  2. 224×224:ResNet 网络固定输入尺寸;
  3. 归一化均值方差是 ImageNet 数据集标准,预训练权重基于该分布训练,必须统一。

3.4 自定义 Dataset 数据集

读取train.txt/test.txt标注文件,文件格式要求:

plaintext

复制代码
./data/img001.jpg 0
./data/img002.jpg 1
./data/img003.jpg 2
...

每行用空格分割:图片相对路径 类别数字标签

  • __len__:返回样本总数,len(数据集)可调用;
  • __getitem__:索引取单张图片与标签,自动执行图像预处理。

3.5 训练 / 测试流程关键点

  1. model.train():训练模式,Dropout、BatchNorm 启用更新;
  2. model.eval():验证模式,关闭随机层,固定归一化参数;
  3. with torch.no_grad():验证阶段关闭梯度计算,大幅节省显存;
  4. StepLR学习率衰减:每 5 轮学习率减半,后期收敛更稳定;
  5. CrossEntropyLoss:多分类专用损失,标签无需 one-hot 编码,直接输入数字标签。

四、数据集文件配置说明

  1. 新建train.txttest.txt放在代码同级目录;
  2. 文本每行格式:图片路径 类别编号,类别从 0 开始依次递增;
  3. 图片路径支持相对路径,确保路径无中文、无空格。

五、拓展作业:单张图片推理预测(输入图片输出分类结果)

在代码末尾追加推理函数,实现单图输入输出类别:

python

运行

python 复制代码
def predict_one_img(img_path, model, transform, device):
    model.eval()
    img = Image.open(img_path).convert("RGB")
    img = transform(img).unsqueeze(0)  # 增加batch维度 [1,3,224,224]
    img = img.to(device)
    with torch.no_grad():
        pred = model(img)
        pred_cls = pred.argmax(1).item()
    return pred_cls

# 测试推理
test_transform = transforms.Compose([
    transforms.Resize([224,224]),
    transforms.ToTensor(),
    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])
result = predict_one_img("./test_food.jpg", model, test_transform, device)
print(f"图片预测类别:{result}")

六、常见问题

  1. CUDA out of memory 显存溢出 调小batch_size=64改为 16/32,或使用 CPU 运行;
  2. test.txt 读取报错 检查 txt 每行分隔符是空格,末尾无空行,图片路径存在;
  3. 精度持续很低 确认归一化参数正确、训练集数据增强正常,可切换全量微调模式;
  4. MPS 设备报错(Mac) PyTorch 版本更新至 2.0 以上,MPS 仅支持新版 torch。

七、总结

  1. 迁移学习核心逻辑:复用预训练卷积特征提取器,仅替换输出层适配自定义分类任务;
  2. 两种训练方案按需选择:小数据集冻结主干,大数据集全量微调;
  3. 完整工程化流程:自定义数据集→数据增强→模型构建→训练循环→验证评估;
  4. 代码可直接拓展:增加模型保存、绘制 loss/acc 曲线、单图推理功能。
相关推荐
宿州派大星1 小时前
[NLP实战] 基于PyTorch实现N-gram词嵌入模型:输入4个词预测第5个词
人工智能·pytorch·深度学习·nlp
Liaiyang664 小时前
空圈容错视角下的无人机全链路审计:从理论框架到耦合式检验
人工智能·pytorch·python·深度学习·系统架构·自动驾驶·无人机
AI模力圈8 小时前
Pytorch图模式技术原理解析
pytorch·深度学习·torch.compile
Tancenter9 小时前
gather和scatter API
pytorch·tensor
TLA技术11 小时前
LogMiner vs 裸日志解析(三):Oracle日志解析中的“前镜像”和“后镜像”,到底怎么用?
数据库·oracle·flink·dba·迁移学习
磁场转动100万匹12 小时前
基于 dlib 与 OpenCV 的疲劳驾驶检测:眼睛纵横比(EAR)原理与代码逐段解析
pytorch·python
Dr_Fourier13 小时前
AWQ量化
c++·人工智能·pytorch·ai
就叫你天选之人啦1 天前
安装torch+vllm+flash_attn的prompt
人工智能·pytorch·python
估值探索者2 天前
【Python量化系统工程化 #08】关了 SSH 就停?systemd 让脚本开机自启 + 异常自动拉起
java·c++·人工智能·分类·数据挖掘
Thomas.Sir2 天前
第16课:PyTorch|循环神经网络RNN与序列数据处理【让模型拥有“记忆”】
人工智能·pytorch·rnn