增加了数据增强离线的程序(可扩充数据集),将图片压缩从32*32,调整为224*224.
This commit is contained in:
@@ -30,7 +30,7 @@ class FoodClassifierApp:
|
|||||||
self.root.geometry("1400x800")
|
self.root.geometry("1400x800")
|
||||||
|
|
||||||
# 食物类别(根据您的数据集)
|
# 食物类别(根据您的数据集)
|
||||||
self.food_classes = ["回锅肉", "西红柿鸡蛋","麻辣小面"]
|
self.food_classes = ['回锅肉', '炒细面', '西红柿鸡蛋', '麻辣小面']
|
||||||
|
|
||||||
# 设备设置
|
# 设备设置
|
||||||
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
|
|||||||
@@ -0,0 +1,50 @@
|
|||||||
|
import os
|
||||||
|
import random
|
||||||
|
from PIL import Image
|
||||||
|
from torchvision import transforms
|
||||||
|
|
||||||
|
# 原始数据目录
|
||||||
|
input_root = "../dataset/train"
|
||||||
|
|
||||||
|
# 增强后保存目录
|
||||||
|
output_root = "../dataset/train_aug"
|
||||||
|
|
||||||
|
# 每类目标张数
|
||||||
|
target_num = 500
|
||||||
|
|
||||||
|
# 定义数据增强
|
||||||
|
transform = transforms.Compose([
|
||||||
|
transforms.RandomHorizontalFlip(p=0.5),
|
||||||
|
transforms.RandomRotation(15),
|
||||||
|
transforms.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.1, hue=0.02),
|
||||||
|
transforms.RandomResizedCrop(size=(224, 224), scale=(0.8, 1.0)),
|
||||||
|
])
|
||||||
|
|
||||||
|
os.makedirs(output_root, exist_ok=True)
|
||||||
|
|
||||||
|
# 遍历每个类别文件夹
|
||||||
|
for class_name in os.listdir(input_root):
|
||||||
|
input_dir = os.path.join(input_root, class_name)
|
||||||
|
output_dir = os.path.join(output_root, class_name)
|
||||||
|
os.makedirs(output_dir, exist_ok=True)
|
||||||
|
|
||||||
|
# 读取类别下所有图片路径
|
||||||
|
img_files = [f for f in os.listdir(input_dir) if f.lower().endswith(('.jpg', '.png', '.jpeg'))]
|
||||||
|
img_paths = [os.path.join(input_dir, f) for f in img_files]
|
||||||
|
|
||||||
|
print(f"类别 {class_name} 原始图片数: {len(img_paths)}")
|
||||||
|
|
||||||
|
count = 0
|
||||||
|
while count < target_num:
|
||||||
|
img_path = random.choice(img_paths)
|
||||||
|
img = Image.open(img_path).convert("RGB")
|
||||||
|
# 生成增强图
|
||||||
|
aug_img = transform(img)
|
||||||
|
# 保存
|
||||||
|
save_path = os.path.join(output_dir, f"aug_{count:03d}.jpg")
|
||||||
|
aug_img.save(save_path)
|
||||||
|
count += 1
|
||||||
|
|
||||||
|
print(f"类别 {class_name} 已扩充到 {target_num} 张,保存于 {output_dir}")
|
||||||
|
|
||||||
|
print("✅ 数据增强完成!")
|
||||||
+11
-6
@@ -17,9 +17,9 @@ class FoodCNN(nn.Module):
|
|||||||
super(FoodCNN, self).__init__()
|
super(FoodCNN, self).__init__()
|
||||||
self.use_internal_preprocess = use_internal_preprocess
|
self.use_internal_preprocess = use_internal_preprocess
|
||||||
|
|
||||||
# 图片预处理变换(仅在推理时使用)
|
# 图片预处理变换(仅在推理时使用)(32)改为(224)
|
||||||
self.preprocess = transforms.Compose([
|
self.preprocess = transforms.Compose([
|
||||||
transforms.Resize((32, 32),interpolation=transforms.InterpolationMode.BILINEAR),
|
transforms.Resize((224, 224),interpolation=transforms.InterpolationMode.BILINEAR),
|
||||||
transforms.ToTensor(),
|
transforms.ToTensor(),
|
||||||
transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))
|
transforms.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))
|
||||||
])
|
])
|
||||||
@@ -49,7 +49,8 @@ class FoodCNN(nn.Module):
|
|||||||
self.dropout3 = nn.Dropout2d(0.25)
|
self.dropout3 = nn.Dropout2d(0.25)
|
||||||
|
|
||||||
# 全连接层
|
# 全连接层
|
||||||
self.fc1 = nn.Linear(128 * 4 * 4, 512)
|
# self.fc1 = nn.Linear(128 * 4 * 4, 512)
|
||||||
|
self.fc1 = nn.Linear(128 * 28 * 28, 512)
|
||||||
self.dropout4 = nn.Dropout(0.5)
|
self.dropout4 = nn.Dropout(0.5)
|
||||||
self.fc2 = nn.Linear(512, settings.NUM_CLASSES)
|
self.fc2 = nn.Linear(512, settings.NUM_CLASSES)
|
||||||
|
|
||||||
@@ -76,7 +77,8 @@ class FoodCNN(nn.Module):
|
|||||||
tensor = tensor.unsqueeze(0)
|
tensor = tensor.unsqueeze(0)
|
||||||
|
|
||||||
# 使用与移动端相同的插值方法缩放到32x32
|
# 使用与移动端相同的插值方法缩放到32x32
|
||||||
tensor = F.interpolate(tensor, size=(32, 32), mode='bilinear', align_corners=False)
|
# tensor = F.interpolate(tensor, size=(32, 32), mode='bilinear', align_corners=False)
|
||||||
|
tensor = F.interpolate(tensor, size=(224, 224), mode='bilinear', align_corners=False)
|
||||||
|
|
||||||
# ImageNet标准化
|
# ImageNet标准化
|
||||||
mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)
|
mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)
|
||||||
@@ -100,7 +102,8 @@ class FoodCNN(nn.Module):
|
|||||||
# x = x.float() / 255.0
|
# x = x.float() / 255.0
|
||||||
|
|
||||||
# 缩放到 32x32
|
# 缩放到 32x32
|
||||||
x = F.interpolate(x, size=(32, 32), mode='bilinear', align_corners=False)
|
# x = F.interpolate(x, size=(32, 32), mode='bilinear', align_corners=False)
|
||||||
|
x = F.interpolate(x, size=(224, 224), mode='bilinear', align_corners=False)
|
||||||
# 缩放到 32x32 - 使用 align_corners=True 来匹配 PIL 的默认行为
|
# 缩放到 32x32 - 使用 align_corners=True 来匹配 PIL 的默认行为
|
||||||
# x = F.interpolate(x, size=(32, 32), mode='bilinear', align_corners=True)
|
# x = F.interpolate(x, size=(32, 32), mode='bilinear', align_corners=True)
|
||||||
|
|
||||||
@@ -140,7 +143,9 @@ class FoodCNN(nn.Module):
|
|||||||
x = self.dropout3(x)
|
x = self.dropout3(x)
|
||||||
|
|
||||||
# 展平
|
# 展平
|
||||||
x = x.view(-1, 128 * 4 * 4)
|
# x = x.view(-1, 128 * 4 * 4)
|
||||||
|
x = x.view(-1, 128 * 28 * 28)
|
||||||
|
# x = x.view(-1, 65536)
|
||||||
|
|
||||||
# 全连接层
|
# 全连接层
|
||||||
x = F.relu(self.fc1(x))
|
x = F.relu(self.fc1(x))
|
||||||
|
|||||||
@@ -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', '12')
|
MODEL_DIR = os.path.join(BASE_DIR, 'model', '15')
|
||||||
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')
|
||||||
@@ -26,6 +26,7 @@ INFERENCE_BEST_MODEL_PATH = BEST_MODEL_PATH
|
|||||||
# 训练参数
|
# 训练参数
|
||||||
NUM_EPOCHS = 100
|
NUM_EPOCHS = 100
|
||||||
BATCH_SIZE = 32
|
BATCH_SIZE = 32
|
||||||
|
# BATCH_SIZE = 128
|
||||||
LEARNING_RATE = 0.001
|
LEARNING_RATE = 0.001
|
||||||
WEIGHT_DECAY = 1e-4
|
WEIGHT_DECAY = 1e-4
|
||||||
|
|
||||||
|
|||||||
@@ -30,17 +30,27 @@ print(f"使用设备: {device}")
|
|||||||
|
|
||||||
|
|
||||||
# 数据预处理 - 必须包含Resize以保证batch中tensor尺寸一致,只有Normalize由模型内部完成
|
# 数据预处理 - 必须包含Resize以保证batch中tensor尺寸一致,只有Normalize由模型内部完成
|
||||||
|
# 把32,32换成224,224
|
||||||
|
# transform_train = transforms.Compose([
|
||||||
|
# transforms.Resize((224, 224)),
|
||||||
|
# transforms.RandomHorizontalFlip(p=0.5),
|
||||||
|
# transforms.RandomRotation(10),
|
||||||
|
# transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1),
|
||||||
|
# transforms.ToTensor(),
|
||||||
|
# # 注意:只有Normalize由模型内部处理
|
||||||
|
# ])
|
||||||
transform_train = transforms.Compose([
|
transform_train = transforms.Compose([
|
||||||
transforms.Resize((32, 32)), # 必须保留,确保batch中tensor尺寸一致
|
transforms.Resize((224, 224)),
|
||||||
transforms.RandomHorizontalFlip(p=0.5),
|
transforms.RandomHorizontalFlip(p=0.5),
|
||||||
transforms.RandomRotation(10),
|
transforms.RandomRotation(10),
|
||||||
transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1),
|
transforms.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.1, hue=0.02),
|
||||||
transforms.ToTensor(),
|
transforms.ToTensor(),
|
||||||
# 注意:只有Normalize由模型内部处理
|
# 注意:只有Normalize由模型内部处理
|
||||||
])
|
])
|
||||||
|
|
||||||
|
# 把缩放成32*32,修改为224*224
|
||||||
transform_test = transforms.Compose([
|
transform_test = transforms.Compose([
|
||||||
transforms.Resize((32, 32)), # 必须保留,确保batch中tensor尺寸一致
|
transforms.Resize((224, 224)), # 必须保留,确保batch中tensor尺寸一致
|
||||||
transforms.ToTensor(),
|
transforms.ToTensor(),
|
||||||
# 注意:只有Normalize由模型内部处理
|
# 注意:只有Normalize由模型内部处理
|
||||||
])
|
])
|
||||||
@@ -129,7 +139,7 @@ def test(model, test_loader, device, class_names):
|
|||||||
})
|
})
|
||||||
|
|
||||||
print(f'\n测试集总体准确率: {100.*correct/total:.2f}%')
|
print(f'\n测试集总体准确率: {100.*correct/total:.2f}%')
|
||||||
for i in range(3):
|
for i in range(settings.NUM_CLASSES):
|
||||||
if class_total[i] > 0:
|
if class_total[i] > 0:
|
||||||
print(f'{class_names[i]} 准确率: {100.*class_correct[i]/class_total[i]:.2f}%')
|
print(f'{class_names[i]} 准确率: {100.*class_correct[i]/class_total[i]:.2f}%')
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user