import mlx.core as mx import mlx.nn as nn class ChessPieceCNN(nn.Module): def __init__(self, num_classes=13): super().__init__() self.conv1 = nn.Conv2d(3, 32, 3, padding=1) self.bn1 = nn.BatchNorm(32) self.conv2 = nn.Conv2d(32, 64, 3, padding=1) self.bn2 = nn.BatchNorm(64) self.conv3 = nn.Conv2d(64, 128, 3, padding=1) self.bn3 = nn.BatchNorm(128) self.pool = nn.MaxPool2d(2, 2) self.fc1 = nn.Linear(128 * 4 * 4, 256) self.fc2 = nn.Linear(256, num_classes) def __call__(self, x): # x input is (B, H, W, C) x = nn.relu(self.bn1(self.conv1(x))) x = self.pool(x) x = nn.relu(self.bn2(self.conv2(x))) x = self.pool(x) x = nn.relu(self.bn3(self.conv3(x))) x = self.pool(x) x = mx.flatten(x, start_axis=1) x = nn.relu(self.fc1(x)) x = self.fc2(x) return x class ChessPieceCNN64(nn.Module): """Stronger per-square classifier for 64x64 crops. Still classifies each square independently into the same 13 labels, but keeps 4x more input pixels than the original 32x32 model and uses two convs per stage so piece color/outline detail survives much better. """ def __init__(self, num_classes=13): super().__init__() self.conv1a = nn.Conv2d(3, 32, 3, padding=1) self.bn1a = nn.BatchNorm(32) self.conv1b = nn.Conv2d(32, 32, 3, padding=1) self.bn1b = nn.BatchNorm(32) self.conv2a = nn.Conv2d(32, 64, 3, padding=1) self.bn2a = nn.BatchNorm(64) self.conv2b = nn.Conv2d(64, 64, 3, padding=1) self.bn2b = nn.BatchNorm(64) self.conv3a = nn.Conv2d(64, 128, 3, padding=1) self.bn3a = nn.BatchNorm(128) self.conv3b = nn.Conv2d(128, 128, 3, padding=1) self.bn3b = nn.BatchNorm(128) self.conv4a = nn.Conv2d(128, 256, 3, padding=1) self.bn4a = nn.BatchNorm(256) self.conv4b = nn.Conv2d(256, 256, 3, padding=1) self.bn4b = nn.BatchNorm(256) self.pool = nn.MaxPool2d(2, 2) self.fc1 = nn.Linear(256 * 4 * 4, 512) self.fc2 = nn.Linear(512, num_classes) def __call__(self, x): # x input is (B, 64, 64, C) x = nn.relu(self.bn1a(self.conv1a(x))) x = nn.relu(self.bn1b(self.conv1b(x))) x = self.pool(x) x = nn.relu(self.bn2a(self.conv2a(x))) x = nn.relu(self.bn2b(self.conv2b(x))) x = self.pool(x) x = nn.relu(self.bn3a(self.conv3a(x))) x = nn.relu(self.bn3b(self.conv3b(x))) x = self.pool(x) x = nn.relu(self.bn4a(self.conv4a(x))) x = nn.relu(self.bn4b(self.conv4b(x))) x = self.pool(x) x = mx.flatten(x, start_axis=1) x = nn.relu(self.fc1(x)) x = self.fc2(x) return x