""" Fashion-MNIST CNN Model Similar architecture to MNIST CNN but trained on fashion data """ import torch import torch.nn as nn import torch.nn.functional as F class FashionCNN(nn.Module): """ CNN for Fashion-MNIST classification ~250K parameters, optimized for more complex features """ def __init__(self, num_classes=10, dropout_rate=0.3): super(FashionCNN, self).__init__() # Enhanced feature extraction for fashion items self.conv1 = nn.Conv2d(1, 64, kernel_size=3, padding=1) self.bn1 = nn.BatchNorm2d(64) self.conv2 = nn.Conv2d(64, 128, kernel_size=3, padding=1) self.bn2 = nn.BatchNorm2d(128) self.conv3 = nn.Conv2d(128, 256, kernel_size=3, padding=1) self.bn3 = nn.BatchNorm2d(256) # Fully connected layers self.fc1 = nn.Linear(256 * 3 * 3, 256) self.dropout1 = nn.Dropout(dropout_rate) self.fc2 = nn.Linear(256, 128) self.dropout2 = nn.Dropout(dropout_rate) self.fc3 = nn.Linear(128, num_classes) # Initialize weights self._initialize_weights() def _initialize_weights(self): """Initialize weights using Kaiming initialization for ReLU""" for m in self.modules(): if isinstance(m, nn.Conv2d): nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu') if m.bias is not None: nn.init.zeros_(m.bias) elif isinstance(m, nn.BatchNorm2d): nn.init.ones_(m.weight) nn.init.zeros_(m.bias) elif isinstance(m, nn.Linear): nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu') nn.init.zeros_(m.bias) def forward(self, x): # Conv block 1 x = F.relu(self.bn1(self.conv1(x))) x = F.max_pool2d(x, 2) # Conv block 2 x = F.relu(self.bn2(self.conv2(x))) x = F.max_pool2d(x, 2) # Conv block 3 x = F.relu(self.bn3(self.conv3(x))) x = F.max_pool2d(x, 2) # Flatten x = x.view(x.size(0), -1) # Fully connected x = F.relu(self.fc1(x)) x = self.dropout1(x) x = F.relu(self.fc2(x)) x = self.dropout2(x) x = self.fc3(x) return x # Factory function def create_fashion_cnn(num_classes=10, dropout_rate=0.3): """Factory function to create Fashion-MNIST CNN""" return FashionCNN(num_classes=num_classes, dropout_rate=dropout_rate)