首页/新闻资讯/正文详情

食物图像分类代码实战

发布时间:2026/9/24 12:30:12 来源:云帆数科 栏目:资讯中心
食物图像分类代码实战
前言延续之前所讲基本上项目代码都是数据集的读入和处理模型定义、训练之前的各种准备设置以及训练流程那么接下来也是按照这个顺序进行。数据集读入和处理train_transform transforms.Compose( [ transforms.ToPILImage(), # Convert image to PIL format (224,224,3) - (3,224,224) transforms.RandomResizedCrop(224), # Random crop and resize to 224x224 transforms.RandomRotation(50), # Apply random rotation up to 50 degrees transforms.ToTensor() # Convert to tensor ] ) val_transform transforms.Compose( [ transforms.ToPILImage(), # Convert image to PIL format transforms.ToTensor() # Convert to tensor ] )图像数据集有点区别于前面的回归模型特征数据集通常在读入数据集阶段可以选用数据增强技术对图像处理这种技术对图片进行随机放大裁剪旋转等操作通过内置强化学习算法自动选择最优方式这个过程类似在分类任务中让模型见识不同角度各种各样的某类物体可以提高模型识别能力。另一方面也可以拓宽训练集抑制模型过拟合。但是在测试阶段不使用该技术测试集上数据模型都没见过可以检验模型泛化能力。class FoodDataset(Dataset): def __init__(self, path, modetrain): self.mode mode self.transform train_transform if mode train else val_transform self.X, self.Y self._load_data(path) def _load_data(self, path): X, Y None, None for class_idx in range(11): class_dir os.path.join(path, f{class_idx:02d}) img_files os.listdir(class_dir) class_images np.zeros((len(img_files), HW, HW, 3), dtypenp.uint8) class_labels np.full(len(img_files), class_idx, dtypenp.uint8) for idx, filename in enumerate(img_files): img_path os.path.join(class_dir, filename) img Image.open(img_path).resize((HW, HW)) class_images[idx] img if class_idx 0: X, Y class_images, class_labels else: X np.concatenate((X, class_images), axis0) Y np.concatenate((Y, class_labels), axis0) print(fLoaded {len(Y)} samples) return X, Y def __getitem__(self, index): return self.transform(self.X[index]), self.Y[index] def __len__(self): return len(self.X)这次项目训练集和测试集在不同文件因此不需要像之前一样拆分直接根据mode用读取不同的文件即可。不同项目的实现方式各有差异文件读取功能会根据数据集在本地的存储路径进行配置。Dataset类负责数据读取将原始数据转换为三维数值表示而Dataloader则用于处理这些数据集如批量加载数据支持数据打乱和多批次处理功能。模型定义class MyModel(nn.Module): def __init__(self, num_class): super(MyModel, self).__init__() # Input: 3x224x224 - Output: 512x7x7 - Flatten - Fully connected layers # Initial convolution block self.conv1 nn.Conv2d(3, 64, kernel_size3, stride1, padding1) # Output: 64x224x224 self.bn1 nn.BatchNorm2d(64) self.relu nn.ReLU() self.pool1 nn.MaxPool2d(2) # Output: 64x112x112 # Feature extraction layers self.layer1 nn.Sequential( nn.Conv2d(64, 128, kernel_size3, stride1, padding1), # 128x112x112 nn.BatchNorm2d(128), nn.ReLU(), nn.MaxPool2d(2) # 128x56x56 ) self.layer2 nn.Sequential( nn.Conv2d(128, 256, kernel_size3, stride1, padding1), # 256x56x56 nn.BatchNorm2d(256), nn.ReLU(), nn.MaxPool2d(2) # 256x28x28 ) self.layer3 nn.Sequential( nn.Conv2d(256, 512, kernel_size3, stride1, padding1), # 512x28x28 nn.BatchNorm2d(512), nn.ReLU(), nn.MaxPool2d(2) # 512x14x14 ) # Final pooling and classifier self.pool2 nn.MaxPool2d(2) # 512x7x7 self.fc1 nn.Linear(512*7*7, 1000) # 25088 - 1000 self.relu2 nn.ReLU() self.fc2 nn.Linear(1000, num_class) # 1000 - num_class def forward(self, x): x self.pool1(self.relu(self.bn1(self.conv1(x)))) x self.layer1(x) x self.layer2(x) x self.layer3(x) x self.pool2(x) x x.view(x.size(0), -1) # Flatten x self.relu2(self.fc1(x)) x self.fc2(x) return x模型定义相对简单且相似。值得一提的是数据作为新时代的石油资源我们训练的模型通常难以与投入数百万美元训练的大模型相媲美因此可以采用迁移学习策略。迁移学习主要分为微调和线性探测两种方式二者的核心区别在于是否冻结主干网络的参数。训练流程def train_val(model, train_loader, val_loader, no_label_loader, device, epochs, optimizer, loss, thres, save_path): model model.to(device) plt_train_loss [] plt_val_loss [] plt_train_acc [] plt_val_acc [] max_acc 0.0 for epoch in range(epochs): train_loss 0.0 val_loss 0.0 train_acc 0.0 val_acc 0.0 start_time time.time() # Training phase model.train() for batch_x, batch_y in train_loader: x, target batch_x.to(device), batch_y.to(device) pred model(x) train_bat_loss loss(pred, target) train_bat_loss.backward() optimizer.step() optimizer.zero_grad() train_loss train_bat_loss.item() train_acc (pred.argmax(dim1) target).sum().item() avg_train_loss train_loss / len(train_loader) avg_train_acc train_acc / len(train_loader.dataset) plt_train_loss.append(avg_train_loss) plt_train_acc.append(avg_train_acc) # Validation phase model.eval() with torch.no_grad(): for batch_x, batch_y in val_loader: x, target batch_x.to(device), batch_y.to(device) pred model(x) val_bat_loss loss(pred, target) val_loss val_bat_loss.item() val_acc (pred.argmax(dim1) target).sum().item() avg_val_loss val_loss / len(val_loader) avg_val_acc val_acc / len(val_loader.dataset) plt_val_loss.append(avg_val_loss) plt_val_acc.append(avg_val_acc) # Semi-supervised learning if epoch % 3 0 and avg_val_acc 0.6: semi_loader get_semi_loader(no_label_loader, model, device, thres) # Save best model if avg_val_acc max_acc: torch.save(model, save_path) max_acc avg_val_acc # Print progress elapsed time.time() - start_time print(f[{epoch:03d}/{epochs:03d}] {elapsed:.2f}s | fTrainLoss: {avg_train_loss:.6f} | ValLoss: {avg_val_loss:.6f} | fTrainAcc: {avg_train_acc:.6f} | ValAcc: {avg_val_acc:.6f}) # Plot training curves plt.figure(figsize(12, 4)) plt.subplot(1, 2, 1) plt.plot(plt_train_loss, labelTrain) plt.plot(plt_val_loss, labelVal) plt.title(Loss Curve) plt.legend() plt.subplot(1, 2, 2) plt.plot(plt_train_acc, labelTrain) plt.plot(plt_val_acc, labelVal) plt.title(Accuracy Curve) plt.legend() plt.show()主干训练流程也是大体相似这里和回归模型主要区别在于多了计算准确率识别正确图片占比因为分类任务输出 label 可以直接知道整体预测正确率。

相关推荐

深入解析MC68336/376微控制器:CPU32核心与集成外设实战指南
深入解析MC68336/376微控制器:CPU32核心与集成外设实战指南

1. 项目概述:深入MC68336/376的微控制器世界在嵌入式系统开发的早期黄金时代,Motorola(后为Freescale,现属NXP)的68K系列处理器以其优雅的架构和强大的性能,占据了工业控制、汽车电子和通信设备等领域的半壁… · 2026/9/23 0:36:41

如何实现VR设备跨品牌兼容:OpenVR空间校准器完整指南
如何实现VR设备跨品牌兼容:OpenVR空间校准器完整指南

如何实现VR设备跨品牌兼容:OpenVR空间校准器完整指南 【免费下载链接】OpenVR-SpaceCalibrator Use tracked VR devices from one company with any other. 项目地址: https://gitcode.com/gh_mirrors/op/OpenVR-SpaceCalibrator 你是否曾想过将HTC Vive的控… · 2026/9/24 17:46:53

Crawl4AI:为AI时代重新定义智能网页爬取的开源利器
Crawl4AI:为AI时代重新定义智能网页爬取的开源利器

Crawl4AI:为AI时代重新定义智能网页爬取的开源利器 【免费下载链接】crawl4ai 🚀🤖 Crawl4AI: Open-source LLM Friendly Web Crawler & Scraper. Dont be shy, join here: https://discord.gg/jP8KfhDhyN 项目地址: https://gitcode.c… · 2026/9/24 17:49:15

OpenClaw QQ插件v0.5.0:非机器人通道与全媒体+权限控制
OpenClaw QQ插件v0.5.0:非机器人通道与全媒体+权限控制

OpenClaw 的 QQ 插件 v0.5.0 终于发布正式版本了。这个版本最让我意外的是,它把“非机器人”这条路线坚持了下来,并且把全媒体消息和精细化权限控制这两个此前最难受的短板一次补齐。文章不聊虚的,就说说这个插件到底改了什么、权限规则怎么写… · 2026/9/24 23:41:12

腹部CT五器官分割:FCN-8s实战指南与避坑手册
腹部CT五器官分割:FCN-8s实战指南与避坑手册

简介:本资源是一套基于全卷积网络(FCN)实现腹部多脏器五类语义分割的完整实战项目,面向医学图像分析初学者与深度学习实践者,解决腹部CT影像中肝脏、脾脏、肾脏、胰腺及胃等器官的像素级精准分割问题。压缩包共1025个文… · 2026/9/24 23:41:12

OpenClaw v0.5.0 QQ插件:全媒体消息与精细化权限控制实战
OpenClaw v0.5.0 QQ插件:全媒体消息与精细化权限控制实战

OpenClaw QQ插件发到v0.5.0了,这次带上了全媒体消息和精细化权限控制。标题里“非机器人”三个字,我觉得是整篇最该聊清楚的地方——很多人一听“QQ插件”,第一反应是申请个QQ机器人接口,实际上这条路在自托管AI Agent的场景里远不… · 2026/9/24 23:41:12

I2C总线物理层与多主仲裁:从开漏输出到RTL实现的深度避坑指南
I2C总线物理层与多主仲裁:从开漏输出到RTL实现的深度避坑指南

I2C这东西,刚入行的时候觉得它简单得不行——两根线,一根时钟一根数据,挂一堆从设备,地址一喊谁应答谁说话,能有多难?结果真到了调试现场,波形抓出来一看,上升沿软塌塌像条抛物线&am… · 2026/9/24 23:41:12

Windows图标转换:从PNG到专业.ico的完整指南
Windows图标转换:从PNG到专业.ico的完整指南

1. 项目概述:一张图到.ico文件,到底在解决什么问题?“怎么把图片转换成ico图标文件?”——这句提问背后藏着的,不是单纯的技术操作,而是一整套Windows生态下的视觉一致性需求。我做桌面应用开发、系统工具打… · 2026/9/24 23:41:12

35岁求职寒冬自救指南:从简历到面试的转型策略
35岁求职寒冬自救指南:从简历到面试的转型策略

1. 先承认:到了这个阶段,找工作这件事完全变样了三十五岁那年冬天,我记得特别清楚。早上七点准时醒来,第一件事是摸手机看邮箱。收件箱里躺着三封新邮件——两封是订阅的行业资讯,一封是某个招聘网站系统自动推送的职位… · 2026/9/24 23:41:06

基于YOLOv8的渔船作业监控系统:从环境搭建到边缘部署全流程
基于YOLOv8的渔船作业监控系统:从环境搭建到边缘部署全流程

简介:这是一套面向计算机、人工智能、自动化等专业学生与教师的毕业设计级项目资源,围绕YOLOv8实现渔船作业监控系统,可用于毕设、课程设计、大作业或项目立项演示。压缩包共97个文件,约24.21MB,以70个Python源码文件为… · 2026/9/24 0:00:13

1D-CNN时间序列建模实战:从Conv1d原理到工业落地
1D-CNN时间序列建模实战:从Conv1d原理到工业落地

简介:面向时间序列数据建模的一维卷积神经网络完整实现,适合深度学习入门者及需要快速验证时序模型的研究者,能够从音频、文本、传感器或股价等序列中挖掘局部特征与时间依赖。压缩包体积很小,只有3KB,内含3个Python脚… · 2026/9/24 0:00:26

柔软的L:汉语语流中被忽视的舌肌张力控制
柔软的L:汉语语流中被忽视的舌肌张力控制

1. 这个“L”不是字母表里的L,而是舌尖上的L最近在几个方言群和语音教学社群里,反复看到有人发一句:“也说字母L:柔软的长舌”。初看以为是英语发音课笔记,点开才发现全是方言爱好者、播音系学生、语言康复师甚至戏曲演… · 2026/9/24 0:00:44

了解更多?预约专属演示

我们的顾问将为您一对一讲解产品与方案

企业微信二维码