File size: 2,871 Bytes
d7661da
 
 
935d968
d7661da
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
935d968
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
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