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 曲线、单图推理功能。

相关新闻

可交换性在统计证据聚合中的应用与实操指南

可交换性在统计证据聚合中的应用与实操指南

1. 先搞清楚“可交换性”在统计证据聚合中到底解决什么问题 如果你处理过多个来源的统计检验结果,比如医学研究中不同临床试验的 p 值、工业质检中多批次抽样的异常分数、或者金融风控中多个模型的预警信号,你肯定遇到过这样的困境:每个独立检…

2026/7/27 9:22:24 阅读更多 →
AI低代码开发:从自然语言到系统原型的革命

AI低代码开发:从自然语言到系统原型的革命

1. AI低代码开发的技术革命2026年的软件开发领域正在经历一场前所未有的范式转移。当我第一次用自然语言描述需求,5分钟后就看到完整可运行的管理系统时,意识到传统编程方式正在被重新定义。AI低代码的组合,正在将软件开发从"手工作坊&q…

2026/7/27 3:26:03 阅读更多 →
n8n开源自动化工具:从入门到企业级部署

n8n开源自动化工具:从入门到企业级部署

1. 为什么你需要n8n自动化工具每天面对重复的数据搬运、表单填写、邮件发送,你是否感觉自己在做"数字流水线工人"?我曾在电商公司负责运营报表工作,每天要手动从5个平台导出数据,再用Excel做合并计算,整个过…

2026/7/26 20:31:31 阅读更多 →

最新新闻

嵌入式低功耗设计:TI CP3SP33 MIWU模块与GPIO配置实战指南

嵌入式低功耗设计:TI CP3SP33 MIWU模块与GPIO配置实战指南

1. 项目概述:低功耗嵌入式系统的“守夜人”在电池供电的嵌入式设备里,比如你手腕上的智能手环、家里的温湿度传感器,或者工业现场的无线采集节点,最核心的矛盾是什么?是“随时待命”的实时响应需求与“寸电寸金”的有限…

2026/7/27 23:01:20 阅读更多 →
虚拟机网卡显示10G速率,实际传输仅1G带宽完整排查方案

虚拟机网卡显示10G速率,实际传输仅1G带宽完整排查方案

虚拟机内部操作系统识别vmxnet3虚拟网卡为10Gbps规格,不代表实际带宽能跑满10G,核心瓶颈集中在物理层协商与虚拟网卡配置;优先排查两点:1、ESXi主机上联物理网卡、接入交换机端口速率协商是否锁定为1G,物理链路速率上限…

2026/7/27 23:01:20 阅读更多 →
如何高效使用R代码格式化工具styler:专业开发者的终极指南

如何高效使用R代码格式化工具styler:专业开发者的终极指南

如何高效使用R代码格式化工具styler:专业开发者的终极指南 【免费下载链接】styler Non-invasive pretty printing of R code 项目地址: https://gitcode.com/gh_mirrors/st/styler styler是一款专为R语言设计的非侵入式代码格式化神器,它能自动美…

2026/7/27 23:01:20 阅读更多 →
进口门窗五金品牌怎么选?2025年十大进口品牌实力盘点,看完不踩坑!

进口门窗五金品牌怎么选?2025年十大进口品牌实力盘点,看完不踩坑!

很多人装修门窗的时候,都会把注意力放在型材、玻璃上,却容易忽略门窗五金的重要性——其实五金件就像门窗的"心脏",开关、密封、承重全靠它支撑,占到了门窗总成本的20%~30%,却决定了门窗80%以上的使用体验和…

2026/7/27 23:01:20 阅读更多 →
Replit生态解析:从在线IDE到云原生开发平台的演进

Replit生态解析:从在线IDE到云原生开发平台的演进

最近在开发者圈子里,Replit 这个名字出现的频率越来越高。但很多人对它的认知还停留在"一个在线 IDE"的层面,以为它只是又一个云端代码编辑器。实际上,Replit 正在做的事情,远比表面看起来要深远得多。 上周 Replit 举…

