将分类的类别也应用在训练程序中,更好的解耦。
This commit is contained in:
+2
-1
@@ -3,6 +3,7 @@ import torch.nn as nn
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from torchvision import transforms
|
from torchvision import transforms
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
from settings import settings
|
||||||
|
|
||||||
|
|
||||||
class FoodCNN(nn.Module):
|
class FoodCNN(nn.Module):
|
||||||
@@ -50,7 +51,7 @@ class FoodCNN(nn.Module):
|
|||||||
# 全连接层
|
# 全连接层
|
||||||
self.fc1 = nn.Linear(128 * 4 * 4, 512)
|
self.fc1 = nn.Linear(128 * 4 * 4, 512)
|
||||||
self.dropout4 = nn.Dropout(0.5)
|
self.dropout4 = nn.Dropout(0.5)
|
||||||
self.fc2 = nn.Linear(512, 3) # 3分类
|
self.fc2 = nn.Linear(512, settings.NUM_CLASSES)
|
||||||
|
|
||||||
def preprocess_image(self, image):
|
def preprocess_image(self, image):
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ VAL_DATA_DIR = os.path.join(DATASET_DIR, 'val')
|
|||||||
TEST_DATA_DIR = os.path.join(DATASET_DIR, 'test')
|
TEST_DATA_DIR = os.path.join(DATASET_DIR, 'test')
|
||||||
|
|
||||||
# 模型保存路径
|
# 模型保存路径
|
||||||
MODEL_DIR = os.path.join(BASE_DIR, 'model', '10')
|
MODEL_DIR = os.path.join(BASE_DIR, 'model', '12')
|
||||||
BEST_MODEL_PATH = os.path.join(MODEL_DIR, 'best_food_model.pth')
|
BEST_MODEL_PATH = os.path.join(MODEL_DIR, 'best_food_model.pth')
|
||||||
TRAINING_CURVES_PATH = os.path.join(MODEL_DIR, 'training_curves.png')
|
TRAINING_CURVES_PATH = os.path.join(MODEL_DIR, 'training_curves.png')
|
||||||
TRAINING_RESULTS_PATH = os.path.join(MODEL_DIR, 'training_results.txt')
|
TRAINING_RESULTS_PATH = os.path.join(MODEL_DIR, 'training_results.txt')
|
||||||
@@ -47,7 +47,7 @@ COLOR_JITTER_SATURATION = 0.2
|
|||||||
COLOR_JITTER_HUE = 0.1
|
COLOR_JITTER_HUE = 0.1
|
||||||
|
|
||||||
# 模型参数
|
# 模型参数
|
||||||
NUM_CLASSES = 3
|
NUM_CLASSES = 4
|
||||||
|
|
||||||
# 其他配置
|
# 其他配置
|
||||||
NUM_WORKERS = 0 # Windows下建议设为0
|
NUM_WORKERS = 0 # Windows下建议设为0
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ def main():
|
|||||||
|
|
||||||
# 1. 加载训练好的基础模型权重
|
# 1. 加载训练好的基础模型权重
|
||||||
base_model = create_food_cnn()
|
base_model = create_food_cnn()
|
||||||
model_path = "../model/09/best_food_model.pth"
|
model_path = "../model/10/best_food_model.pth"
|
||||||
|
|
||||||
if not os.path.exists(model_path):
|
if not os.path.exists(model_path):
|
||||||
print(f"错误:模型文件不存在 {model_path}")
|
print(f"错误:模型文件不存在 {model_path}")
|
||||||
@@ -37,7 +37,7 @@ def main():
|
|||||||
traced_model = torch.jit.trace(mobile_model, example_input)
|
traced_model = torch.jit.trace(mobile_model, example_input)
|
||||||
|
|
||||||
# 保存模型
|
# 保存模型
|
||||||
output_path = "../model/09/best_food_model_mobile.pt"
|
output_path = "../model/10/best_food_model_mobile.pt"
|
||||||
traced_model.save(output_path)
|
traced_model.save(output_path)
|
||||||
print(f"✓ TorchScript模型保存成功: {output_path}")
|
print(f"✓ TorchScript模型保存成功: {output_path}")
|
||||||
|
|
||||||
|
|||||||
@@ -105,8 +105,8 @@ def test(model, test_loader, device, class_names):
|
|||||||
model.eval()
|
model.eval()
|
||||||
correct = 0
|
correct = 0
|
||||||
total = 0
|
total = 0
|
||||||
class_correct = list(0. for i in range(3))
|
class_correct = list(0. for i in range(settings.NUM_CLASSES))
|
||||||
class_total = list(0. for i in range(3))
|
class_total = list(0. for i in range(settings.NUM_CLASSES))
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
test_bar = tqdm(test_loader, desc='测试中')
|
test_bar = tqdm(test_loader, desc='测试中')
|
||||||
@@ -305,8 +305,8 @@ if __name__ == '__main__':
|
|||||||
f.write("各类别测试准确率:\n")
|
f.write("各类别测试准确率:\n")
|
||||||
f.write("-" * 30 + "\n")
|
f.write("-" * 30 + "\n")
|
||||||
model.eval()
|
model.eval()
|
||||||
class_correct = list(0. for i in range(3))
|
class_correct = list(0. for i in range(settings.NUM_CLASSES))
|
||||||
class_total = list(0. for i in range(3))
|
class_total = list(0. for i in range(settings.NUM_CLASSES))
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
for data, target in test_loader:
|
for data, target in test_loader:
|
||||||
|
|||||||
Reference in New Issue
Block a user