MNIST Handwritten Digit Classification using CNN with PyTorch
Here is the code to load the MNIST dataset, normalize the data, split it into training, validation, and testing sets, build a CNN network, and train it using PyTorch:
import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
import matplotlib.pyplot as plt
# Set random seed for reproducibility
torch.manual_seed(42)
# Define batch size
batch_size = 64
# Load MNIST dataset and apply transformations
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,))
])
train_dataset = datasets.MNIST(root='./data', train=True, transform=transform, download=True)
test_dataset = datasets.MNIST(root='./data', train=False, transform=transform)
# Split 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])
# Create data loaders
train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=batch_size, shuffle=True)
val_loader = torch.utils.data.DataLoader(val_dataset, batch_size=batch_size, shuffle=True)
test_loader = torch.utils.data.DataLoader(test_dataset, batch_size=batch_size, shuffle=True)
# Define CNN model
class Net(nn.Module):
def __init__(self):
super(Net, self).__init__()
self.conv1 = nn.Conv2d(1, 32, kernel_size=3)
self.conv2 = nn.Conv2d(32, 64, kernel_size=3)
self.dropout1 = nn.Dropout2d(0.25)
self.dropout2 = nn.Dropout2d(0.5)
self.fc1 = nn.Linear(64 * 12 * 12, 128)
self.fc2 = nn.Linear(128, 10)
def forward(self, x):
x = self.conv1(x)
x = nn.functional.relu(x)
x = self.conv2(x)
x = nn.functional.relu(x)
x = nn.functional.max_pool2d(x, 2)
x = self.dropout1(x)
x = torch.flatten(x, 1)
x = self.fc1(x)
x = nn.functional.relu(x)
x = self.dropout2(x)
x = self.fc2(x)
return x
# Create an instance of the CNN model
model = Net()
# Define loss function and optimizer
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=0.001)
# Train the model
epochs = 10
train_losses = []
val_losses = []
train_accs = []
val_accs = []
for epoch in range(epochs):
train_loss = 0.0
val_loss = 0.0
train_acc = 0.0
val_acc = 0.0
# Training phase
model.train()
for images, labels in train_loader:
optimizer.zero_grad()
output = model(images)
loss = criterion(output, labels)
loss.backward()
optimizer.step()
train_loss += loss.item() * images.size(0)
_, predicted = torch.max(output.data, 1)
train_acc += (predicted == labels).sum().item()
# Validation phase
model.eval()
with torch.no_grad():
for images, labels in val_loader:
output = model(images)
loss = criterion(output, labels)
val_loss += loss.item() * images.size(0)
_, predicted = torch.max(output.data, 1)
val_acc += (predicted == labels).sum().item()
# Calculate average loss and accuracy
train_loss /= len(train_loader.dataset)
val_loss /= len(val_loader.dataset)
train_acc /= len(train_loader.dataset)
val_acc /= len(val_loader.dataset)
train_losses.append(train_loss)
val_losses.append(val_loss)
train_accs.append(train_acc)
val_accs.append(val_acc)
print(f'Epoch {epoch+1}/{epochs}, Train Loss: {train_loss:.4f}, Val Loss: {val_loss:.4f}, Train Acc: {train_acc:.4f}, Val Acc: {val_acc:.4f}')
# Test the model
model.eval()
test_acc = 0.0
with torch.no_grad():
for images, labels in test_loader:
output = model(images)
_, predicted = torch.max(output.data, 1)
test_acc += (predicted == labels).sum().item()
test_acc /= len(test_loader.dataset)
print(f'Testing Accuracy: {test_acc:.4f}')
# Plot training and validation losses
plt.figure()
plt.plot(range(1, epochs+1), train_losses, label='Train')
plt.plot(range(1, epochs+1), val_losses, label='Validation')
plt.xlabel('Epochs')
plt.ylabel('Loss')
plt.legend()
plt.show()
# Plot training and validation accuracies
plt.figure()
plt.plot(range(1, epochs+1), train_accs, label='Train')
plt.plot(range(1, epochs+1), val_accs, label='Validation')
plt.xlabel('Epochs')
plt.ylabel('Accuracy')
plt.legend()
plt.show()
This code will print the testing accuracy and plot the training loss, validation loss, training accuracy, and validation accuracy as a function of epochs.
原文地址: https://www.cveoy.top/t/topic/h5GX 著作权归作者所有。请勿转载和采集!