zimengxiong commited on
Commit
d7661da
·
verified ·
1 Parent(s): 22ab823

Upload Lichess FEN square classifier

Browse files
README.md ADDED
@@ -0,0 +1,86 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: agpl-3.0
3
+ tags:
4
+ - chess
5
+ - computer-vision
6
+ - fen
7
+ - lichess
8
+ - mlx
9
+ - safetensors
10
+ library_name: mlx
11
+ pipeline_tag: image-classification
12
+ ---
13
+
14
+ # ChessCV Lichess FEN Square Classifier
15
+
16
+ This repository contains an MLX chess-board square classifier for FEN recognition from Lichess-style 2D board screenshots.
17
+
18
+ The model classifies each 32×32 RGB board-square crop into one of 13 classes: empty square plus the 12 standard chess pieces. A full FEN can be produced by running the classifier over all 64 square crops in board order and converting the predicted class sequence into piece-placement notation.
19
+
20
+ ## Model file
21
+
22
+ - `mlx_model_lichess_fen.safetensors` — MLX weights for `ChessPieceCNN`
23
+ - `mlx_model.py` — minimal MLX model definition
24
+ - `labels.json` — class-index mapping
25
+ - `training_config.json` — high-level training metadata
26
+
27
+ ## Labels
28
+
29
+ ```text
30
+ 0: empty
31
+ 1: white pawn
32
+ 2: white rook
33
+ 3: white knight
34
+ 4: white bishop
35
+ 5: white queen
36
+ 6: white king
37
+ 7: black pawn
38
+ 8: black rook
39
+ 9: black knight
40
+ 10: black bishop
41
+ 11: black queen
42
+ 12: black king
43
+ ```
44
+
45
+ ## Input format
46
+
47
+ - Shape: `(N, 32, 32, 3)`
48
+ - Color: RGB
49
+ - Value range: float32 normalized to `[0, 1]`
50
+
51
+ For board-level FEN recognition, crop the detected board into 64 squares, resize each crop to 32×32, run a batch inference pass, then map argmax outputs to FEN piece-placement symbols.
52
+
53
+ ## Example usage
54
+
55
+ ```python
56
+ import mlx.core as mx
57
+ import numpy as np
58
+ from mlx_model import ChessPieceCNN
59
+
60
+ model = ChessPieceCNN()
61
+ model.load_weights("mlx_model_lichess_fen.safetensors")
62
+
63
+ # batch: np.ndarray shaped (64, 32, 32, 3), RGB, float32 in [0, 1]
64
+ logits = model(mx.array(batch))
65
+ mx.eval(logits)
66
+ labels = np.argmax(np.array(logits), axis=1)
67
+ ```
68
+
69
+ ## Training data summary
70
+
71
+ The model was trained on synthetic square crops generated from Lichess piece assets, with extra emphasis on the default Lichess piece style. The dataset includes occupied and empty squares, varied Lichess-like board colors, scale changes, compression artifacts, blur, lighting/gamma variation, and small-screen resampling effects.
72
+
73
+ Final training run:
74
+
75
+ - Samples: 240,000
76
+ - Epochs: 16
77
+ - Empty-square samples: ~42%
78
+ - Final training accuracy: 94.96%
79
+
80
+ ## Intended use
81
+
82
+ This model is intended as a component in chess computer-vision pipelines that detect a 2D chessboard on screen and reconstruct the piece-placement part of a FEN string.
83
+
84
+ ## License and source assets
85
+
86
+ The model was trained using Lichess piece assets from the open-source Lichess project. Lichess is distributed under AGPL-3.0-or-later; individual piece sets may have their own original author/license history upstream.
labels.json ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "0": "empty",
3
+ "1": "wP",
4
+ "2": "wR",
5
+ "3": "wN",
6
+ "4": "wB",
7
+ "5": "wQ",
8
+ "6": "wK",
9
+ "7": "bP",
10
+ "8": "bR",
11
+ "9": "bN",
12
+ "10": "bB",
13
+ "11": "bQ",
14
+ "12": "bK"
15
+ }
mlx_model.py ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import mlx.core as mx
2
+ import mlx.nn as nn
3
+
4
+ class ChessPieceCNN(nn.Module):
5
+ def __init__(self, num_classes=13):
6
+ super().__init__()
7
+ self.conv1 = nn.Conv2d(3, 32, 3, padding=1)
8
+ self.bn1 = nn.BatchNorm(32)
9
+ self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
10
+ self.bn2 = nn.BatchNorm(64)
11
+ self.conv3 = nn.Conv2d(64, 128, 3, padding=1)
12
+ self.bn3 = nn.BatchNorm(128)
13
+ self.pool = nn.MaxPool2d(2, 2)
14
+
15
+ self.fc1 = nn.Linear(128 * 4 * 4, 256)
16
+ self.fc2 = nn.Linear(256, num_classes)
17
+
18
+ def __call__(self, x):
19
+ # x input is (B, H, W, C)
20
+ x = nn.relu(self.bn1(self.conv1(x)))
21
+ x = self.pool(x)
22
+ x = nn.relu(self.bn2(self.conv2(x)))
23
+ x = self.pool(x)
24
+ x = nn.relu(self.bn3(self.conv3(x)))
25
+ x = self.pool(x)
26
+
27
+ x = mx.flatten(x, start_axis=1)
28
+ x = nn.relu(self.fc1(x))
29
+ x = self.fc2(x)
30
+ return x
mlx_model_lichess_fen.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:9395c7b2121bae026182e7169af3fb907ff1b2752e9b5db2458dcdb0a7c76d38
3
+ size 2489811
training_config.json ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architecture": "ChessPieceCNN",
3
+ "framework": "MLX",
4
+ "input_shape": [32, 32, 3],
5
+ "num_classes": 13,
6
+ "target_platform": "Lichess-style 2D chess boards",
7
+ "samples": 240000,
8
+ "epochs": 16,
9
+ "empty_square_probability": 0.42,
10
+ "default_lichess_piece_set": "cburnett",
11
+ "final_training_loss": 0.1234,
12
+ "final_training_accuracy": 0.9496,
13
+ "normalization": "RGB float32 in [0, 1]"
14
+ }