import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from data_transform import load_mnist, load_and_preview_mnist import matplotlib.pyplot as plt import torchvision.utils as vutils # 定义卷积神经网络模型 class Model(nn.Module): def __init__(self): super(Model, self).__init__() self.conv1 = nn.Sequential( nn.Conv2d(1, 64, kernel_size=3, stride=2, padding=1), nn.ReLU(), nn.Conv2d(64, 128, kernel_size=3, stride=2, padding=1), nn.ReLU(), # nn.MaxPool2d(stride=2, kernel_size=2) ) self.dense = nn.Sequential( nn.Linear(7 * 7 * 128, 512), nn.ReLU(), nn.Dropout(p=0.8), nn.Linear(512, 10) ) def forward(self, x): x = self.conv1(x) # 卷积处理 x = x.view(-1, 7 * 7 * 128) # 对参数实行扁平化处理 x = self.dense(x) return x # 加载和预览数据集 def main(): # 参数设置 data_size = 1000 batch_size = 4 seed = 42 root = '../data' epochs_n = 5 learning_rate = 0.001 # 加载数据加载器 train_loader = load_mnist(data_size=data_size, batch_size=batch_size, seed=seed, root=root, train=True) test_loader = load_mnist(data_size=data_size, batch_size=batch_size, seed=seed, root=root, train=False) # 预览数据集(可选) load_and_preview_mnist(data_size=data_size, batch_size=batch_size, seed=seed, root=root) # 初始化模型 model = Model() cost = nn.CrossEntropyLoss() optimizer = optim.Adam(model.parameters(), lr=learning_rate) # 训练模型 for epoch in range(epochs_n): model.train() running_loss = 0.0 running_correct = 0 print(f"Epoch {epoch + 1}/{epochs_n}") print("-" * 10) for data in train_loader: images, labels = data optimizer.zero_grad() outputs = model(images) _, preds = torch.max(outputs.data, 1) loss = cost(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() running_correct += (preds == labels).sum().item() epoch_loss = running_loss / len(train_loader.dataset) epoch_acc = 100 * running_correct / len(train_loader.dataset) print(f"Loss: {epoch_loss:.4f}, Train Accuracy: {epoch_acc:.4f}%") # 测试模型 model.eval() testing_correct = 0 with torch.no_grad(): for data in test_loader: images, labels = data outputs = model(images) _, preds = torch.max(outputs.data, 1) testing_correct += (preds == labels).sum().item() test_acc = 100 * testing_correct / len(test_loader.dataset) print(f"Test Accuracy: {test_acc:.4f}%") # 测试数据可视化 X_test, y_test = next(iter(test_loader)) with torch.no_grad(): pred = model(X_test) _, predicted_labels = torch.max(pred, 1) print("Predicted Labels:", [i.item() for i in predicted_labels]) print("True Labels:", [i.item() for i in y_test]) # 显示图像 img_grid = vutils.make_grid(X_test, nrow=batch_size // 4) img = img_grid.numpy().transpose(1, 2, 0) mean = [0.5] std = [0.5] img = img * std + mean plt.imshow(img, cmap='gray') plt.axis('off') plt.show() if __name__ == '__main__': main()