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/22 6:01:02 阅读更多 →
AI低代码开发:从自然语言到系统原型的革命

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

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

2026/7/22 6:01:02 阅读更多 →
n8n开源自动化工具:从入门到企业级部署

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

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

2026/7/22 6:00:02 阅读更多 →

最新新闻

苹果涨价背后的消费电子行业变革

苹果涨价背后的消费电子行业变革

1. 涨价现象背后的行业信号上周三凌晨,苹果官网悄然更新了Mac产品线的价格标签。MacBook Air基础款从7999元调整为8499元,MacBook Pro 14英寸入门型号从14999元涨至15999元,iMac 24英寸版本也有500-800元不等的涨幅。这波平均6-8%的调价幅度&…

2026/7/22 6:44:19 阅读更多 →
2026-07-22:最大化特殊下标数目的最少增加次数。用go语言,给定一个长度为 n 的整数数组,如果某个下标 i(不是第一个也不是最后一个)满足它对应的元素比左右邻居都大,那么这个位置就算作“特殊

2026-07-22:最大化特殊下标数目的最少增加次数。用go语言,给定一个长度为 n 的整数数组,如果某个下标 i(不是第一个也不是最后一个)满足它对应的元素比左右邻居都大,那么这个位置就算作“特殊

2026-07-22:最大化特殊下标数目的最少增加次数。用go语言,给定一个长度为 n 的整数数组,如果某个下标 i(不是第一个也不是最后一个)满足它对应的元素比左右邻居都大,那么这个位置就算作“特殊位置”。 你可…

2026/7/22 6:44:19 阅读更多 →
导购比价小程序场景|京东联盟商品详情对接方案|多规格图文拉取技术实操

导购比价小程序场景|京东联盟商品详情对接方案|多规格图文拉取技术实操

一、业务背景:导购比价小程序核心落地痛点随着私域流量、内容种草、电商导购模式快速普及,轻量化导购比价小程序成为个人创业者、自媒体团队、电商服务商的主流变现载体。这类小程序核心能力是聚合京东海量商品、展示完整商品信息、实现多商品比价、精准…

2026/7/22 6:44:19 阅读更多 →
Java与Lua集成实战:构建可热更新的动态规则引擎

Java与Lua集成实战:构建可热更新的动态规则引擎

1. 项目概述:当Java遇见Lua,静态架构的动态革命 在传统的Java开发世界里,我们习惯了“编译-打包-部署”的固定流程。每次业务逻辑的微小变动,都可能意味着一次繁琐的发布、重启和验证。尤其是在需要快速响应市场变化、频繁调整策略…

2026/7/22 6:44:19 阅读更多 →
WinDbg Preview与KDNET v2协议详解及配置指南

WinDbg Preview与KDNET v2协议详解及配置指南

1. WinDbg Preview与KDNET v2协议概述微软商店版WinDbg Preview近期迎来重要更新,正式加入对KDNET v2协议的支持。作为Windows内核调试的核心工具,这一升级显著改善了远程调试体验。KDNET(Kernel Debugging over Network)是微软推…

2026/7/22 6:44:19 阅读更多 →
AR数字孪生:重塑工业运维的下一代交互范式

AR数字孪生:重塑工业运维的下一代交互范式

随着工业4.0进程的深入,传统运维模式正面临数据孤岛、专家资源稀缺以及现场作业效率低下等严峻挑战。增强现实(AR)技术与数字孪生(Digital Twin)的深度融合,正在构建一种全新的工业交互范式。这种范式不仅实…

2026/7/22 6:43:19 阅读更多 →

日新闻

TI DSP系统配置模块SYSCFG详解:中断机制与主设备优先级配置实战

TI DSP系统配置模块SYSCFG详解:中断机制与主设备优先级配置实战

1. 项目概述与SYSCFG模块的核心价值在嵌入式系统,尤其是像TI C6000系列这样的高性能DSP开发中,我们常常会与芯片手册里那些密密麻麻的寄存器打交道。很多开发者可能更关注算法实现、内存优化或者外设驱动,但对于一个稳定、高效的系统而言&…

2026/7/22 0:00:26 阅读更多 →
微信Server酱:高到达率的应急通知方案实践

微信Server酱:高到达率的应急通知方案实践

1. 为什么我们需要"最次"的通知方案? 在数字化协作环境中,消息通知系统的重要性不言而喻明。但现实情况是,企业级通知方案往往需要复杂的API对接(如企业微信、钉钉、飞书),个人开发者的小项目又经…

2026/7/22 0:00:26 阅读更多 →
甲方要的“简洁“PPT,到底是简洁还是省事?

甲方要的“简洁“PPT,到底是简洁还是省事?

甲方说"简洁一点",乙方听到的是"少做几页"。甲方说"不要太复杂",乙方理解成"别放图表了"。结果交过去,甲方说"我说的简洁不是这个意思"。"简洁"这个词在PPT语境里,是…

2026/7/22 0:00:26 阅读更多 →

周新闻

Go语言静态资源打包方案对比与实践指南

Go语言静态资源打包方案对比与实践指南

1. 项目背景与核心需求在Go语言开发中,我们经常需要处理静态资源文件的打包问题。无论是Web应用的模板文件、前端资源,还是配置文件、证书等,都需要随程序一起分发。传统做法是将这些文件与编译后的二进制文件放在同一目录下,但这…

2026/7/21 8:48:31 阅读更多 →
Go语言实现高性能LDAP认证服务的架构与实践

Go语言实现高性能LDAP认证服务的架构与实践

1. 项目背景与核心价值LDAP(轻量级目录访问协议)作为企业级身份认证的黄金标准,已经服务了超过80%的财富500强公司。我在金融科技领域实施统一认证体系时,发现传统Java方案存在启动慢、内存占用高等痛点。而Go语言凭借其协程并发模…

2026/7/21 5:34:47 阅读更多 →
【AI面试官实战指南】:用ChatGPT模拟10类高频技术岗面试,3天提升应答精准度92%

【AI面试官实战指南】:用ChatGPT模拟10类高频技术岗面试,3天提升应答精准度92%

更多请点击: https://intelliparadigm.com 第一章:AI面试官实战指南的核心价值与适用场景 AI面试官并非替代人类HR的“黑箱工具”,而是以可解释、可审计、可迭代的方式,赋能招聘全链路的关键基础设施。其核心价值在于将主观经验沉…

2026/7/21 8:25:39 阅读更多 →

月新闻