chesscv-lichess-fen / mlx_model.py
zimengxiong's picture
Upload stronger 64x64 Lichess FEN classifier
935d968 verified
Raw History Blame Contribute Delete
2.87 kB
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