Spaces:
Sleeping
Sleeping
| import torch | |
| import torch.nn as nn | |
| import torch.optim as optim | |
| import torchvision | |
| import torchvision.transforms as transforms | |
| # ----------------- TRANSFORMS ----------------- | |
| transform = transforms.Compose([ | |
| transforms.ToTensor(), | |
| transforms.Normalize((0.5,), (0.5,)) | |
| ]) | |
| # ----------------- LOAD DATASET ----------------- | |
| train_set = torchvision.datasets.CIFAR10( | |
| root="./data", train=True, download=True, transform=transform | |
| ) | |
| test_set = torchvision.datasets.CIFAR10( | |
| root="./data", train=False, download=True, transform=transform | |
| ) | |
| train_loader = torch.utils.data.DataLoader(train_set, batch_size=64, shuffle=True) | |
| test_loader = torch.utils.data.DataLoader(test_set, batch_size=64, shuffle=False) | |
| # ----------------- BUILD CNN MODEL ----------------- | |
| class CNN(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| self.conv_layer = nn.Sequential( | |
| nn.Conv2d(3, 32, kernel_size=3, padding=1), | |
| nn.ReLU(), | |
| nn.MaxPool2d(2, 2), | |
| nn.Conv2d(32, 64, kernel_size=3, padding=1), | |
| nn.ReLU(), | |
| nn.MaxPool2d(2, 2) | |
| ) | |
| self.fc_layer = nn.Sequential( | |
| nn.Linear(64 * 8 * 8, 256), | |
| nn.ReLU(), | |
| nn.Linear(256, 10) | |
| ) | |
| def forward(self, x): | |
| x = self.conv_layer(x) | |
| x = x.view(x.size(0), -1) | |
| x = self.fc_layer(x) | |
| return x | |
| model = CNN() | |
| criterion = nn.CrossEntropyLoss() | |
| optimizer = optim.Adam(model.parameters(), lr=0.001) | |
| # ----------------- TRAIN LOOP ----------------- | |
| for epoch in range(5): | |
| running_loss = 0.0 | |
| for images, labels in train_loader: | |
| optimizer.zero_grad() | |
| outputs = model(images) | |
| loss = criterion(outputs, labels) | |
| loss.backward() | |
| optimizer.step() | |
| running_loss += loss.item() | |
| print(f"Epoch {epoch+1}, Loss: {running_loss/len(train_loader)}") | |
| # ----------------- SAVE MODEL ----------------- | |
| torch.save(model.state_dict(), "model.pth") | |
| print("Model saved as model.pth") | |