import torch import random import torchvision import torch.nn as nn import torch.optim as optim import matplotlib.pyplot as plt from torchvision.models import resnet18 import torchvision.transforms as transforms from torch.utils.data import DataLoader, random_split #-------------------------------------------------- SEED EVERYTHING -------------------------------------------------- def seed_everything(seed): '''Set the seed for all random generators to be the provided seed''' random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) #-------------------------------------------------- COMPUTE MEAN + STD -------------------------------------------------- def compute_mean_std(dataset): loader = DataLoader(dataset, batch_size=512, shuffle=False) mean = 0. std = 0. total = 0 for images, _ in loader: images = images.view(images.size(0), images.size(1), -1) # (B, C, H*W) mean += images.mean(2).sum(0) std += images.std(2).sum(0) total += images.size(0) mean /= total std /= total return mean, std #-------------------------------------------------- DATASET + PREPROCESSING -------------------------------------------------- def get_dataloaders(batch_size=128, subset_size=20000): base_transform = transforms.Compose([transforms.ToTensor()]) dataset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=base_transform) if subset_size: dataset, _ = random_split(dataset, [subset_size, len(dataset)-subset_size]) mean, std = compute_mean_std(dataset) mean, std = tuple(mean.tolist()), tuple(std.tolist()) print("Computed mean:", mean) print("Computed std:", std) train_transform = transforms.Compose([ transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.RandomCrop(32, padding=4), transforms.ToTensor(), transforms.Normalize(mean, std)]) test_transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean, std)]) dataset = torchvision.datasets.CIFAR10(root='./data', train=True, download=False, transform=train_transform) if subset_size: dataset, _ = random_split(dataset, [subset_size, len(dataset)-subset_size]) train_size = int(0.8 * len(dataset)) val_size = len(dataset) - train_size train_set, val_set = random_split(dataset, [train_size, val_size]) val_set.dataset.transform = test_transform test_set = torchvision.datasets.CIFAR10(root='./data', train=False, download=True, transform=test_transform) train_loader = DataLoader(train_set, batch_size=batch_size, shuffle=True) val_loader = DataLoader(val_set, batch_size=batch_size) test_loader = DataLoader(test_set, batch_size=batch_size) return train_loader, val_loader, test_loader #-------------------------------------------------- CNN BLOCK -------------------------------------------------- class ConvBlock(nn.Module): def __init__(self, in_channels, out_channels, num_layers, stride=1): super().__init__() layers = [] for i in range(num_layers): s = stride if i == 0 else 1 layers.append(nn.Conv2d(in_channels if i == 0 else out_channels, out_channels, kernel_size=3, stride=s, padding=1)) layers.append(nn.BatchNorm2d(out_channels)) layers.append(nn.ReLU(inplace=True)) self.block = nn.Sequential(*layers) def forward(self, x): return self.block(x) #-------------------------------------------------- SIMPLE CNN -------------------------------------------------- class SimpleCNN(nn.Module): def __init__(self, depth=2): super().__init__() self.module1 = ConvBlock(3, 32, depth, stride=1) self.module2 = ConvBlock(32, 64, depth, stride=2) self.module3 = ConvBlock(64, 128, depth, stride=2) self.global_pool = nn.AdaptiveAvgPool2d((1, 1)) self.fc = nn.Linear(128, 10) self._initialize_weights() def forward(self, x): x = self.module1(x) x = self.module2(x) x = self.module3(x) x = self.global_pool(x) x = x.view(x.size(0), -1) return self.fc(x) def _initialize_weights(self): for m in self.modules(): if isinstance(m, nn.Conv2d): nn.init.kaiming_normal_(m.weight) #-------------------------------------------------- TRAINING LOOP -------------------------------------------------- def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss = 0 for inputs, targets in loader: inputs, targets = inputs.to(device), targets.to(device) optimizer.zero_grad() outputs = model(inputs) loss = criterion(outputs, targets) loss.backward() optimizer.step() total_loss += loss.item() return total_loss / len(loader) #-------------------------------------------------- EVALUATION LOOP -------------------------------------------------- def evaluate(model, loader, criterion, device): model.eval() total_loss = 0 correct = 0 total = 0 with torch.no_grad(): for inputs, targets in loader: inputs, targets = inputs.to(device), targets.to(device) outputs = model(inputs) loss = criterion(outputs, targets) total_loss += loss.item() _, predicted = outputs.max(1) total += targets.size(0) correct += predicted.eq(targets).sum().item() accuracy = 100. * correct / total return total_loss / len(loader), accuracy #-------------------------------------------------- OPTIMIZER SELECTION -------------------------------------------------- def get_optimizer(model, name="sgd", lr=0.01): match name.lower(): case "sgd": return optim.SGD(model.parameters(), lr=lr, momentum=0.9) case "adam": return optim.Adam(model.parameters(), lr=lr) case _: raise ValueError("Unknown optimizer") #-------------------------------------------------- TRAINING -------------------------------------------------- def run_training(model, train_loader, val_loader, optimizer_name="sgd", epochs=10, device="cpu", debug=True): model = model.to(device) optimizer = get_optimizer(model, optimizer_name) criterion = nn.CrossEntropyLoss() train_losses, val_losses, val_accs = [], [], [] for epoch in range(epochs): train_loss = train_one_epoch(model, train_loader, optimizer, criterion, device) val_loss, val_acc = evaluate(model, val_loader, criterion, device) train_losses.append(train_loss) val_losses.append(val_loss) val_accs.append(val_acc) if debug: print(f"[{epoch+1}/{epochs}] Train Loss: {train_loss:.4f} | Val Acc: {val_acc:.2f}%") return train_losses, val_losses, val_accs #-------------------------------------------------- RESNET18 -------------------------------------------------- def get_resnet(): model = resnet18(weights=None) model.fc = nn.Linear(model.fc.in_features, 10) return model #-------------------------------------------------- EXPERIMENTS -------------------------------------------------- def run_experiment_A(train_loader, val_loader, debug=False): print("========================= EXPERIMENT A: DEPTH COMPARISON =========================") device = "cuda" if torch.cuda.is_available() else "cpu" model = SimpleCNN(depth=2) train_s, val_s, acc_s = run_training(model, train_loader, val_loader, optimizer_name="sgd", device=device, debug=debug) model = SimpleCNN(depth=3) train_d, val_d, acc_d = run_training(model, train_loader, val_loader, optimizer_name="sgd", device=device, debug=debug) #------------------------- ACCURACY ------------------------- plt.figure() plt.plot(acc_s, label="Shallow") plt.plot(acc_d, label="Deep") plt.xlabel("Epoch") plt.ylabel("Validation Accuracy") plt.title("Depth Comparison (Accuracy)") plt.legend() plt.show() #------------------------- LOSS ------------------------- plt.figure() plt.plot(train_s, label="Train Shallow") plt.plot(val_s, label="Val Shallow") plt.plot(train_d, label="Train Deep") plt.plot(val_d, label="Val Deep") plt.xlabel("Epoch") plt.ylabel("Loss") plt.title("Depth Comparison (Loss)") plt.legend() plt.show() print("==================================================================================") def run_experiment_B(train_loader, val_loader, debug=False): print("========================= EXPERIMENT B: ARCHITECTURE COMPARISON =========================") device = "cuda" if torch.cuda.is_available() else "cpu" cnn_model = SimpleCNN(depth=3) train_c, val_c, acc_c = run_training(cnn_model, train_loader, val_loader, device=device, debug=debug) resnet_model = get_resnet() train_r, val_r, acc_r = run_training(resnet_model, train_loader, val_loader, device=device, debug=debug) #------------------------- ACCURACY ------------------------- plt.figure() plt.plot(acc_c, label="SimpleCNN-Deep") plt.plot(acc_r, label="ResNet18") plt.xlabel("Epoch") plt.ylabel("Validation Accuracy") plt.title("Architecture Comparison (Accuracy)") plt.legend() plt.show() #------------------------- LOSS ------------------------- plt.figure() plt.plot(train_c, label="Train CNN") plt.plot(val_c, label="Val CNN") plt.plot(train_r, label="Train ResNet") plt.plot(val_r, label="Val ResNet") plt.xlabel("Epoch") plt.ylabel("Loss") plt.title("Architecture Comparison (Loss)") plt.legend() plt.show() print("=========================================================================================") def run_experiment_C(train_loader, val_loader, debug=False): print("========================= EXPERIMENT C: OPTIMIZER COMPARISON =========================") device = "cuda" if torch.cuda.is_available() else "cpu" model_sgd = SimpleCNN(depth=3) train_sgd, _, acc_sgd = run_training(model_sgd, train_loader, val_loader, optimizer_name="sgd", device=device, debug=debug) model_adam = SimpleCNN(depth=3) train_adam, _, acc_adam = run_training(model_adam, train_loader, val_loader, optimizer_name="adam", device=device, debug=debug) #------------------------- LOSS ------------------------- plt.figure() plt.plot(train_sgd, label="SGD") plt.plot(train_adam, label="Adam") plt.xlabel("Epoch") plt.ylabel("Training Loss") plt.title("Optimizer Comparison") plt.legend() plt.show() #------------------------- ACCURACY ------------------------- plt.figure() plt.plot(acc_sgd, label="SGD") plt.plot(acc_adam, label="Adam") plt.xlabel("Epoch") plt.ylabel("Accuracy") plt.title("Optimizer Comparison (Accuracy)") plt.legend() plt.show() print("=====================================================================================") #-------------------------------------------------- MAIN -------------------------------------------------- if __name__ == "__main__": seed_everything(2026) debug = False train_loader, val_loader, _ = get_dataloaders() run_experiment_A(train_loader, val_loader, debug) run_experiment_B(train_loader, val_loader, debug) run_experiment_C(train_loader, val_loader, debug)