From 4c2fa0e533ab6ce1804d8a9476a25ce7f3a4af72 Mon Sep 17 00:00:00 2001 From: zhanghuan <1262329256@qq.com> Date: Wed, 10 Sep 2025 14:48:43 +0800 Subject: [PATCH] =?UTF-8?q?=E6=8A=8A=E7=BD=91=E7=BB=9C=E6=A8=A1=E5=9E=8B?= =?UTF-8?q?=E5=8D=95=E7=8B=AC=E6=8B=8E=E5=87=BA=E6=9D=A5=E4=BA=86=EF=BC=8C?= =?UTF-8?q?=E6=96=B9=E4=BE=BF=E8=A7=A3=E8=80=A6=E3=80=82=E4=B8=8B=E4=B8=80?= =?UTF-8?q?=E6=AD=A5=E5=87=86=E5=A4=87=E6=8A=8A=E9=87=8D=E8=A6=81=E7=9A=84?= =?UTF-8?q?=E9=85=8D=E7=BD=AE=E5=85=A8=E9=83=A8=E6=8B=8E=E5=87=BA=E6=9D=A5?= =?UTF-8?q?=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- classifier/food_classifier_app.py | 64 +++---------------------------- toAndroid/toAndroid.py | 61 ++--------------------------- train/food_classifier.py | 64 ++++--------------------------- 3 files changed, 18 insertions(+), 171 deletions(-) diff --git a/classifier/food_classifier_app.py b/classifier/food_classifier_app.py index ba37b87..ff46079 100644 --- a/classifier/food_classifier_app.py +++ b/classifier/food_classifier_app.py @@ -13,66 +13,14 @@ from typing import List, Optional, Tuple from tkinterdnd2 import DND_FILES, TkinterDnD import threading import time +from net import create_food_cnn # 设置customtkinter的外观 ctk.set_appearance_mode("System") ctk.set_default_color_theme("blue") -# 定义CNN模型(与训练代码中的结构相同) -class FoodCNN(nn.Module): - def __init__(self): - super(FoodCNN, self).__init__() - # 第一个卷积块 - self.conv1 = nn.Conv2d(3, 32, 3, padding=1) - self.conv2 = nn.Conv2d(32, 32, 3, padding=1) - self.pool1 = nn.MaxPool2d(2, 2) - self.dropout1 = nn.Dropout2d(0.25) - - # 第二个卷积块 - self.conv3 = nn.Conv2d(32, 64, 3, padding=1) - self.conv4 = nn.Conv2d(64, 64, 3, padding=1) - self.pool2 = nn.MaxPool2d(2, 2) - self.dropout2 = nn.Dropout2d(0.25) - - # 第三个卷积块 - self.conv5 = nn.Conv2d(64, 128, 3, padding=1) - self.conv6 = nn.Conv2d(128, 128, 3, padding=1) - self.pool3 = nn.MaxPool2d(2, 2) - self.dropout3 = nn.Dropout2d(0.25) - - # 全连接层 - self.fc1 = nn.Linear(128 * 4 * 4, 512) - self.dropout4 = nn.Dropout(0.5) - self.fc2 = nn.Linear(512, 2) # 2分类 - - def forward(self, x): - # 第一个卷积块 - x = F.relu(self.conv1(x)) - x = F.relu(self.conv2(x)) - x = self.pool1(x) - x = self.dropout1(x) - - # 第二个卷积块 - x = F.relu(self.conv3(x)) - x = F.relu(self.conv4(x)) - x = self.pool2(x) - x = self.dropout2(x) - - # 第三个卷积块 - x = F.relu(self.conv5(x)) - x = F.relu(self.conv6(x)) - x = self.pool3(x) - x = self.dropout3(x) - - # 展平 - x = x.view(-1, 128 * 4 * 4) - - # 全连接层 - x = F.relu(self.fc1(x)) - x = self.dropout4(x) - x = self.fc2(x) - - return x + + class FoodClassifierApp: def __init__(self, root): @@ -81,7 +29,7 @@ class FoodClassifierApp: self.root.geometry("1400x800") # 食物类别(根据您的数据集) - self.food_classes = ["回锅肉", "西红柿鸡蛋"] + self.food_classes = ["回锅肉", "西红柿鸡蛋","麻辣小面"] # 设备设置 self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") @@ -106,10 +54,10 @@ class FoodClassifierApp: def load_model(self): """加载训练好的PyTorch模型""" try: - model_path = "../model/01/best_food_model.pth" + model_path = "../model/02/best_food_model.pth" if os.path.exists(model_path): # 创建模型实例 - self.model = FoodCNN() + self.model = create_food_cnn() # 加载模型权重 self.model.load_state_dict(torch.load(model_path, map_location=self.device)) self.model.to(self.device) diff --git a/toAndroid/toAndroid.py b/toAndroid/toAndroid.py index 8c293d6..04fca23 100644 --- a/toAndroid/toAndroid.py +++ b/toAndroid/toAndroid.py @@ -1,66 +1,13 @@ import torch import torch.nn as nn import torch.nn.functional as F +from net import create_food_cnn -class FoodCNN(nn.Module): - def __init__(self): - super(FoodCNN, self).__init__() - # 第一个卷积块 - self.conv1 = nn.Conv2d(3, 32, 3, padding=1) - self.conv2 = nn.Conv2d(32, 32, 3, padding=1) - self.pool1 = nn.MaxPool2d(2, 2) - self.dropout1 = nn.Dropout2d(0.25) - - # 第二个卷积块 - self.conv3 = nn.Conv2d(32, 64, 3, padding=1) - self.conv4 = nn.Conv2d(64, 64, 3, padding=1) - self.pool2 = nn.MaxPool2d(2, 2) - self.dropout2 = nn.Dropout2d(0.25) - - # 第三个卷积块 - self.conv5 = nn.Conv2d(64, 128, 3, padding=1) - self.conv6 = nn.Conv2d(128, 128, 3, padding=1) - self.pool3 = nn.MaxPool2d(2, 2) - self.dropout3 = nn.Dropout2d(0.25) - - # 全连接层 - self.fc1 = nn.Linear(128 * 4 * 4, 512) - self.dropout4 = nn.Dropout(0.5) - self.fc2 = nn.Linear(512, 2) # 2分类 - - def forward(self, x): - # 第一个卷积块 - x = F.relu(self.conv1(x)) - x = F.relu(self.conv2(x)) - x = self.pool1(x) - x = self.dropout1(x) - - # 第二个卷积块 - x = F.relu(self.conv3(x)) - x = F.relu(self.conv4(x)) - x = self.pool2(x) - x = self.dropout2(x) - - # 第三个卷积块 - x = F.relu(self.conv5(x)) - x = F.relu(self.conv6(x)) - x = self.pool3(x) - x = self.dropout3(x) - - # 展平 - x = x.view(-1, 128 * 4 * 4) - - # 全连接层 - x = F.relu(self.fc1(x)) - x = self.dropout4(x) - x = self.fc2(x) - - return x # 1. 初始化模型 -model = FoodCNN() +model = create_food_cnn() # 2. 加载训练好的权重 -model.load_state_dict(torch.load("../model/01/best_food_model.pth", map_location='cpu')) +model.load_state_dict(torch.load("../model/02/best_food_model.pth", map_location='cpu')) model.eval() # 设置为推理模式 # 3. 创建示例输入 (假设输入是 3x224x224 的图片) @@ -68,4 +15,4 @@ example_input = torch.randn(1, 3, 224, 224) # 4. 转换为 TorchScript traced_script_module = torch.jit.trace(model, example_input) -traced_script_module.save("../model/01/best_food_model_mobile.pt") +traced_script_module.save("../model/02/best_food_model_mobile.pt") diff --git a/train/food_classifier.py b/train/food_classifier.py index d43f51d..2bd6997 100644 --- a/train/food_classifier.py +++ b/train/food_classifier.py @@ -9,8 +9,14 @@ import matplotlib import numpy as np from tqdm import tqdm import os +import sys import time +# 添加net目录到路径 +# sys.path.append(os.path.join(os.path.dirname(__file__), '..', 'net')) +# from food_net import create_food_cnn +from net import create_food_cnn + # 设置matplotlib支持中文显示 plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'DejaVu Sans'] # 指定默认字体 plt.rcParams['axes.unicode_minus'] = False # 解决保存图像是负号'-'显示为方块的问题 @@ -19,61 +25,7 @@ plt.rcParams['axes.unicode_minus'] = False # 解决保存图像是负号'-'显 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"使用设备: {device}") -# 定义CNN模型(基于CIFAR10结构,输出层改为3分类) -class FoodCNN(nn.Module): - def __init__(self): - super(FoodCNN, self).__init__() - # 第一个卷积块 - self.conv1 = nn.Conv2d(3, 32, 3, padding=1) - self.conv2 = nn.Conv2d(32, 32, 3, padding=1) - self.pool1 = nn.MaxPool2d(2, 2) - self.dropout1 = nn.Dropout2d(0.25) - - # 第二个卷积块 - self.conv3 = nn.Conv2d(32, 64, 3, padding=1) - self.conv4 = nn.Conv2d(64, 64, 3, padding=1) - self.pool2 = nn.MaxPool2d(2, 2) - self.dropout2 = nn.Dropout2d(0.25) - - # 第三个卷积块 - self.conv5 = nn.Conv2d(64, 128, 3, padding=1) - self.conv6 = nn.Conv2d(128, 128, 3, padding=1) - self.pool3 = nn.MaxPool2d(2, 2) - self.dropout3 = nn.Dropout2d(0.25) - - # 全连接层 - self.fc1 = nn.Linear(128 * 4 * 4, 512) - self.dropout4 = nn.Dropout(0.5) - self.fc2 = nn.Linear(512, 3) # 改为2分类 - - def forward(self, x): - # 第一个卷积块 - x = F.relu(self.conv1(x)) - x = F.relu(self.conv2(x)) - x = self.pool1(x) - x = self.dropout1(x) - - # 第二个卷积块 - x = F.relu(self.conv3(x)) - x = F.relu(self.conv4(x)) - x = self.pool2(x) - x = self.dropout2(x) - - # 第三个卷积块 - x = F.relu(self.conv5(x)) - x = F.relu(self.conv6(x)) - x = self.pool3(x) - x = self.dropout3(x) - - # 展平 - x = x.view(-1, 128 * 4 * 4) - - # 全连接层 - x = F.relu(self.fc1(x)) - x = self.dropout4(x) - x = self.fc2(x) - - return x + # 数据预处理 transform_train = transforms.Compose([ @@ -201,7 +153,7 @@ if __name__ == '__main__': print(f"测试集大小: {len(test_dataset)}") # 创建模型 - model = FoodCNN().to(device) + model = create_food_cnn().to(device) print(f"模型参数数量: {sum(p.numel() for p in model.parameters() if p.requires_grad)}") # 定义损失函数和优化器(使用与CIFAR10相同的超参数)