2026/7/27 23:01:20 阅读更多 →
3分钟学会用AI制作短视频:从新手到专家的完整指南

3分钟学会用AI制作短视频:从新手到专家的完整指南

3分钟学会用AI制作短视频:从新手到专家的完整指南 【免费下载链接】MoneyPrinterTurbo 利用 AI 大模型和自动化工作流,根据主题或关键词一键生成高清短视频。Generate HD short videos from a topic or keyword with an automated AI workflow. 项目地…

2026/7/27 23:00:20 阅读更多 →

日新闻

【JAVA毕设源码分享】基于SpringBoot的社区智能垃圾管理系统的设计与实现(程序+文档+代码讲解+一条龙定制)

【JAVA毕设源码分享】基于SpringBoot的社区智能垃圾管理系统的设计与实现(程序+文档+代码讲解+一条龙定制)

博主介绍:✌️码农一枚 ,专注于大学生项目实战开发、讲解和毕业🚢文撰写修改等。全栈领域优质创作者,博客之星、掘金/华为云/阿里云/InfoQ等平台优质作者、专注于Java、小程序技术领域和毕业项目实战 ✌️技术范围:&am…

2026/7/27 0:00:54 阅读更多 →
SPI实战指南:从时钟模式到寄存器配置,解决嵌入式通信难题

SPI实战指南:从时钟模式到寄存器配置,解决嵌入式通信难题

1. 项目概述:从寄存器手册到实战指南 如果你手头有一份类似德州仪器(TI)TMS320x240xA系列DSP的SPI模块技术手册,看着里面密密麻麻的寄存器位定义、时序图和公式,是不是感觉头大?这份资料虽然权威&#xff0…

2026/7/27 0:00:54 阅读更多 →
【JAVA毕设源码分享】基于springboot的水果购物管理系统的设计与实现(程序+文档+代码讲解+一条龙定制)

【JAVA毕设源码分享】基于springboot的水果购物管理系统的设计与实现(程序+文档+代码讲解+一条龙定制)

博主介绍:✌️码农一枚 ,专注于大学生项目实战开发、讲解和毕业🚢文撰写修改等。全栈领域优质创作者,博客之星、掘金/华为云/阿里云/InfoQ等平台优质作者、专注于Java、小程序技术领域和毕业项目实战 ✌️技术范围:&am…

2026/7/27 0:00:54 阅读更多 →

周新闻

深度学习道路桥梁裂缝检测系统 道路桥梁裂缝检测数据集 道路桥梁病害识别检测数据集

深度学习道路桥梁裂缝检测系统 道路桥梁裂缝检测数据集 道路桥梁病害识别检测数据集

深度学习道路桥梁裂缝检测系统 数据集6000张 完整源码已标注数据集训练好的模型环境配置教程程序运行说明文档,可以直接使用!系统支持图片、视频、摄像头等多种方式检测裂缝,功能强大实用。 1数据集6000张 8各类别

2026/7/27 4:33:59 阅读更多 →
深度学习YOLO模型如何训练 PUBG 绝地求生目标检测数据集

深度学习YOLO模型如何训练 PUBG 绝地求生目标检测数据集

pubg数据集 精选原图1.42万数据 1.49万标签 无任何重复、算法增强或冗余图像! pubg绝地求生目标检测数据集 1分类:e_body,14905个标签,txt格式 共计14244张图,99%为640*640尺寸图像 适合yolo目标检测、AI训练关键词&am…

2026/7/27 6:31:56 阅读更多 →
Apex英雄目标检测数据集 深度学习框架YOLO如何训练APEX数据集

Apex英雄目标检测数据集 深度学习框架YOLO如何训练APEX数据集

Apex检测数据集数据集详情检测类别: allies enemy tag图片总量:7247张训练集:5139张验证集:1425张测试集:683张标注状态:全部已标注,即拿即用数据格式:支持YOLO格式及其他格式&#…

2026/7/27 4:01:12 阅读更多 →

月新闻