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

📅 2026/7/22 6:01:00
PyTorch 迁移学习实战:ResNet18 实现 20 类食物图像分类(完整可运行代码)
目录一、项目前言环境依赖二、完整源码三、代码分模块深度解析3.1 迁移学习核心冻结主干网络两种训练模式切换3.2 答疑model resnet_model.to(device) 为什么不用加括号3.3 数据增强与归一化说明3.4 自定义 Dataset 数据集3.5 训练 / 测试流程关键点四、数据集文件配置说明五、拓展作业单张图片推理预测输入图片输出分类结果六、常见问题七、总结一、项目前言传统从零搭建 CNN 训练图像分类需要海量数据、长时间迭代收敛速度慢。迁移学习可以直接复用 ImageNet 预训练好的 ResNet 残差网络仅微调最后一层全连接层即可适配自定义数据集大幅降低训练成本、提升精度。本文基于ResNet18搭建 20 分类食物识别模型完整包含数据集自定义、数据增强、模型冻结、优化器 学习率衰减、训练 / 测试循环、最优精度保存逻辑附带两种训练模式冻结主干 / 全量训练适合深度学习入门学习迁移学习。环境依赖bash运行pip install torch torchvision pillow numpy二、完整源码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(weightsmodels.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×224ResNet标准输入尺寸 transforms.RandomHorizontalFlip(p0.5),# 随机水平翻转 transforms.RandomVerticalFlip(p0.5), # 随机垂直翻转 transforms.RandomGrayscale(p0.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, transformNone): self.file_path file_path self.imgs [] self.labels [] self.transform transform # 读取txt标注文件每行格式 图片路径 类别标签 with open(self.file_path, r, encodingutf-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, dtypenp.int64)) return image, label # 4. 构建DataLoader数据加载器 training_data food_dataset(file_path./train.txt, transformdata_transforms[trainda]) test_data food_dataset(file_path./test.txt, transformdata_transforms[valid]) train_dataloader DataLoader(training_data, batch_size64, shuffleTrue) test_dataloader DataLoader(test_data, batch_size64, shuffleTrue) # 5. 设备自动适配GPU/CUDA/MPS/CPU device cuda if torch.cuda.is_available() else mps if torch.backends.mps.is_available() else cpu print(fUsing {device} device) # 模型移至GPU/CPU无需括号原因下文详解 model resnet_model.to(device) # 6. 损失函数、优化器、学习率衰减 loss_fn nn.CrossEntropyLoss() # 多分类标准损失函数 # 仅更新解冻的全连接层参数 optimizer torch.optim.Adam(params_to_update, lr0.001) # 每5轮epoch学习率×0.5逐步降低学习率 scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size5, gamma0.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(fTest 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(fEpoch {t1}\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运行resnet_model models.resnet18(weightsmodels.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)weightsmodels.ResNet18_Weights.DEFAULT加载 ImageNet 百万图像预训练权重网络已经学会通用边缘、纹理、色彩特征param.requires_grad False冻结参数反向传播时不会更新卷积层权重只训练最后自定义全连接层ResNet18 默认输出 1000 类替换fc层将输出改为 20适配食物 20 分类任务。两种训练模式切换模式 1代码默认冻结主干仅微调全连接层适合数据集较小、硬件算力不足训练快、不易过拟合模式 2解冻全部参数全量微调注释冻结循环代码优化器改为读取全部参数python运行# 注释冻结代码 # for param in resnet_model.parameters(): # param.requires_grad False # 优化器传入全部参数 optimizer torch.optim.Adam(resnet_model.parameters(), lr0.001)适合数据集量大、算力充足整体精度上限更高。3.2 答疑model resnet_model.to(device)为什么不用加括号新手自定义 CNN 网络时写法model CNN().to(device)CNN()实例化网络创建新对象 本文代码resnet_model已经提前实例化完成不需要再次调用构造函数直接调用.to(device)迁移设备即可。python运行# 分步拆解 # 1. 实例化预训练模型已完成 resnet_model models.resnet18(...) # 2. 直接迁移至GPU无需再次实例化 model resnet_model.to(device)3.3 数据增强与归一化说明训练集使用大量随机变换扩充样本防止过拟合验证集仅做基础缩放不添加随机操作旋转、翻转、灰度化模拟真实场景拍摄角度、光线变化224×224ResNet 网络固定输入尺寸归一化均值方差是 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 训练 / 测试流程关键点model.train()训练模式Dropout、BatchNorm 启用更新model.eval()验证模式关闭随机层固定归一化参数with torch.no_grad()验证阶段关闭梯度计算大幅节省显存StepLR学习率衰减每 5 轮学习率减半后期收敛更稳定CrossEntropyLoss多分类专用损失标签无需 one-hot 编码直接输入数字标签。四、数据集文件配置说明新建train.txt、test.txt放在代码同级目录文本每行格式图片路径 类别编号类别从 0 开始依次递增图片路径支持相对路径确保路径无中文、无空格。五、拓展作业单张图片推理预测输入图片输出分类结果在代码末尾追加推理函数实现单图输入输出类别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})六、常见问题CUDA out of memory 显存溢出调小batch_size64改为 16/32或使用 CPU 运行test.txt 读取报错检查 txt 每行分隔符是空格末尾无空行图片路径存在精度持续很低确认归一化参数正确、训练集数据增强正常可切换全量微调模式MPS 设备报错MacPyTorch 版本更新至 2.0 以上MPS 仅支持新版 torch。七、总结迁移学习核心逻辑复用预训练卷积特征提取器仅替换输出层适配自定义分类任务两种训练方案按需选择小数据集冻结主干大数据集全量微调完整工程化流程自定义数据集→数据增强→模型构建→训练循环→验证评估代码可直接拓展增加模型保存、绘制 loss/acc 曲线、单图推理功能。