PyTorch玩转Fashion-MNIST:构建现代深度学习模型实战
PyTorch玩转Fashion-MNIST:构建现代深度学习模型实战
引言:告别MNIST,迎接更具挑战性的时尚识别任务
你是否已经厌倦了用手写数字数据集(MNIST)来测试你的深度学习模型?想要一个更贴近真实世界计算机视觉任务的基准数据集?Fashion-MNIST正是为解决这些问题而生。作为MNIST的直接替代品,它包含10个类别的时尚产品图片,分辨率同样为28×28,却能更真实地反映现代计算机视觉系统面临的挑战。本文将带你从零开始,使用PyTorch构建、训练和评估多个深度学习模型,掌握解决图像分类问题的核心技术。
读完本文后,你将能够:
- 熟练使用PyTorch加载和预处理Fashion-MNIST数据集
- 构建多层感知机(MLP)和卷积神经网络(CNN)模型
- 实现数据增强、批归一化和正则化等高级技术
- 优化模型性能并避免过拟合
- 可视化训练过程和结果分析
- 部署训练好的模型进行预测
1. Fashion-MNIST数据集详解
1.1 数据集概述
Fashion-MNIST是由Zalando研究团队开发的图像数据集,旨在替代过度简化的MNIST手写数字数据集。它包含60,000个训练样本和10,000个测试样本,每个样本都是28×28像素的灰度图像,分属10个时尚产品类别。
1.2 类别标签对应表
| 标签 | 类别描述 | 样本示例 |
|---|---|---|
| 0 | T-shirt/top (T恤/上衣) | ![T-shirt/top示例] |
| 1 | Trouser (裤子) | ![Trouser示例] |
| 2 | Pullover (套衫) | ![Pullover示例] |
| 3 | Dress (连衣裙) | ![Dress示例] |
| 4 | Coat (外套) | ![Coat示例] |
| 5 | Sandal (凉鞋) | ![Sandal示例] |
| 6 | Shirt (衬衫) | ![Shirt示例] |
| 7 | Sneaker (运动鞋) | ![Sneaker示例] |
| 8 | Bag (包) | ![Bag示例] |
| 9 | Ankle boot (短靴) | ![Ankle boot示例] |
1.3 与MNIST的对比优势
Fashion-MNIST相比传统MNIST具有以下优势:
MNIST过于简单,现代模型很容易达到99.7%以上的准确率,无法有效评估模型的泛化能力。而Fashion-MNIST更能反映真实世界的视觉识别任务难度,是检验模型性能的更好选择。
2. 环境准备与PyTorch基础
2.1 开发环境配置
首先,确保你的系统中安装了以下依赖库:
# 克隆仓库
git clone https://gitcode.com/gh_mirrors/fa/fashion-mnist
cd fashion-mnist
# 安装依赖
pip install torch torchvision matplotlib numpy pandas scikit-learn tqdm
2.2 PyTorch核心概念
在开始之前,让我们快速回顾PyTorch的几个核心概念:
- 张量(Tensor): 多维数组,类似于NumPy数组,但可以在GPU上运行
- 自动求导(Autograd): 自动计算张量的梯度,是反向传播的基础
- 神经网络模块(Module): 构建神经网络的基本组件
- 优化器(Optimizer): 实现各种优化算法,如SGD、Adam等
- 数据加载器(DataLoader): 高效加载和预处理数据的工具
2.3 检查PyTorch安装
import torch
import torchvision
import torch.nn as nn
import torch.optim as optim
import torch.nn.functional as F
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
# 检查PyTorch版本和GPU可用性
print(f"PyTorch版本: {torch.__version__}")
print(f"CUDA可用: {torch.cuda.is_available()}")
print(f"GPU数量: {torch.cuda.device_count()}")
if torch.cuda.is_available():
print(f"GPU名称: {torch.cuda.get_device_name(0)}")
3. 数据加载与预处理
3.1 使用PyTorch内置数据集
PyTorch的torchvision.datasets模块已内置Fashion-MNIST数据集,我们可以直接使用:
# 定义数据转换
transform = transforms.Compose([
transforms.ToTensor(), # 将PIL图像转换为Tensor (0-1范围)
transforms.Normalize((0.5,), (0.5,)) # 标准化到[-1, 1]范围
])
# 加载训练集和测试集
train_dataset = datasets.FashionMNIST(
root='./data', # 数据保存路径
train=True, # 训练集
download=True, # 如果本地没有则下载
transform=transform # 应用转换
)
test_dataset = datasets.FashionMNIST(
root='./data',
train=False, # 测试集
download=True,
transform=transform
)
3.2 自定义数据加载器
对于自定义数据集或需要更复杂的数据处理,我们可以使用项目中提供的mnist_reader模块:
# 自定义数据加载器示例
import sys
sys.path.append('./utils')
import mnist_reader
# 加载Fashion-MNIST数据
X_train, y_train = mnist_reader.load_mnist('data/fashion', kind='train')
X_test, y_test = mnist_reader.load_mnist('data/fashion', kind='t10k')
# 转换为PyTorch张量
X_train = torch.tensor(X_train, dtype=torch.float32).reshape(-1, 1, 28, 28) / 255.0
y_train = torch.tensor(y_train, dtype=torch.long)
X_test = torch.tensor(X_test, dtype=torch.float32).reshape(-1, 1, 28, 28) / 255.0
y_test = torch.tensor(y_test, dtype=torch.long)
# 创建数据集和数据加载器
train_dataset = torch.utils.data.TensorDataset(X_train, y_train)
test_dataset = torch.utils.data.TensorDataset(X_test, y_test)
3.3 创建数据加载器
使用DataLoader可以高效地批量加载数据,并支持打乱顺序和多线程加载:
# 定义批次大小
BATCH_SIZE = 64
# 创建数据加载器
train_loader = DataLoader(
dataset=train_dataset,
batch_size=BATCH_SIZE,
shuffle=True, # 训练集打乱顺序
num_workers=2 # 使用2个线程加载数据
)
test_loader = DataLoader(
dataset=test_dataset,
batch_size=BATCH_SIZE,
shuffle=False, # 测试集不需要打乱
num_workers=2
)
3.4 数据可视化
让我们可视化一些样本,直观了解数据集:
import matplotlib.pyplot as plt
import numpy as np
# 类别名称
class_names = ['T-shirt/top', 'Trouser', 'Pullover', 'Dress', 'Coat',
'Sandal', 'Shirt', 'Sneaker', 'Bag', 'Ankle boot']
# 创建图像网格
fig, axes = plt.subplots(5, 5, figsize=(15, 15))
axes = axes.flatten()
# 随机选择样本
indices = np.random.choice(len(train_dataset), 25, replace=False)
for i, idx in enumerate(indices):
img, label = train_dataset[idx]
# 将张量转换为图像格式
img = img.squeeze().numpy()
axes[i].imshow(img, cmap='gray')
axes[i].set_title(f"{class_names[label]} ({label})")
axes[i].axis('off')
plt.tight_layout()
plt.savefig('fashion_mnist_samples.png')
plt.show()
3.5 高级数据增强
为了提高模型的泛化能力,我们可以添加数据增强:
# 高级数据增强
train_transform = transforms.Compose([
transforms.RandomHorizontalFlip(p=0.2), # 随机水平翻转
transforms.RandomCrop(28, padding=2), # 随机裁剪
transforms.RandomRotation(10), # 随机旋转
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
# 使用增强后的数据
train_dataset_augmented = datasets.FashionMNIST(
root='./data',
train=True,
download=True,
transform=train_transform
)
train_loader_augmented = DataLoader(
dataset=train_dataset_augmented,
batch_size=BATCH_SIZE,
shuffle=True,
num_workers=2
)
4. 构建深度学习模型
4.1 简单多层感知机(MLP)
首先,让我们构建一个简单的多层感知机模型:
class SimpleMLP(nn.Module):
def __init__(self, input_size=784, hidden_size=128, num_classes=10):
super(SimpleMLP, self).__init__()
self.fc1 = nn.Linear(input_size, hidden_size) # 输入层到隐藏层
self.fc2 = nn.Linear(hidden_size, hidden_size) # 隐藏层到隐藏层
self.fc3 = nn.Linear(hidden_size, num_classes) # 隐藏层到输出层
self.dropout = nn.Dropout(0.25) # Dropout层防止过拟合
def forward(self, x):
# 将图像展平
x = x.view(-1, 28*28) # 28x28=784
# 前向传播
x = F.relu(self.fc1(x))
x = self.dropout(x)
x = F.relu(self.fc2(x))
x = self.dropout(x)
x = self.fc3(x)
return x
# 初始化模型
model_mlp = SimpleMLP()
# 如果有GPU则使用GPU
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model_mlp.to(device)
4.2 卷积神经网络(CNN)
卷积神经网络在图像识别任务上通常表现更好:
class SimpleCNN(nn.Module):
def __init__(self, num_classes=10):
super(SimpleCNN, self).__init__()
# 卷积层1: 1 -> 32通道,5x5卷积核
self.conv1 = nn.Conv2d(1, 32, kernel_size=5, padding=2)
# 池化层: 2x2最大池化
self.pool = nn.MaxPool2d(2, 2)
# 卷积层2: 32 -> 64通道,5x5卷积核
self.conv2 = nn.Conv2d(32, 64, kernel_size=5, padding=2)
# 全连接层1: 7x7x64 -> 1024
self.fc1 = nn.Linear(7 * 7 * 64, 1024)
# 全连接层2: 1024 -> 10 (输出层)
self.fc2 = nn.Linear(1024, num_classes)
# Dropout层
self.dropout = nn.Dropout(0.4)
def forward(self, x):
# 卷积层1 -> ReLU -> 池化
x = self.pool(F.relu(self.conv1(x))) # 输出: (32, 14, 14)
# 卷积层2 -> ReLU -> 池化
x = self.pool(F.relu(self.conv2(x))) # 输出: (64, 7, 7)
# 展平特征图
x = x.view(-1, 7 * 7 * 64)
# 全连接层1 -> ReLU -> Dropout
x = F.relu(self.fc1(x))
x = self.dropout(x)
# 输出层
x = self.fc2(x)
return x
# 初始化CNN模型
model_cnn = SimpleCNN()
model_cnn.to(device)
4.3 高级CNN模型(带批归一化)
让我们构建一个更复杂的CNN模型,包含批归一化和更多卷积层:
class AdvancedCNN(nn.Module):
def __init__(self, num_classes=10):
super(AdvancedCNN, self).__init__()
# 第一个卷积块
self.conv_block1 = nn.Sequential(
nn.Conv2d(1, 32, kernel_size=3, padding=1),
nn.BatchNorm2d(32), # 批归一化
nn.ReLU(inplace=True),
nn.Conv2d(32, 64, kernel_size=3, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.MaxPool2d(kernel_size=2, stride=2)
)
# 第二个卷积块
self.conv_block2 = nn.Sequential(
nn.Conv2d(64, 128, kernel_size=3, padding=1),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
nn.Conv2d(128, 128, kernel_size=3, padding=1),
nn.BatchNorm2d(128),
nn.ReLU(inplace=True),
nn.MaxPool2d(kernel_size=2, stride=2)
)
# 第三个卷积块
self.conv_block3 = nn.Sequential(
nn.Conv2d(128, 256, kernel_size=3, padding=1),
nn.BatchNorm2d(256),
nn.ReLU(inplace=True),
nn.MaxPool2d(kernel_size=2, stride=2)
)
# 全连接层
self.fc_layers = nn.Sequential(
nn.Dropout(0.5),
nn.Linear(256 * 3 * 3, 1024),
nn.BatchNorm1d(1024),
nn.ReLU(inplace=True),
nn.Dropout(0.5),
nn.Linear(1024, 512),
nn.BatchNorm1d(512),
nn.ReLU(inplace=True),
nn.Dropout(0.3),
nn.Linear(512, num_classes)
)
def forward(self, x):
x = self.conv_block1(x) # 输出: (64, 14, 14)
x = self.conv_block2(x) # 输出: (128, 7, 7)
x = self.conv_block3(x) # 输出: (256, 3, 3)
x = x.view(-1, 256 * 3 * 3) # 展平
x = self.fc_layers(x) # 全连接层处理
return x
# 初始化高级CNN模型
model_advanced = AdvancedCNN()
model_advanced.to(device)
4.4 模型架构可视化
5. 模型训练与优化
5.1 训练配置
首先,我们定义训练所需的超参数、损失函数和优化器:
# 超参数设置
LEARNING_RATE = 0.001
MOMENTUM = 0.9
WEIGHT_DECAY = 1e-4 # 权重衰减(正则化)
NUM_EPOCHS = 30
# 损失函数和优化器
criterion = nn.CrossEntropyLoss()
# 使用Adam优化器
optimizer_cnn = optim.Adam(
model_cnn.parameters(),
lr=LEARNING_RATE,
weight_decay=WEIGHT_DECAY
)
# 学习率调度器 - 动态调整学习率
scheduler = optim.lr_scheduler.ReduceLROnPlateau(
optimizer_cnn,
mode='max', # 当验证准确率不再提高时
factor=0.5, # 学习率减半
patience=3, # 等待3个epoch
verbose=True
)
5.2 训练循环实现
import time
import copy
import numpy as np
def train_model(model, criterion, optimizer, dataloader, num_epochs=25, scheduler=None):
"""
训练模型的通用函数
参数:
- model: 要训练的模型
- criterion: 损失函数
- optimizer: 优化器
- dataloader: 包含训练集和验证集的数据加载器字典
- num_epochs: 训练轮数
- scheduler: 学习率调度器
返回:
- model: 训练好的模型
- history: 包含训练过程中指标的字典
"""
# 初始化记录
history = {
'train_loss': [], 'train_acc': [],
'val_loss': [], 'val_acc': []
}
best_model_weights = None
best_acc = 0.0
start_time = time.time()
for epoch in range(num_epochs):
print(f'Epoch {epoch+1}/{num_epochs}')
print('-' * 50)
# 每个epoch包含训练和验证阶段
for phase in ['train', 'val']:
if phase == 'train':
model.train() # 设置为训练模式
dataloader = train_loader_augmented
else:
model.eval() # 设置为评估模式
dataloader = test_loader
running_loss = 0.0
running_corrects = 0
# 迭代数据
for inputs, labels in dataloader:
inputs = inputs.to(device) # 将数据移到GPU
labels = labels.to(device)
# 清零梯度
optimizer.zero_grad()
# 前向传播
with torch.set_grad_enabled(phase == 'train'):
outputs = model(inputs)
_, preds = torch.max(outputs, 1)
loss = criterion(outputs, labels)
# 训练阶段才有反向传播和优化
if phase == 'train':
loss.backward()
optimizer.step()
# 统计
batch_loss = loss.item() * inputs.size(0)
running_loss += batch_loss
running_corrects += torch.sum(preds == labels.data)
# 计算每个epoch的损失和准确率
epoch_loss = running_loss / len(dataloader.dataset)
epoch_acc = running_corrects.double() / len(dataloader.dataset)
print(f'{phase} Loss: {epoch_loss:.4f} Acc: {epoch_acc:.4f}')
# 保存历史记录
history[f'{phase}_loss'].append(epoch_loss)
history[f'{phase}_acc'].append(epoch_acc.item())
# 如果是验证阶段且准确率更高,则保存最佳模型权重
if phase == 'val' and epoch_acc > best_acc:
best_acc = epoch_acc
best_model_weights = copy.deepcopy(model.state_dict())
# 保存模型
torch.save(model.state_dict(), 'best_model.pth')
# 更新学习率调度器
if scheduler:
scheduler.step(epoch_acc)
print()
# 计算训练总时间
time_elapsed = time.time() - start_time
print(f'Training complete in {time_elapsed//60:.0f}m {time_elapsed%60:.0f}s')
print(f'Best val Acc: {best_acc:.4f}')
# 加载最佳模型权重
model.load_state_dict(best_model_weights)
return model, history
5.3 训练模型
现在我们可以开始训练模型了。这里我们以高级CNN模型为例:
# 准备数据加载器字典
dataloaders = {
'train': train_loader_augmented,
'val': test_loader
}
# 训练高级CNN模型
trained_model, history = train_model(
model=model_advanced,
criterion=criterion,
optimizer=optimizer_cnn,
dataloader=dataloaders,
num_epochs=NUM_EPOCHS,
scheduler=scheduler
)
5.4 多模型比较训练
为了比较不同模型的性能,我们可以训练多个模型并比较结果:
# 定义要训练的模型列表
models = {
'SimpleMLP': model_mlp,
'SimpleCNN': model_cnn,
'AdvancedCNN': model_advanced
}
# 为每个模型定义优化器
optimizers = {
'SimpleMLP': optim.Adam(model_mlp.parameters(), lr=0.001, weight_decay=1e-4),
'SimpleCNN': optim.Adam(model_cnn.parameters(), lr=0.001, weight_decay=1e-4),
'AdvancedCNN': optim.Adam(model_advanced.parameters(), lr=0.001, weight_decay=1e-4)
}
# 存储所有模型的训练历史
all_histories = {}
# 训练每个模型
for name, model in models.items():
print(f"\n\nTraining {name}...")
model.to(device)
# 创建调度器
scheduler = optim.lr_scheduler.ReduceLROnPlateau(
optimizers[name], mode='max', factor=0.5, patience=3, verbose=True
)
# 训练模型
_, history = train_model(
model=model,
criterion=criterion,
optimizer=optimizers[name],
dataloader=dataloaders,
num_epochs=20,
scheduler=scheduler
)
all_histories[name] = history
6. 结果可视化与分析
6.1 绘制训练曲线
import matplotlib.pyplot as plt
def plot_training_history(history, title):
"""绘制训练损失和准确率曲线"""
plt.figure(figsize=(12, 4))
# 绘制损失曲线
plt.subplot(1, 2, 1)
plt.plot(history['train_loss'], label='Training Loss')
plt.plot(history['val_loss'], label='Validation Loss')
plt.title(f'{title} - Loss Curves')
plt.xlabel('Epoch')
plt.ylabel('Loss')
plt.legend()
# 绘制准确率曲线
plt.subplot(1, 2, 2)
plt.plot(history['train_acc'], label='Training Accuracy')
plt.plot(history['val_acc'], label='Validation Accuracy')
plt.title(f'{title} - Accuracy Curves')
plt.xlabel('Epoch')
plt.ylabel('Accuracy')
plt.legend()
plt.tight_layout()
plt.savefig(f'{title}_training_curves.png')
plt.show()
# 绘制高级CNN模型的训练曲线
plot_training_history(history, 'Advanced CNN')
6.2 多模型性能比较
def compare_models(histories):
"""比较多个模型的性能"""
plt.figure(figsize=(12, 5))
# 比较验证准确率
plt.subplot(1, 2, 1)
for name, history in histories.items():
plt.plot(history['val_acc'], label=name)
plt.title('Validation Accuracy Comparison')
plt.xlabel('Epoch')
plt.ylabel('Accuracy')
plt.legend()
# 比较验证损失
plt.subplot(1, 2, 2)
for name, history in histories.items():
plt.plot(history['val_loss'], label=name)
plt.title('Validation Loss Comparison')
plt.xlabel('Epoch')
plt.ylabel('Loss')
plt.legend()
plt.tight_layout()
plt.savefig('model_comparison.png')
plt.show()
# 比较所有模型
compare_models(all_histories)
6.3 混淆矩阵分析
from sklearn.metrics import confusion_matrix, classification_report
import seaborn as sns
def plot_confusion_matrix(model, dataloader):
"""绘制混淆矩阵"""
model.eval()
all_preds = []
all_labels = []
with torch.no_grad():
for inputs, labels in dataloader:
inputs = inputs.to(device)
outputs = model(inputs)
_, preds = torch.max(outputs, 1)
all_preds.extend(preds.cpu().numpy())
all_labels.extend(labels.numpy())
# 计算混淆矩阵
cm = confusion_matrix(all_labels, all_preds)
# 绘制混淆矩阵
plt.figure(figsize=(12, 10))
sns.heatmap(
cm,
annot=True,
fmt='d',
cmap='Blues',
xticklabels=class_names,
yticklabels=class_names
)
plt.xlabel('Predicted Label')
plt.ylabel('True Label')
plt.title('Confusion Matrix')
plt.tight_layout()
plt.savefig('confusion_matrix.png')
plt.show()
# 打印分类报告
print("Classification Report:")
print(classification_report(all_labels, all_preds, target_names=class_names))
# 分析最佳模型
plot_confusion_matrix(trained_model, test_loader)
6.4 错误案例分析
def analyze_errors(model, dataloader, num_examples=25):
"""分析模型预测错误的案例"""
model.eval()
errors = []
with torch.no_grad():
for inputs, labels in dataloader:
inputs = inputs.to(device)
outputs = model(inputs)
_, preds = torch.max(outputs, 1)
# 找出预测错误的样本
for i, (pred, label) in enumerate(zip(preds, labels)):
if pred != label:
errors.append({
'image': inputs[i].cpu().squeeze().numpy(),
'true_label': label.item(),
'pred_label': pred.item(),
'confidence': F.softmax(outputs[i], dim=0)[pred].item()
})
if len(errors) >= num_examples:
break
# 显示错误案例
fig, axes = plt.subplots(5, 5, figsize=(15, 15))
axes = axes.flatten()
for i, error in enumerate(errors[:num_examples]):
axes[i].imshow(error['image'], cmap='gray')
axes[i].set_title(
f"True: {class_names[error['true_label']]}\n"
f"Pred: {class_names[error['pred_label']]}\n"
f"Confidence: {error['confidence']:.2f}"
)
axes[i].axis('off')
plt.tight_layout()
plt.savefig('error_analysis.png')
plt.show()
# 分析错误案例
analyze_errors(trained_model, test_loader)
7. 模型优化与调参
7.1 学习率搜索
def find_optimal_lr(model, criterion, optimizer, train_loader, init_lr=1e-7, max_lr=1, num_iter=100):
"""使用循环学习率寻找最优学习率"""
model.train()
losses = []
lrs = []
best_loss = float('inf')
# 设置初始学习率
lr_scheduler = optim.lr_scheduler.LinearLR(
optimizer, start_factor=init_lr/max_lr, total_iters=num_iter
)
for iter_num, (inputs, labels) in enumerate(train_loader):
if iter_num >= num_iter:
break
inputs = inputs.to(device)
labels = labels.to(device)
optimizer.zero_grad()
outputs = model(inputs)
loss = criterion(outputs, labels)
# 如果损失开始爆炸,停止搜索
if loss.item() > 4 * best_loss:
break
if loss.item() < best_loss and iter_num > 0:
best_loss = loss.item()
losses.append(loss.item())
lrs.append(lr_scheduler.get_last_lr()[0])
loss.backward()
optimizer.step()
lr_scheduler.step()
# 绘制学习率-损失曲线
plt.figure(figsize=(10, 6))
plt.plot(lrs, losses)
plt.xscale('log')
plt.xlabel('Learning Rate (log scale)')
plt.ylabel('Loss')
plt.title('Learning Rate vs. Loss')
plt.grid(True, which="both", ls="-")
plt.savefig('lr_finder.png')
plt.show()
return lrs, losses
# 寻找最优学习率
lrs, losses = find_optimal_lr(model_advanced, criterion, optimizer_cnn, train_loader)
7.2 正则化技术比较
# 比较不同正则化技术的效果
regularization_comparison = {
'No Regularization': {'loss': [], 'acc': []},
'L2 Regularization': {'loss': [], 'acc': []},
'Dropout': {'loss': [], 'acc': []},
'L2 + Dropout': {'loss': [], 'acc': []}
}
# 这里仅展示框架,实际运行需要分别训练不同正则化配置的模型
7.3 超参数调优
from sklearn.model_selection import ParameterGrid
# 定义超参数网格
param_grid = {
'learning_rate': [0.001, 0.005, 0.01],
'batch_size': [32, 64, 128],
'dropout_rate': [0.3, 0.4, 0.5],
'weight_decay': [1e-4, 1e-3]
}
# 网格搜索
best_params = None
best_val_acc = 0.0
for params in ParameterGrid(param_grid):
print(f"Testing parameters: {params}")
# 创建模型
model = AdvancedCNN()
model.to(device)
# 创建优化器
optimizer = optim.Adam(
model.parameters(),
lr=params['learning_rate'],
weight_decay=params['weight_decay']
)
# 调整模型的dropout率
for module in model.modules():
if isinstance(module, nn.Dropout):
module.p = params['dropout_rate']
# 创建数据加载器
batch_size = params['batch_size']
temp_loader = DataLoader(
train_dataset_augmented,
batch_size=batch_size,
shuffle=True,
num_workers=2
)
# 训练模型(使用较少的epoch进行快速评估)
_, temp_history = train_model(
model, criterion, optimizer, {'train': temp_loader, 'val': test_loader},
num_epochs=10, scheduler=None
)
# 获取验证准确率
current_val_acc = max(temp_history['val_acc'])
print(f"Validation Accuracy: {current_val_acc:.4f}")
# 更新最佳参数
if current_val_acc > best_val_acc:
best_val_acc = current_val_acc
best_params = params
print(f"New best parameters: {best_params}, Accuracy: {best_val_acc:.4f}")
print(f"Best parameters found: {best_params} with accuracy: {best_val_acc:.4f}")
8. 模型部署与预测
8.1 保存和加载模型
# 保存完整模型
torch.save(trained_model, 'fashion_mnist_cnn.pth
更多推荐



所有评论(0)