Here is a code snippet that demonstrates loading the MNIST dataset, normalizing the data, splitting it into training, validation, and testing sets, building a CNN network, training the network, and evaluating its performance on the testing set:

import torch
import torch.nn as nn
import torch.optim as optim
import torchvision.datasets as datasets
import torchvision.transforms as transforms
import matplotlib.pyplot as plt

# Load MNIST dataset
train_dataset = datasets.MNIST(root='data/', train=True, transform=transforms.ToTensor(), download=True)
test_dataset = datasets.MNIST(root='data/', train=False, transform=transforms.ToTensor())

# Normalize the data
mean = train_dataset.data.float().mean() / 255
std = train_dataset.data.float().std() / 255
train_dataset.data = (train_dataset.data.float() / 255 - mean) / std
test_dataset.data = (test_dataset.data.float() / 255 - mean) / std

# Split the data into training, validation, and testing sets
train_size = int(0.8 * len(train_dataset))
val_size = len(train_dataset) - train_size
train_dataset, val_dataset = torch.utils.data.random_split(train_dataset, [train_size, val_size])

# Define the CNN network
class CNN(nn.Module):
    def __init__(self):
        super(CNN, self).__init__()
        self.conv1 = nn.Conv2d(1, 32, kernel_size=3, stride=1, padding=1)
        self.relu = nn.ReLU()
        self.maxpool = nn.MaxPool2d(kernel_size=2, stride=2)
        self.conv2 = nn.Conv2d(32, 64, kernel_size=3, stride=1, padding=1)
        self.fc = nn.Linear(64 * 7 * 7, 10)
    
    def forward(self, x):
        x = self.conv1(x)
        x = self.relu(x)
        x = self.maxpool(x)
        x = self.conv2(x)
        x = self.relu(x)
        x = self.maxpool(x)
        x = x.view(x.size(0), -1)
        x = self.fc(x)
        return x

# Create the CNN model and define the loss function and optimizer
model = CNN()
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)

# Training loop
num_epochs = 10
train_losses = []
val_losses = []
train_accs = []
val_accs = []

for epoch in range(num_epochs):
    train_loss = 0.0
    val_loss = 0.0
    train_acc = 0.0
    val_acc = 0.0

    # Training
    model.train()
    for images, labels in train_dataset:
        optimizer.zero_grad()
        outputs = model(images.unsqueeze(0))
        loss = criterion(outputs, labels.unsqueeze(0))
        loss.backward()
        optimizer.step()
        train_loss += loss.item()
        _, predicted = torch.max(outputs.data, 1)
        train_acc += (predicted == labels).sum().item()

    # Validation
    model.eval()
    with torch.no_grad():
        for images, labels in val_dataset:
            outputs = model(images.unsqueeze(0))
            loss = criterion(outputs, labels.unsqueeze(0))
            val_loss += loss.item()
            _, predicted = torch.max(outputs.data, 1)
            val_acc += (predicted == labels).sum().item()

    # Calculate average loss and accuracy
    train_loss /= len(train_dataset)
    val_loss /= len(val_dataset)
    train_acc /= len(train_dataset)
    val_acc /= len(val_dataset)
    
    # Print progress
    print('Epoch [{}/{}], Train Loss: {:.4f}, Val Loss: {:.4f}, Train Acc: {:.4f}, Val Acc: {:.4f}' 
          .format(epoch+1, num_epochs, train_loss, val_loss, train_acc, val_acc))
    
    # Save loss and accuracy for plotting
    train_losses.append(train_loss)
    val_losses.append(val_loss)
    train_accs.append(train_acc)
    val_accs.append(val_acc)

# Plot training and validation loss
plt.plot(train_losses, label='Train Loss')
plt.plot(val_losses, label='Val Loss')
plt.xlabel('Epoch')
plt.ylabel('Loss')
plt.legend()
plt.show()

# Plot training and validation accuracy
plt.plot(train_accs, label='Train Acc')
plt.plot(val_accs, label='Val Acc')
plt.xlabel('Epoch')
plt.ylabel('Accuracy')
plt.legend()
plt.show()

# Testing
model.eval()
test_acc = 0.0
with torch.no_grad():
    for images, labels in test_dataset:
        outputs = model(images.unsqueeze(0))
        _, predicted = torch.max(outputs.data, 1)
        test_acc += (predicted == labels).sum().item()

test_acc /= len(test_dataset)
print('Testing Accuracy: {:.4f}'.format(test_acc))

Make sure you have PyTorch and torchvision installed before running the code.

MNIST Handwritten Digit Classification using CNN in PyTorch

原文地址: https://www.cveoy.top/t/topic/lAoL 著作权归作者所有。请勿转载和采集!

免费AI点我,无需注册和登录