增加了数据增强离线的程序(可扩充数据集),将图片压缩从32*32,调整为224*224.

This commit is contained in:
zhanghuan
2025-09-15 17:26:14 +08:00
parent d73331010c
commit 264d01e9b9
5 changed files with 80 additions and 14 deletions
+14 -4
View File
@@ -30,17 +30,27 @@ print(f"使用设备: {device}")
# 数据预处理 - 必须包含Resize以保证batch中tensor尺寸一致,只有Normalize由模型内部完成
# 把3232换成224224
# 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([
transforms.Resize((32, 32)), # 必须保留,确保batch中tensor尺寸一致
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.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.1, hue=0.02),
transforms.ToTensor(),
# 注意:只有Normalize由模型内部处理
])
# 把缩放成32*32,修改为224*224
transform_test = transforms.Compose([
transforms.Resize((32, 32)), # 必须保留,确保batch中tensor尺寸一致
transforms.Resize((224, 224)), # 必须保留,确保batch中tensor尺寸一致
transforms.ToTensor(),
# 注意:只有Normalize由模型内部处理
])
@@ -129,7 +139,7 @@ def test(model, test_loader, device, class_names):
})
print(f'\n测试集总体准确率: {100.*correct/total:.2f}%')
for i in range(3):
for i in range(settings.NUM_CLASSES):
if class_total[i] > 0:
print(f'{class_names[i]} 准确率: {100.*class_correct[i]/class_total[i]:.2f}%')