Spaces:
Sleeping
Sleeping
Commit ·
20d7fde
0
Parent(s):
Initial commit
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +1 -0
- .gitignore +21 -0
- .python-version +1 -0
- README.md +159 -0
- app.py +174 -0
- data/real_crops/labels.csv +11 -0
- docs/superpowers/specs/2026-07-08-hf-space-gradio-demo-design.md +70 -0
- eval.py +72 -0
- legacy/a.py +43 -0
- legacy/b.py +93 -0
- legacy/b2.py +124 -0
- legacy/b3.py +80 -0
- legacy/b4.py +143 -0
- legacy/book.txt +8 -0
- legacy/c.py +108 -0
- legacy/class_mapping.json +1 -0
- legacy/d.py +64 -0
- legacy/d2.py +118 -0
- legacy/d3.py +105 -0
- legacy/d4.py +78 -0
- legacy/d5.py +80 -0
- legacy/debug_AI_eyes.png +0 -0
- legacy/e4.py +28 -0
- legacy/image copy.png +0 -0
- legacy/image.png +0 -0
- legacy/what_the_model_actually_sees.png +0 -0
- models/best.classes.json +1 -0
- models/best.onnx +3 -0
- models/kazan.classes.json +1 -0
- models/kazan.onnx +3 -0
- notebooks/colab_train.ipynb +93 -0
- predict.py +68 -0
- read_game.py +95 -0
- requirements-train.txt +9 -0
- requirements.txt +8 -0
- scripts/export_onnx.py +82 -0
- scripts/get_glyphs.py +182 -0
- tests/fixtures/shortest_candidate_replay.tsv +14 -0
- tests/fixtures/shortest_terminal_game.txt +6 -0
- tests/test_rules.py +110 -0
- togyz/__init__.py +0 -0
- togyz/classes.py +23 -0
- togyz/dataset.py +71 -0
- togyz/glyphs.py +167 -0
- togyz/model.py +46 -0
- togyz/pipeline.py +355 -0
- togyz/preprocess.py +52 -0
- togyz/rules.py +143 -0
- togyz/sheet.py +298 -0
- togyz/synth.py +238 -0
.gitattributes
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
*.onnx filter=lfs diff=lfs merge=lfs -text
|
.gitignore
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
__pycache__/
|
| 2 |
+
.DS_Store
|
| 3 |
+
venv/
|
| 4 |
+
venv311/
|
| 5 |
+
|
| 6 |
+
# datasets & training artifacts (regenerable)
|
| 7 |
+
ingredients/
|
| 8 |
+
glyph_data/
|
| 9 |
+
checkpoints/
|
| 10 |
+
train_data/
|
| 11 |
+
train_data2/
|
| 12 |
+
train_data3/
|
| 13 |
+
train_data4/
|
| 14 |
+
train_data5/
|
| 15 |
+
preview*.png
|
| 16 |
+
out/
|
| 17 |
+
project.zip
|
| 18 |
+
|
| 19 |
+
# archived experiments: keep code, ignore heavy model binaries
|
| 20 |
+
legacy/*.pkl
|
| 21 |
+
legacy/*.pth
|
.python-version
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
3.11.8
|
README.md
ADDED
|
@@ -0,0 +1,159 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
title: Togyzkumalak Scoresheet Reader
|
| 3 |
+
emoji: ♟️
|
| 4 |
+
colorFrom: indigo
|
| 5 |
+
colorTo: green
|
| 6 |
+
sdk: gradio
|
| 7 |
+
sdk_version: 6.20.0
|
| 8 |
+
app_file: app.py
|
| 9 |
+
pinned: false
|
| 10 |
+
---
|
| 11 |
+
|
| 12 |
+
# Togyzkumalak Move OCR
|
| 13 |
+
|
| 14 |
+
Classifies a photo of a single scoresheet cell into one of **163 classes**:
|
| 15 |
+
162 moves (`11` … `99x` — start hole 1–9, end hole 1–9, optional capture `x`)
|
| 16 |
+
plus `empty`. A full-sheet reader turns a whole scoresheet photo into PGN game
|
| 17 |
+
records, and a **Gradio demo** (`app.py`) packages it for a HuggingFace Space.
|
| 18 |
+
|
| 19 |
+
## Live demo (HuggingFace Space)
|
| 20 |
+
|
| 21 |
+
`app.py` is a Gradio app: upload up to **5** scoresheet photos, optionally tag
|
| 22 |
+
each game's known result, and download the reconstructed PGNs (beam / raw /
|
| 23 |
+
legal) as a zip. Inference runs on the exported **ONNX** models (torch-free),
|
| 24 |
+
so the Space needs only the light serving deps in `requirements.txt`.
|
| 25 |
+
|
| 26 |
+
```bash
|
| 27 |
+
python scripts/export_onnx.py # checkpoints/*.pt -> single-file .onnx (+ .classes.json)
|
| 28 |
+
cp checkpoints/best.onnx checkpoints/best.classes.json models/
|
| 29 |
+
cp checkpoints/kazan/best.onnx models/kazan.onnx
|
| 30 |
+
cp checkpoints/kazan/best.classes.json models/kazan.classes.json
|
| 31 |
+
python app.py # serve locally at http://127.0.0.1:7860
|
| 32 |
+
```
|
| 33 |
+
|
| 34 |
+
Deploy: create a Gradio Space and push this repo to its `origin` remote (the
|
| 35 |
+
`.onnx` files are tracked with Git LFS via `.gitattributes`). The Space reads
|
| 36 |
+
the YAML front-matter above and installs `requirements.txt`. No secrets or
|
| 37 |
+
database are required — it is a stateless demo capped at 5 images per run, with
|
| 38 |
+
requests serialized through Gradio's queue so a busy free Space degrades into a
|
| 39 |
+
short wait rather than 429 errors.
|
| 40 |
+
|
| 41 |
+
## Why the old pipeline failed (and what this one does differently)
|
| 42 |
+
|
| 43 |
+
| Problem | Old pipeline | This pipeline |
|
| 44 |
+
|---|---|---|
|
| 45 |
+
| Glyph style | EMNIST (American digits) — but players write **European/Kazakh** digits: crossed 7, serif 1, cursive 9 | ARDIS dataset (real European handwriting) + procedural crossbars/serifs on EMNIST glyphs |
|
| 46 |
+
| Preprocessing | Train and inference normalized **differently** (`0.5/0.5` vs ImageNet stats) | One shared module, `togyz/preprocess.py`, imported by both |
|
| 47 |
+
| Framing | Small digits with wide margins in fixed 80×40 cells | Random framing from tight crops to loose cells, matching real photos |
|
| 48 |
+
| Data | 48 900 fixed PNGs on disk | Infinite on-the-fly synthesis, deterministic validation set |
|
| 49 |
+
|
| 50 |
+
## Layout
|
| 51 |
+
|
| 52 |
+
```
|
| 53 |
+
togyz/ core library
|
| 54 |
+
classes.py canonical 163-class list (indices match legacy class_mapping.json)
|
| 55 |
+
preprocess.py THE single image->tensor preprocessing (train AND inference)
|
| 56 |
+
glyphs.py per-character glyph pools + procedural style edits
|
| 57 |
+
synth.py cell synthesizer; python -m togyz.synth --preview preview.png
|
| 58 |
+
dataset.py on-the-fly synthetic dataset + real-crop eval dataset
|
| 59 |
+
model.py resnet18 (grayscale, 163 outputs), checkpoint helpers
|
| 60 |
+
train.py training CLI (auto device: cuda/mps/cpu)
|
| 61 |
+
eval.py synthetic val + per-file real-crop report
|
| 62 |
+
predict.py classify images; --allowed restricts to legal moves
|
| 63 |
+
scripts/get_glyphs.py downloads/prepares glyph pools (ARDIS, optional EMNIST)
|
| 64 |
+
data/real_crops/ 10 labeled real crops (labels.csv) — evaluation only
|
| 65 |
+
legacy/ previous experiments, kept for reference
|
| 66 |
+
```
|
| 67 |
+
|
| 68 |
+
## Quickstart (local)
|
| 69 |
+
|
| 70 |
+
```bash
|
| 71 |
+
source venv311/bin/activate # or: pip install -r requirements.txt
|
| 72 |
+
python scripts/get_glyphs.py # one-time: download+prepare ARDIS glyphs
|
| 73 |
+
python -m togyz.synth --preview preview.png # eyeball synthetic vs data/real_crops
|
| 74 |
+
python train.py --epochs 1 --samples-per-epoch 4000 --batch-size 64 # smoke test
|
| 75 |
+
python eval.py --ckpt checkpoints/best.pt
|
| 76 |
+
python predict.py "data/real_crops/*.jpg"
|
| 77 |
+
```
|
| 78 |
+
|
| 79 |
+
A real training run needs a GPU — see Colab below. Defaults
|
| 80 |
+
(`--epochs 20 --samples-per-epoch 50000`) are a sensible full run.
|
| 81 |
+
|
| 82 |
+
## Training on Google Colab (or any GPU machine)
|
| 83 |
+
|
| 84 |
+
Open `notebooks/colab_train.ipynb` in Colab, or manually:
|
| 85 |
+
|
| 86 |
+
```bash
|
| 87 |
+
git clone <this-repo> && cd 9OCR # or upload a zip of the project
|
| 88 |
+
pip install -r requirements.txt
|
| 89 |
+
python scripts/get_glyphs.py --emnist # ingredients/ is gitignored, so also
|
| 90 |
+
# rebuild the EMNIST pool from torchvision
|
| 91 |
+
python train.py --epochs 30
|
| 92 |
+
# download checkpoints/best.pt when done
|
| 93 |
+
```
|
| 94 |
+
|
| 95 |
+
`predict.py`/`eval.py` run anywhere (CPU is fine) with the downloaded
|
| 96 |
+
checkpoint.
|
| 97 |
+
|
| 98 |
+
## Inference
|
| 99 |
+
|
| 100 |
+
```bash
|
| 101 |
+
python predict.py cell.jpg --topk 3
|
| 102 |
+
python predict.py cell.jpg --allowed "12,34x,56"
|
| 103 |
+
```
|
| 104 |
+
|
| 105 |
+
At any game state at most 9 moves are legal. If an upstream game tracker
|
| 106 |
+
passes them via `--allowed`, probabilities are renormalized over just those
|
| 107 |
+
moves — a large, free accuracy boost when reading whole games.
|
| 108 |
+
|
| 109 |
+
## Reading a whole scoresheet
|
| 110 |
+
|
| 111 |
+
```bash
|
| 112 |
+
python read_game.py "data/2026-07-06 00.00.20.jpg" --out out/sheet1 --result 0-1
|
| 113 |
+
```
|
| 114 |
+
|
| 115 |
+
Pass `--result` (1-0 / 0-1 / draw, from the sheet footer) when known: the
|
| 116 |
+
beam's final pool is re-ranked to prefer reconstructions whose end state is
|
| 117 |
+
consistent with how the game actually ended.
|
| 118 |
+
|
| 119 |
+
The summary strips between the tables record both kazan counts every 10
|
| 120 |
+
moves (Black's box above the strip, White's below, always two digits 00-81).
|
| 121 |
+
A second classifier reads them (`python train.py --task kazan`, saved to
|
| 122 |
+
`checkpoints/kazan/best.pt`) and the beam gains likelihood when a
|
| 123 |
+
hypothesis' computed kazans match these written checkpoints.
|
| 124 |
+
|
| 125 |
+
Finds the 8 printed move tables, classifies every cell, and writes:
|
| 126 |
+
|
| 127 |
+
- `game.json` — per ply: all 163 class probabilities, top-5, legality info
|
| 128 |
+
- `raw.pgn` — pure classifier argmax per ply, even if illegal
|
| 129 |
+
- `legal.pgn` — strict replay, stops at the first illegal argmax
|
| 130 |
+
- `beam.pgn` — **best reconstruction**: beam search over all legal
|
| 131 |
+
continuations per ply; each legal move is scored by blending its exact
|
| 132 |
+
class probability with the probability mass of its source digit
|
| 133 |
+
(the landing digit is derived from the position, so a misread second
|
| 134 |
+
digit doesn't discard the right source hole). Longest fully legal
|
| 135 |
+
sequence wins, ties broken by joint probability.
|
| 136 |
+
- `annotated.jpg`, `cells/` — visual debugging
|
| 137 |
+
|
| 138 |
+
Move annotations: `+` capture, `x` tuzdyk creation. Strip `+` to feed the
|
| 139 |
+
moves to the 9Q engine. Rules implementation: `togyz/rules.py`, validated
|
| 140 |
+
against the 9Q C++ engine fixture (`python tests/test_rules.py`).
|
| 141 |
+
|
| 142 |
+
Accuracy depends heavily on photo resolution — the CLI warns when cells are
|
| 143 |
+
under 35 px tall; photograph sheets at full camera resolution.
|
| 144 |
+
|
| 145 |
+
## Improving accuracy further
|
| 146 |
+
|
| 147 |
+
1. **Label more real cells** — append rows to `data/real_crops/labels.csv`.
|
| 148 |
+
Even ~50 real crops make the reported real-crop accuracy meaningful; a few
|
| 149 |
+
hundred would allow fine-tuning on them.
|
| 150 |
+
2. Add glyph styles: put white-on-black PNG masks under
|
| 151 |
+
`glyph_data/<source>/<char>/` and register the source in `togyz/glyphs.py`.
|
| 152 |
+
3. Full-game decoding with a Togyzkumalak rules engine (choose the most
|
| 153 |
+
probable *legal* move per cell) — hook already exists via `--allowed`.
|
| 154 |
+
|
| 155 |
+
## Notes
|
| 156 |
+
|
| 157 |
+
- `train_data5/` (the old pre-generated dataset) is no longer used and can be
|
| 158 |
+
deleted; synthesis now happens on the fly.
|
| 159 |
+
- Old models/scripts live in `legacy/` (`.pkl`/`.pth` files are gitignored).
|
app.py
ADDED
|
@@ -0,0 +1,174 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Gradio demo for the Togyzkumalak scoresheet reader (HuggingFace Space).
|
| 2 |
+
|
| 3 |
+
Upload up to 5 scoresheet photos, optionally tag each with its known result,
|
| 4 |
+
and download the reconstructed PGNs. Inference runs on the exported ONNX
|
| 5 |
+
models (torch-free) via `togyz.pipeline.run_pipeline`.
|
| 6 |
+
|
| 7 |
+
This is a demo, not production: state is per-session, work is capped at 5
|
| 8 |
+
images per run, and requests are serialized through Gradio's queue so a shared
|
| 9 |
+
free Space degrades into a wait rather than a flurry of 429s.
|
| 10 |
+
"""
|
| 11 |
+
|
| 12 |
+
import os
|
| 13 |
+
import sys
|
| 14 |
+
import tempfile
|
| 15 |
+
import zipfile
|
| 16 |
+
from pathlib import Path
|
| 17 |
+
|
| 18 |
+
# Unbuffered stdout so boot progress actually shows in the Space container logs
|
| 19 |
+
# (otherwise a slow import/model-load looks like a silent hang).
|
| 20 |
+
try:
|
| 21 |
+
sys.stdout.reconfigure(line_buffering=True)
|
| 22 |
+
sys.stderr.reconfigure(line_buffering=True)
|
| 23 |
+
except Exception:
|
| 24 |
+
pass
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def _log(msg):
|
| 28 |
+
print(f"[app] {msg}", flush=True)
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
_log("importing gradio ...")
|
| 32 |
+
import gradio as gr
|
| 33 |
+
|
| 34 |
+
_log("importing pipeline ...")
|
| 35 |
+
from togyz.pipeline import load_classifier, run_pipeline
|
| 36 |
+
|
| 37 |
+
MAX_IMAGES = 5
|
| 38 |
+
MODEL_DIR = Path(__file__).parent / "models"
|
| 39 |
+
RESULT_CHOICES = ["unknown", "1-0", "0-1", "draw"]
|
| 40 |
+
|
| 41 |
+
# Load the ONNX sessions once at import - warm for the whole process lifetime.
|
| 42 |
+
_log(f"loading move model from {MODEL_DIR / 'best.onnx'} ...")
|
| 43 |
+
_MOVES = load_classifier(MODEL_DIR / "best.onnx")
|
| 44 |
+
_KAZAN = None
|
| 45 |
+
_kazan_path = MODEL_DIR / "kazan.onnx"
|
| 46 |
+
if _kazan_path.exists():
|
| 47 |
+
_log("loading kazan model ...")
|
| 48 |
+
_KAZAN = load_classifier(_kazan_path)
|
| 49 |
+
_log("models loaded")
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
def _safe_slug(text: str) -> str:
|
| 53 |
+
keep = "".join(c if c.isalnum() else "_" for c in (text or "").strip())
|
| 54 |
+
return keep.strip("_")
|
| 55 |
+
|
| 56 |
+
|
| 57 |
+
def _process_one(image_path, result_choice, base_name, out_dir: Path):
|
| 58 |
+
"""Run the pipeline on one image; return (row, gallery_item, pgn_files)."""
|
| 59 |
+
result = None if result_choice in (None, "unknown") else result_choice
|
| 60 |
+
out = run_pipeline(image_path, _MOVES, _KAZAN, result=result)
|
| 61 |
+
|
| 62 |
+
stop = out["stopped"]
|
| 63 |
+
stop_txt = stop.get("reason", "")
|
| 64 |
+
if "winner" in stop:
|
| 65 |
+
stop_txt += f" ({stop['winner']})"
|
| 66 |
+
note = " ⚠ low-res" if out["low_resolution"] else ""
|
| 67 |
+
|
| 68 |
+
files = []
|
| 69 |
+
for kind in ("beam", "raw", "legal"):
|
| 70 |
+
f = out_dir / f"{base_name}_{kind}.pgn"
|
| 71 |
+
f.write_text(out[f"{kind}_pgn"])
|
| 72 |
+
files.append(str(f))
|
| 73 |
+
|
| 74 |
+
row = [base_name, out["beam_plies"], stop_txt + note, out["beam_pgn"].strip()]
|
| 75 |
+
caption = f"{base_name}: {out['beam_plies']} plies"
|
| 76 |
+
return row, (out["annotated_image"], caption), files
|
| 77 |
+
|
| 78 |
+
|
| 79 |
+
def convert(round_label, *slot_values):
|
| 80 |
+
"""slot_values = [img1, res1, img2, res2, ...] for the 5 fixed slots."""
|
| 81 |
+
images = slot_values[0::2]
|
| 82 |
+
results = slot_values[1::2]
|
| 83 |
+
provided = [(img, res) for img, res in zip(images, results) if img]
|
| 84 |
+
|
| 85 |
+
if not provided:
|
| 86 |
+
raise gr.Error("Please upload at least one scoresheet image.")
|
| 87 |
+
if len(provided) > MAX_IMAGES: # defensive; the UI only exposes 5 slots
|
| 88 |
+
raise gr.Error(f"This demo handles at most {MAX_IMAGES} images per run.")
|
| 89 |
+
|
| 90 |
+
out_dir = Path(tempfile.mkdtemp(prefix="togyz_"))
|
| 91 |
+
round_slug = _safe_slug(round_label)
|
| 92 |
+
rows, gallery, all_files = [], [], []
|
| 93 |
+
|
| 94 |
+
for i, (img, res) in enumerate(provided, start=1):
|
| 95 |
+
prefix = f"{round_slug}_table{i}" if round_slug else f"table{i}"
|
| 96 |
+
try:
|
| 97 |
+
row, gal, files = _process_one(img, res, prefix, out_dir)
|
| 98 |
+
except Exception as exc: # one bad image must not kill the batch
|
| 99 |
+
rows.append([prefix, 0, f"error: {exc}", ""])
|
| 100 |
+
continue
|
| 101 |
+
rows.append(row)
|
| 102 |
+
gallery.append(gal)
|
| 103 |
+
all_files.extend(files)
|
| 104 |
+
|
| 105 |
+
if not all_files:
|
| 106 |
+
# every image errored - still return the table so the user sees why
|
| 107 |
+
return rows, gallery, None
|
| 108 |
+
|
| 109 |
+
zip_path = out_dir / (f"{round_slug}_pgns.zip" if round_slug else "pgns.zip")
|
| 110 |
+
with zipfile.ZipFile(zip_path, "w") as zf:
|
| 111 |
+
for f in all_files:
|
| 112 |
+
zf.write(f, arcname=Path(f).name)
|
| 113 |
+
|
| 114 |
+
return rows, gallery, str(zip_path)
|
| 115 |
+
|
| 116 |
+
|
| 117 |
+
def _busy_wrapper(*args):
|
| 118 |
+
"""Turn infrastructure overload into a friendly message instead of a 500."""
|
| 119 |
+
try:
|
| 120 |
+
return convert(*args)
|
| 121 |
+
except gr.Error:
|
| 122 |
+
raise
|
| 123 |
+
except Exception as exc: # noqa: BLE001 - surface anything else gracefully
|
| 124 |
+
msg = str(exc).lower()
|
| 125 |
+
if "429" in msg or "too many" in msg or "rate" in msg:
|
| 126 |
+
raise gr.Error("Server busy — please retry in a moment.")
|
| 127 |
+
raise gr.Error(f"Something went wrong: {exc}")
|
| 128 |
+
|
| 129 |
+
|
| 130 |
+
with gr.Blocks(title="Togyzkumalak Scoresheet Reader") as demo:
|
| 131 |
+
gr.Markdown(
|
| 132 |
+
"# Togyzkumalak Scoresheet Reader\n"
|
| 133 |
+
"Upload up to **5** scoresheet photos, optionally tag each game's known "
|
| 134 |
+
"result (improves accuracy), then **Convert** to download the PGNs.\n\n"
|
| 135 |
+
"Outputs per game: `beam` (best legal reconstruction), `raw` (pure OCR), "
|
| 136 |
+
"`legal` (strict replay). Free demo — a first run may wake the Space, and "
|
| 137 |
+
"images are processed one at a time."
|
| 138 |
+
)
|
| 139 |
+
round_label = gr.Textbox(label="Round (optional)", placeholder="e.g. 3",
|
| 140 |
+
scale=1, max_lines=1)
|
| 141 |
+
|
| 142 |
+
slots = []
|
| 143 |
+
for i in range(MAX_IMAGES):
|
| 144 |
+
with gr.Row():
|
| 145 |
+
img = gr.Image(label=f"Table {i + 1}", type="filepath", height=150)
|
| 146 |
+
res = gr.Dropdown(RESULT_CHOICES, value="unknown",
|
| 147 |
+
label="Result (1-0 = White won)", scale=1)
|
| 148 |
+
slots.extend([img, res])
|
| 149 |
+
|
| 150 |
+
convert_btn = gr.Button("Convert", variant="primary")
|
| 151 |
+
|
| 152 |
+
results_table = gr.Dataframe(
|
| 153 |
+
headers=["game", "legal plies", "stopped", "beam PGN"],
|
| 154 |
+
label="Results", wrap=True, interactive=False,
|
| 155 |
+
)
|
| 156 |
+
gallery = gr.Gallery(label="Annotated reconstruction", columns=2, height="auto")
|
| 157 |
+
zip_out = gr.File(label="Download all PGNs (zip)")
|
| 158 |
+
|
| 159 |
+
convert_btn.click(
|
| 160 |
+
_busy_wrapper,
|
| 161 |
+
inputs=[round_label, *slots],
|
| 162 |
+
outputs=[results_table, gallery, zip_out],
|
| 163 |
+
)
|
| 164 |
+
|
| 165 |
+
# Serialize CPU-heavy runs: callers wait in a bounded queue instead of
|
| 166 |
+
# overloading the shared Space (which is what triggers 429s).
|
| 167 |
+
demo.queue(max_size=16, default_concurrency_limit=1)
|
| 168 |
+
|
| 169 |
+
if __name__ == "__main__":
|
| 170 |
+
# Bind explicitly to 0.0.0.0 and the Space's port so HF can detect the
|
| 171 |
+
# running app (the default 127.0.0.1 bind can leave a Space stuck "Starting").
|
| 172 |
+
port = int(os.environ.get("GRADIO_SERVER_PORT", os.environ.get("PORT", 7860)))
|
| 173 |
+
_log(f"launching gradio on 0.0.0.0:{port} ...")
|
| 174 |
+
demo.launch(server_name="0.0.0.0", server_port=port)
|
data/real_crops/labels.csv
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
filename,label
|
| 2 |
+
96_a.jpg,96
|
| 3 |
+
96_b.jpg,96
|
| 4 |
+
77_a.jpg,77
|
| 5 |
+
77_b.jpg,77
|
| 6 |
+
13x_a.jpg,13x
|
| 7 |
+
13x_b.jpg,13x
|
| 8 |
+
37_a.jpg,37
|
| 9 |
+
37_b.jpg,37
|
| 10 |
+
85_a.jpg,85
|
| 11 |
+
42_a.jpg,42
|
docs/superpowers/specs/2026-07-08-hf-space-gradio-demo-design.md
ADDED
|
@@ -0,0 +1,70 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Togyzkumalak Scoresheet Reader — Gradio demo on a HuggingFace Space
|
| 2 |
+
|
| 3 |
+
## Context
|
| 4 |
+
|
| 5 |
+
The OCR pipeline reads a scoresheet photo into PGN records (sheet 1 = 36/39
|
| 6 |
+
exact plies). We want a public, shareable **demo** — not production — where
|
| 7 |
+
anyone with the link uploads a few scoresheet photos and downloads the PGNs.
|
| 8 |
+
|
| 9 |
+
Hosting is a single **HuggingFace Space running Gradio**. A free Space has
|
| 10 |
+
ample RAM (≈16 GB), so no database and no external services are needed —
|
| 11 |
+
outputs are files downloaded in-session. Constraints: **≤5 images per run**
|
| 12 |
+
and **graceful handling of overload (429)** on shared free compute.
|
| 13 |
+
|
| 14 |
+
The blocker on smaller hosts was PyTorch (528 MiB peak RAM, 441 MB on disk).
|
| 15 |
+
The fix — kept here because it also makes the Space faster and lighter — is to
|
| 16 |
+
run inference under **onnxruntime + numpy**, dropping torch from the serving
|
| 17 |
+
path (torch stays only for training/export).
|
| 18 |
+
|
| 19 |
+
## Architecture
|
| 20 |
+
|
| 21 |
+
```
|
| 22 |
+
HuggingFace Space (Gradio SDK)
|
| 23 |
+
app.py UI + handler; loads 2 ONNX sessions once at import
|
| 24 |
+
togyz/pipeline.py run_pipeline() — torch-free core, shared with the CLI
|
| 25 |
+
togyz/{sheet,rules,preprocess,classes}.py already torch-free
|
| 26 |
+
models/*.onnx committed via Git LFS (+ .classes.json sidecars)
|
| 27 |
+
requirements.txt serving deps (torch-free); requirements-train.txt for training
|
| 28 |
+
```
|
| 29 |
+
|
| 30 |
+
## Key decisions
|
| 31 |
+
|
| 32 |
+
- **onnxruntime + numpy inference.** `togyz/preprocess.py` returns numpy;
|
| 33 |
+
`togyz/pipeline.py` reimplements cell classification, legal-move scoring,
|
| 34 |
+
kazan evidence, and beam search with numpy. Verified **byte-identical**
|
| 35 |
+
`beam.pgn` vs the torch baseline on both sample sheets. Peak RAM 394 MiB
|
| 36 |
+
(was 528), wall time 8.8 s (was 12.9).
|
| 37 |
+
- **Single-file ONNX.** `scripts/export_onnx.py` uses the legacy exporter
|
| 38 |
+
(`dynamo=False`) + `onnx.save_model(save_as_external_data=False)` so each
|
| 39 |
+
model is one self-contained ~45 MB `.onnx` (no `.data` sidecar) — clean for
|
| 40 |
+
Git LFS. Class list saved as a `<stem>.classes.json` sidecar.
|
| 41 |
+
- **Fixed 5 slots** (image + result dropdown each) structurally enforce the
|
| 42 |
+
5-image cap and give unambiguous per-game result tagging. Result improves
|
| 43 |
+
accuracy via the beam's end-state re-ranking; `unknown` → no hint.
|
| 44 |
+
- **Graceful overload.** `demo.queue(default_concurrency_limit=1, max_size=16)`
|
| 45 |
+
serializes CPU-heavy runs so callers wait rather than triggering 429s; a
|
| 46 |
+
`_busy_wrapper` maps any 429/rate error to a friendly `gr.Error`.
|
| 47 |
+
- **Per-image isolation.** Each image runs in `try/except`; a failure reports
|
| 48 |
+
an error row and the batch continues.
|
| 49 |
+
|
| 50 |
+
## Outputs
|
| 51 |
+
|
| 52 |
+
Per game: `beam` (best legal reconstruction), `raw` (pure OCR argmax), `legal`
|
| 53 |
+
(strict replay) PGNs, plus an annotated preview image. All PGNs are bundled
|
| 54 |
+
into a downloadable zip; filenames use the optional round label
|
| 55 |
+
(`R3_table1_beam.pgn`).
|
| 56 |
+
|
| 57 |
+
## Verification (all passed)
|
| 58 |
+
|
| 59 |
+
1. ONNX parity: `beam.pgn` identical to torch on both sheets.
|
| 60 |
+
2. RAM 394 MiB / 8.8 s per image via the ONNX CLI.
|
| 61 |
+
3. `app.convert()` headless on both samples → 39 + 77 plies, zip with 6 PGNs,
|
| 62 |
+
beam matches the CLI.
|
| 63 |
+
4. Resilience: a table-less image reports an error row while the good image
|
| 64 |
+
still produces output.
|
| 65 |
+
|
| 66 |
+
## Out of scope (demo)
|
| 67 |
+
|
| 68 |
+
Batches >5, persistence/history, auth, tournament round/table auto-detection.
|
| 69 |
+
Larger production hosting (Render/Supabase/Vercel) was considered and dropped
|
| 70 |
+
in favor of this single-Space demo.
|
eval.py
ADDED
|
@@ -0,0 +1,72 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Evaluate a checkpoint on the synthetic validation set and the real crops.
|
| 2 |
+
|
| 3 |
+
python eval.py --ckpt checkpoints/best.pt
|
| 4 |
+
"""
|
| 5 |
+
|
| 6 |
+
import argparse
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
from torch.utils.data import DataLoader
|
| 10 |
+
|
| 11 |
+
from togyz.dataset import RealCropDataset, SyntheticCellDataset
|
| 12 |
+
from togyz.model import auto_device, load_checkpoint
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def main() -> None:
|
| 16 |
+
parser = argparse.ArgumentParser(description=__doc__)
|
| 17 |
+
parser.add_argument("--ckpt", default="checkpoints/best.pt")
|
| 18 |
+
parser.add_argument("--val-size", type=int, default=4_000)
|
| 19 |
+
parser.add_argument("--batch-size", type=int, default=128)
|
| 20 |
+
parser.add_argument("--workers", type=int, default=4)
|
| 21 |
+
parser.add_argument("--device", default=None)
|
| 22 |
+
args = parser.parse_args()
|
| 23 |
+
|
| 24 |
+
device = torch.device(args.device) if args.device else auto_device()
|
| 25 |
+
model, ckpt = load_checkpoint(args.ckpt, device)
|
| 26 |
+
classes = ckpt["classes"]
|
| 27 |
+
print(f"Loaded {args.ckpt} (epoch {ckpt['epoch'] + 1}, synth val {ckpt['val_acc']:.2%})")
|
| 28 |
+
|
| 29 |
+
# synthetic validation (same fixed seed as train.py)
|
| 30 |
+
val_loader = DataLoader(
|
| 31 |
+
SyntheticCellDataset(args.val_size, seed=1234),
|
| 32 |
+
batch_size=args.batch_size,
|
| 33 |
+
num_workers=args.workers,
|
| 34 |
+
)
|
| 35 |
+
correct = total = 0
|
| 36 |
+
with torch.no_grad():
|
| 37 |
+
for images, labels in val_loader:
|
| 38 |
+
preds = model(images.to(device)).argmax(dim=1).cpu()
|
| 39 |
+
correct += (preds == labels).sum().item()
|
| 40 |
+
total += labels.numel()
|
| 41 |
+
print(f"Synthetic val accuracy: {correct / total:.2%} ({total} samples)")
|
| 42 |
+
|
| 43 |
+
# real crops with per-file report
|
| 44 |
+
real = RealCropDataset()
|
| 45 |
+
if len(real) == 0:
|
| 46 |
+
print("No real crops found in data/real_crops - skipping.")
|
| 47 |
+
return
|
| 48 |
+
images, labels, names = real.batch()
|
| 49 |
+
with torch.no_grad():
|
| 50 |
+
probs = torch.softmax(model(images.to(device)).cpu(), dim=1)
|
| 51 |
+
topk = probs.topk(3, dim=1)
|
| 52 |
+
|
| 53 |
+
print(f"\nReal crops ({len(real)} files):")
|
| 54 |
+
top1 = top3 = 0
|
| 55 |
+
for i, name in enumerate(names):
|
| 56 |
+
truth = classes[labels[i]]
|
| 57 |
+
guesses = [
|
| 58 |
+
f"{classes[idx]} {p:.1%}"
|
| 59 |
+
for idx, p in zip(topk.indices[i].tolist(), topk.values[i].tolist())
|
| 60 |
+
]
|
| 61 |
+
hit1 = topk.indices[i, 0] == labels[i]
|
| 62 |
+
hit3 = (topk.indices[i] == labels[i]).any()
|
| 63 |
+
top1 += int(hit1)
|
| 64 |
+
top3 += int(hit3)
|
| 65 |
+
marker = "OK " if hit1 else ("~3 " if hit3 else "MISS")
|
| 66 |
+
print(f" [{marker}] {name:<14} truth={truth:<4} top3: {', '.join(guesses)}")
|
| 67 |
+
print(f"Real top-1: {top1}/{len(real)} ({top1 / len(real):.0%}) "
|
| 68 |
+
f"top-3: {top3}/{len(real)} ({top3 / len(real):.0%})")
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
if __name__ == "__main__":
|
| 72 |
+
main()
|
legacy/a.py
ADDED
|
@@ -0,0 +1,43 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import torchvision
|
| 3 |
+
from PIL import Image
|
| 4 |
+
|
| 5 |
+
# 1. Create the folder structure
|
| 6 |
+
base_dir = "ingredients"
|
| 7 |
+
classes_we_want = {
|
| 8 |
+
1: "1", 2: "2", 3: "3", 4: "4", 5: "5",
|
| 9 |
+
6: "6", 7: "7", 8: "8", 9: "9", 33: "x"
|
| 10 |
+
}
|
| 11 |
+
|
| 12 |
+
for folder_name in classes_we_want.values():
|
| 13 |
+
os.makedirs(os.path.join(base_dir, folder_name), exist_ok=True)
|
| 14 |
+
|
| 15 |
+
# 2. Load the dataset
|
| 16 |
+
dataset = torchvision.datasets.EMNIST(root='./data', split='balanced', train=True, download=True)
|
| 17 |
+
|
| 18 |
+
print("Extracting and FIXING rotation... please wait.")
|
| 19 |
+
|
| 20 |
+
counts = {k: 0 for k in classes_we_want.keys()}
|
| 21 |
+
limit_per_class = 2000
|
| 22 |
+
|
| 23 |
+
for i in range(len(dataset)):
|
| 24 |
+
img, label = dataset[i] # img is a PIL Image
|
| 25 |
+
|
| 26 |
+
if label in classes_we_want:
|
| 27 |
+
if counts[label] < limit_per_class:
|
| 28 |
+
label_name = classes_we_want[label]
|
| 29 |
+
file_path = os.path.join(base_dir, label_name, f"{label_name}_{counts[label]}.png")
|
| 30 |
+
|
| 31 |
+
# --- THE FIX ---
|
| 32 |
+
# EMNIST is stored (width, height) instead of (height, width)
|
| 33 |
+
# We transpose it to make it human-readable
|
| 34 |
+
fixed_img = img.transpose(Image.TRANSPOSE)
|
| 35 |
+
# ---------------
|
| 36 |
+
|
| 37 |
+
fixed_img.save(file_path)
|
| 38 |
+
counts[label] += 1
|
| 39 |
+
|
| 40 |
+
if all(c >= limit_per_class for c in counts.values()):
|
| 41 |
+
break
|
| 42 |
+
|
| 43 |
+
print("Success! Check your 'ingredients' folder now. They should be upright.")
|
legacy/b.py
ADDED
|
@@ -0,0 +1,93 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import random
|
| 3 |
+
import glob
|
| 4 |
+
import numpy as np
|
| 5 |
+
from PIL import Image, ImageOps, ImageFilter
|
| 6 |
+
|
| 7 |
+
# --- CONFIGURATION ---
|
| 8 |
+
INGREDIENTS_PATH = "ingredients"
|
| 9 |
+
OUTPUT_PATH = "train_data"
|
| 10 |
+
BOX_HEIGHT = 40
|
| 11 |
+
BOX_WIDTH = 120 # 3:1 Proportion
|
| 12 |
+
SAMPLES_PER_CLASS = 300 # Adjust based on your disk space
|
| 13 |
+
|
| 14 |
+
# 1. Generate the 163 Class Names
|
| 15 |
+
# Format: "72" (no x) or "72x" (with x)
|
| 16 |
+
move_classes = []
|
| 17 |
+
for start_hole in range(1, 10):
|
| 18 |
+
for end_hole in range(1, 10):
|
| 19 |
+
move_classes.append(f"{start_hole}{end_hole}") # e.g., "72"
|
| 20 |
+
move_classes.append(f"{start_hole}{end_hole}x") # e.g., "72x"
|
| 21 |
+
|
| 22 |
+
classes = move_classes + ['empty']
|
| 23 |
+
|
| 24 |
+
os.makedirs(OUTPUT_PATH, exist_ok=True)
|
| 25 |
+
|
| 26 |
+
def get_random_ingredient(char):
|
| 27 |
+
# char will be '1'-'9' or 'x'
|
| 28 |
+
files = glob.glob(os.path.join(INGREDIENTS_PATH, char, "*.png"))
|
| 29 |
+
if not files:
|
| 30 |
+
raise ValueError(f"No images found for character: {char}")
|
| 31 |
+
return Image.open(random.choice(files))
|
| 32 |
+
|
| 33 |
+
def create_move_image(class_name):
|
| 34 |
+
# 1. Create the 3:1 paper background (light gray/off-white)
|
| 35 |
+
bg_color = random.randint(220, 250)
|
| 36 |
+
img = Image.new('L', (BOX_WIDTH, BOX_HEIGHT), color=bg_color)
|
| 37 |
+
|
| 38 |
+
if class_name == 'empty':
|
| 39 |
+
return img
|
| 40 |
+
|
| 41 |
+
chars_to_draw = list(class_name)
|
| 42 |
+
|
| 43 |
+
for slot in range(len(chars_to_draw)):
|
| 44 |
+
char = chars_to_draw[slot]
|
| 45 |
+
char_img = get_random_ingredient(char) # This is White-on-Black
|
| 46 |
+
|
| 47 |
+
# Resize
|
| 48 |
+
size = random.randint(28, 36)
|
| 49 |
+
char_img = char_img.resize((size, size), Image.Resampling.LANCZOS)
|
| 50 |
+
|
| 51 |
+
# Rotate
|
| 52 |
+
char_img = char_img.rotate(random.randint(-10, 10), expand=False, fillcolor=0)
|
| 53 |
+
|
| 54 |
+
# --- THE FIX: MASKED PASTING ---
|
| 55 |
+
# Instead of inverting the whole square, we use the original
|
| 56 |
+
# White-on-Black image as a "mask".
|
| 57 |
+
|
| 58 |
+
# Create a solid black square of the same size
|
| 59 |
+
ink_color = random.randint(0, 50) # Dark gray to black ink
|
| 60 |
+
ink_layer = Image.new('L', (size, size), color=ink_color)
|
| 61 |
+
|
| 62 |
+
# Position
|
| 63 |
+
slot_center_x = (slot * 40) + 20
|
| 64 |
+
paste_x = slot_center_x - (size // 2) + random.randint(-4, 4)
|
| 65 |
+
paste_y = (BOX_HEIGHT // 2) - (size // 2) + random.randint(-3, 3)
|
| 66 |
+
|
| 67 |
+
# We paste the "ink_layer" onto the "img" ONLY where "char_img" is white.
|
| 68 |
+
img.paste(ink_layer, (paste_x, paste_y), mask=char_img)
|
| 69 |
+
|
| 70 |
+
# 4. Final touch: Add a little bit of noise to the whole box
|
| 71 |
+
# This makes the "pure" background look more like paper texture
|
| 72 |
+
arr = np.array(img)
|
| 73 |
+
noise = np.random.randint(-5, 5, arr.shape)
|
| 74 |
+
arr = np.clip(arr + noise, 0, 255).astype(np.uint8)
|
| 75 |
+
|
| 76 |
+
return Image.fromarray(arr)
|
| 77 |
+
|
| 78 |
+
# --- EXECUTION ---
|
| 79 |
+
print(f"Generating {len(classes)} classes...")
|
| 80 |
+
|
| 81 |
+
for cls in classes:
|
| 82 |
+
class_dir = os.path.join(OUTPUT_PATH, cls)
|
| 83 |
+
os.makedirs(class_dir, exist_ok=True)
|
| 84 |
+
|
| 85 |
+
# Use fewer samples if you are just testing, increase for final training
|
| 86 |
+
for i in range(SAMPLES_PER_CLASS):
|
| 87 |
+
box_img = create_move_image(cls)
|
| 88 |
+
# We save as '72x_1.png' etc.
|
| 89 |
+
box_img.save(os.path.join(class_dir, f"{cls}_{i}.png"))
|
| 90 |
+
|
| 91 |
+
print(f"Class {cls} generated.")
|
| 92 |
+
|
| 93 |
+
print(f"\nSuccess! Generated {len(classes)} folders in {OUTPUT_PATH}")
|
legacy/b2.py
ADDED
|
@@ -0,0 +1,124 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import random
|
| 3 |
+
import glob
|
| 4 |
+
import numpy as np
|
| 5 |
+
from PIL import Image, ImageOps, ImageFilter
|
| 6 |
+
import cv2
|
| 7 |
+
# --- CONFIGURATION ---
|
| 8 |
+
INGREDIENTS_PATH = "ingredients"
|
| 9 |
+
OUTPUT_PATH = "train_data2"
|
| 10 |
+
BOX_HEIGHT = 40
|
| 11 |
+
BOX_WIDTH = 120 # 3:1 Proportion
|
| 12 |
+
SAMPLES_PER_CLASS = 300 # Adjust based on your disk space
|
| 13 |
+
|
| 14 |
+
# 1. Generate the 163 Class Names
|
| 15 |
+
move_classes = []
|
| 16 |
+
for start_hole in range(1, 10):
|
| 17 |
+
for end_hole in range(1, 10):
|
| 18 |
+
move_classes.append(f"{start_hole}{end_hole}") # e.g., "72"
|
| 19 |
+
move_classes.append(f"{start_hole}{end_hole}x") # e.g., "72x"
|
| 20 |
+
|
| 21 |
+
classes = move_classes + ['empty']
|
| 22 |
+
|
| 23 |
+
os.makedirs(OUTPUT_PATH, exist_ok=True)
|
| 24 |
+
|
| 25 |
+
def get_random_ingredient(char):
|
| 26 |
+
files = glob.glob(os.path.join(INGREDIENTS_PATH, char, "*.png"))
|
| 27 |
+
if not files:
|
| 28 |
+
raise ValueError(f"No images found for character: {char}")
|
| 29 |
+
return Image.open(random.choice(files))
|
| 30 |
+
|
| 31 |
+
def create_move_image(class_name):
|
| 32 |
+
# 1. Create a pure, solid white paper background (255)
|
| 33 |
+
# This matches the clean white background of the post-processed real photos
|
| 34 |
+
img = Image.new('L', (BOX_WIDTH, BOX_HEIGHT), color=255)
|
| 35 |
+
|
| 36 |
+
if class_name == 'empty':
|
| 37 |
+
return img
|
| 38 |
+
|
| 39 |
+
chars_to_draw = list(class_name)
|
| 40 |
+
prepared_chars = []
|
| 41 |
+
total_width = 0
|
| 42 |
+
gaps = []
|
| 43 |
+
|
| 44 |
+
# --- STEP 1: PREPARE AND MEASURE ALL CHARACTERS ---
|
| 45 |
+
for i, char in enumerate(chars_to_draw):
|
| 46 |
+
char_img = get_random_ingredient(char)
|
| 47 |
+
|
| 48 |
+
# Make 'x' slightly smaller than numbers
|
| 49 |
+
if char == 'x':
|
| 50 |
+
size = random.randint(20, 26)
|
| 51 |
+
else:
|
| 52 |
+
size = random.randint(28, 36)
|
| 53 |
+
|
| 54 |
+
char_img = char_img.resize((size, size), Image.Resampling.LANCZOS)
|
| 55 |
+
char_img = char_img.rotate(random.randint(-10, 10), expand=False, fillcolor=0)
|
| 56 |
+
|
| 57 |
+
# Crop the black space around the EMNIST digit
|
| 58 |
+
bbox = char_img.getbbox()
|
| 59 |
+
if bbox:
|
| 60 |
+
char_img = char_img.crop(bbox)
|
| 61 |
+
|
| 62 |
+
prepared_chars.append(char_img)
|
| 63 |
+
total_width += char_img.width
|
| 64 |
+
|
| 65 |
+
# Calculate random gap
|
| 66 |
+
if i < len(chars_to_draw) - 1:
|
| 67 |
+
gap = random.randint(-2, 5)
|
| 68 |
+
gaps.append(gap)
|
| 69 |
+
total_width += gap
|
| 70 |
+
|
| 71 |
+
# --- STEP 2: RANDOM TRANSLATION ---
|
| 72 |
+
max_start_x = BOX_WIDTH - total_width
|
| 73 |
+
if max_start_x <= 2:
|
| 74 |
+
start_x = 2
|
| 75 |
+
else:
|
| 76 |
+
start_x = random.randint(2, max_start_x - 2)
|
| 77 |
+
|
| 78 |
+
# --- STEP 3: PASTE THE CHARACTERS (Using Solid Black Ink) ---
|
| 79 |
+
current_x = start_x
|
| 80 |
+
for i, char_img in enumerate(prepared_chars):
|
| 81 |
+
# We use solid pure black ink (0) to match post-threshold images
|
| 82 |
+
ink_layer = Image.new('L', char_img.size, color=0)
|
| 83 |
+
|
| 84 |
+
max_y = BOX_HEIGHT - char_img.height
|
| 85 |
+
paste_y = random.randint(2, max(2, max_y - 2))
|
| 86 |
+
|
| 87 |
+
img.paste(ink_layer, (current_x, paste_y), mask=char_img)
|
| 88 |
+
|
| 89 |
+
if i < len(gaps):
|
| 90 |
+
current_x += char_img.width + gaps[i]
|
| 91 |
+
|
| 92 |
+
# --- STEP 4: DILATION ALIGNMENT ---
|
| 93 |
+
# Convert to numpy array for OpenCV
|
| 94 |
+
arr = np.array(img)
|
| 95 |
+
|
| 96 |
+
# Invert the image so the ink is white (required for dilation)
|
| 97 |
+
ink_is_white = cv2.bitwise_not(arr)
|
| 98 |
+
|
| 99 |
+
# Apply a 3x3 dilation kernel to thicken the strokes
|
| 100 |
+
# This matches the dilated thickness of the real-world processed pen strokes!
|
| 101 |
+
kernel = np.ones((3,3), np.uint8)
|
| 102 |
+
thick_ink = cv2.dilate(ink_is_white, kernel, iterations=1)
|
| 103 |
+
|
| 104 |
+
# Invert back: Ink is Black (0), Background is Pure White (255)
|
| 105 |
+
final_img = cv2.bitwise_not(thick_ink)
|
| 106 |
+
|
| 107 |
+
return Image.fromarray(final_img)
|
| 108 |
+
|
| 109 |
+
# --- EXECUTION ---
|
| 110 |
+
print(f"Generating {len(classes)} classes with Dynamic Spacing...")
|
| 111 |
+
|
| 112 |
+
for cls in classes:
|
| 113 |
+
class_dir = os.path.join(OUTPUT_PATH, cls)
|
| 114 |
+
os.makedirs(class_dir, exist_ok=True)
|
| 115 |
+
|
| 116 |
+
for i in range(SAMPLES_PER_CLASS):
|
| 117 |
+
box_img = create_move_image(cls)
|
| 118 |
+
box_img.save(os.path.join(class_dir, f"{cls}_{i}.png"))
|
| 119 |
+
|
| 120 |
+
# Optional print to track progress
|
| 121 |
+
if (classes.index(cls) + 1) % 10 == 0:
|
| 122 |
+
print(f"Generated {classes.index(cls) + 1}/{len(classes)} classes...")
|
| 123 |
+
|
| 124 |
+
print(f"\nSuccess! Generated {len(classes)} folders in {OUTPUT_PATH}")
|
legacy/b3.py
ADDED
|
@@ -0,0 +1,80 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import random
|
| 3 |
+
import glob
|
| 4 |
+
import numpy as np
|
| 5 |
+
from PIL import Image, ImageFilter
|
| 6 |
+
import cv2
|
| 7 |
+
|
| 8 |
+
INGREDIENTS_PATH = "ingredients"
|
| 9 |
+
OUTPUT_PATH = "train_data2"
|
| 10 |
+
BOX_HEIGHT = 40
|
| 11 |
+
BOX_WIDTH = 80
|
| 12 |
+
SAMPLES_PER_CLASS = 300
|
| 13 |
+
|
| 14 |
+
move_classes = [f"{s}{e}" for s in range(1, 10) for e in range(1, 10)] + \
|
| 15 |
+
[f"{s}{e}x" for s in range(1, 10) for e in range(1, 10)]
|
| 16 |
+
classes = move_classes + ['empty']
|
| 17 |
+
|
| 18 |
+
os.makedirs(OUTPUT_PATH, exist_ok=True)
|
| 19 |
+
|
| 20 |
+
def crop_and_center_ink(cv2_gray_img, target_w=80, target_h=40, margin=4):
|
| 21 |
+
ink_mask = cv2.bitwise_not(cv2_gray_img)
|
| 22 |
+
coords = cv2.findNonZero(ink_mask)
|
| 23 |
+
if coords is not None:
|
| 24 |
+
x, y, w, h = cv2.boundingRect(coords)
|
| 25 |
+
crop = cv2_gray_img[y:y+h, x:x+w]
|
| 26 |
+
scale = min((target_w - 2*margin) / w, (target_h - 2*margin) / h)
|
| 27 |
+
new_w, new_h = max(1, int(w * scale)), max(1, int(h * scale))
|
| 28 |
+
resized_crop = cv2.resize(crop, (new_w, new_h), interpolation=cv2.INTER_AREA)
|
| 29 |
+
canvas = np.full((target_h, target_w), 255, dtype=np.uint8)
|
| 30 |
+
start_x, start_y = (target_w - new_w) // 2, (target_h - new_h) // 2
|
| 31 |
+
canvas[start_y:start_y+new_h, start_x:start_x+new_w] = resized_crop
|
| 32 |
+
return canvas
|
| 33 |
+
return np.full((target_h, target_w), 255, dtype=np.uint8)
|
| 34 |
+
|
| 35 |
+
def create_move_image(class_name):
|
| 36 |
+
if class_name == 'empty':
|
| 37 |
+
return Image.new('L', (BOX_WIDTH, BOX_HEIGHT), color=255)
|
| 38 |
+
|
| 39 |
+
# Start with a clean canvas
|
| 40 |
+
img = Image.new('L', (200, 200), color=255)
|
| 41 |
+
current_x = 40
|
| 42 |
+
for char in class_name:
|
| 43 |
+
files = glob.glob(os.path.join(INGREDIENTS_PATH, char, "*.png"))
|
| 44 |
+
char_img = Image.open(random.choice(files))
|
| 45 |
+
size = random.randint(22, 28) if char == 'x' else random.randint(32, 42)
|
| 46 |
+
char_img = char_img.resize((size, size), Image.Resampling.LANCZOS)
|
| 47 |
+
char_img = char_img.rotate(random.randint(-12, 12), expand=True, fillcolor=0)
|
| 48 |
+
bbox = char_img.getbbox()
|
| 49 |
+
if bbox: char_img = char_img.crop(bbox)
|
| 50 |
+
|
| 51 |
+
ink_layer = Image.new('L', char_img.size, color=0)
|
| 52 |
+
paste_y = 100 - (char_img.height // 2) + random.randint(-5, 5)
|
| 53 |
+
img.paste(ink_layer, (current_x, paste_y), mask=char_img)
|
| 54 |
+
current_x += char_img.width + random.randint(-1, 4)
|
| 55 |
+
|
| 56 |
+
# --- THE CRITICAL FIX: Center FIRST, then Dilate ---
|
| 57 |
+
temp_arr = np.array(img)
|
| 58 |
+
centered = crop_and_center_ink(temp_arr, BOX_WIDTH, BOX_HEIGHT)
|
| 59 |
+
|
| 60 |
+
# Randomly vary thickness (Some thin, some thick)
|
| 61 |
+
ink_is_white = cv2.bitwise_not(centered)
|
| 62 |
+
thickness = random.choice([1, 2, 2]) # Weight it towards thicker lines
|
| 63 |
+
kernel = np.ones((thickness, thickness), np.uint8)
|
| 64 |
+
thick_ink = cv2.dilate(ink_is_white, kernel, iterations=1)
|
| 65 |
+
|
| 66 |
+
# Add a tiny bit of blur/noise to simulate real camera focus
|
| 67 |
+
final_img = cv2.bitwise_not(thick_ink)
|
| 68 |
+
pil_final = Image.fromarray(final_img)
|
| 69 |
+
if random.random() > 0.5:
|
| 70 |
+
pil_final = pil_final.filter(ImageFilter.GaussianBlur(radius=0.3))
|
| 71 |
+
|
| 72 |
+
return pil_final
|
| 73 |
+
|
| 74 |
+
# ... Execution loop remains same as previous ...
|
| 75 |
+
print("Generating centered data with variable thickness...")
|
| 76 |
+
for cls in classes:
|
| 77 |
+
class_dir = os.path.join(OUTPUT_PATH, cls)
|
| 78 |
+
os.makedirs(class_dir, exist_ok=True)
|
| 79 |
+
for i in range(SAMPLES_PER_CLASS):
|
| 80 |
+
create_move_image(cls).save(os.path.join(class_dir, f"{cls}_{i}.png"))
|
legacy/b4.py
ADDED
|
@@ -0,0 +1,143 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import random
|
| 3 |
+
import glob
|
| 4 |
+
import numpy as np
|
| 5 |
+
from PIL import Image, ImageFilter
|
| 6 |
+
import cv2
|
| 7 |
+
|
| 8 |
+
# --- CONFIGURATION ---
|
| 9 |
+
INGREDIENTS_PATH = "ingredients"
|
| 10 |
+
OUTPUT_PATH = "train_data5"
|
| 11 |
+
BOX_HEIGHT = 40
|
| 12 |
+
BOX_WIDTH = 80
|
| 13 |
+
SAMPLES_PER_CLASS = 300
|
| 14 |
+
|
| 15 |
+
# 1. Generate the 163 Class Names
|
| 16 |
+
move_classes = []
|
| 17 |
+
for start_hole in range(1, 10):
|
| 18 |
+
for end_hole in range(1, 10):
|
| 19 |
+
move_classes.append(f"{start_hole}{end_hole}")
|
| 20 |
+
move_classes.append(f"{start_hole}{end_hole}x")
|
| 21 |
+
|
| 22 |
+
classes = move_classes + ['empty']
|
| 23 |
+
os.makedirs(OUTPUT_PATH, exist_ok=True)
|
| 24 |
+
|
| 25 |
+
def get_random_ingredient(char):
|
| 26 |
+
files = glob.glob(os.path.join(INGREDIENTS_PATH, char, "*.png"))
|
| 27 |
+
if not files:
|
| 28 |
+
raise ValueError(f"No images found for character: {char}")
|
| 29 |
+
# Convert to grayscale to act perfectly as an alpha mask
|
| 30 |
+
return Image.open(random.choice(files)).convert('L')
|
| 31 |
+
|
| 32 |
+
def add_camera_noise(img_array):
|
| 33 |
+
"""Simulates rough paper texture and camera sensor noise"""
|
| 34 |
+
noise = np.random.randint(0, 25, img_array.shape, dtype='uint8')
|
| 35 |
+
# Subtracting noise makes random pixels slightly darker, like paper grain
|
| 36 |
+
noisy_img = np.clip(img_array.astype(int) - noise, 0, 255).astype('uint8')
|
| 37 |
+
return noisy_img
|
| 38 |
+
|
| 39 |
+
def create_move_image(class_name):
|
| 40 |
+
# --- PHILOSOPHY 1: IMPERFECT PAPER BACKGROUND ---
|
| 41 |
+
# Random RGB values mimicking paper under different lighting (off-white, warm, cool)
|
| 42 |
+
bg_gray = random.randint(200, 255)
|
| 43 |
+
|
| 44 |
+
img = Image.new('L', (BOX_WIDTH, BOX_HEIGHT), color=bg_gray)
|
| 45 |
+
|
| 46 |
+
if class_name == 'empty':
|
| 47 |
+
img_arr = add_camera_noise(np.array(img))
|
| 48 |
+
img = Image.fromarray(img_arr)
|
| 49 |
+
return img.filter(ImageFilter.GaussianBlur(radius=random.uniform(0.1, 0.5)))
|
| 50 |
+
|
| 51 |
+
# --- PHILOSOPHY 2: REAL PEN INK COLORS ---
|
| 52 |
+
# Randomly pick black, dark blue, or bright blue ink
|
| 53 |
+
|
| 54 |
+
ink_color = random.choice([0,15,30,45])
|
| 55 |
+
|
| 56 |
+
chars_to_draw = list(class_name)
|
| 57 |
+
prepared_chars = []
|
| 58 |
+
total_width = 0
|
| 59 |
+
gaps = []
|
| 60 |
+
|
| 61 |
+
for i, char in enumerate(chars_to_draw):
|
| 62 |
+
char_img = get_random_ingredient(char)
|
| 63 |
+
|
| 64 |
+
# Sizing
|
| 65 |
+
size = random.randint(14, 18) if char == 'x' else random.randint(20, 26)
|
| 66 |
+
char_img = char_img.resize((size, size), Image.Resampling.LANCZOS)
|
| 67 |
+
|
| 68 |
+
# Rotation
|
| 69 |
+
char_img = char_img.rotate(random.randint(-15, 15), expand=True, fillcolor=0)
|
| 70 |
+
|
| 71 |
+
# --- PHILOSOPHY 3: VARIABLE PEN PRESSURE ---
|
| 72 |
+
# Randomly thicken or thin the stroke using morphology
|
| 73 |
+
char_arr = np.array(char_img)
|
| 74 |
+
kernel = np.ones((2, 2), np.uint8)
|
| 75 |
+
thickness_op = random.choice(['dilate', 'erode', 'none', 'dilate']) # Bias slightly towards thicker
|
| 76 |
+
|
| 77 |
+
if thickness_op == 'dilate':
|
| 78 |
+
char_arr = cv2.dilate(char_arr, kernel, iterations=1)
|
| 79 |
+
elif thickness_op == 'erode':
|
| 80 |
+
char_arr = cv2.erode(char_arr, kernel, iterations=1)
|
| 81 |
+
|
| 82 |
+
char_img = Image.fromarray(char_arr)
|
| 83 |
+
|
| 84 |
+
# Crop tight around the character
|
| 85 |
+
bbox = char_img.getbbox()
|
| 86 |
+
if bbox:
|
| 87 |
+
char_img = char_img.crop(bbox)
|
| 88 |
+
|
| 89 |
+
prepared_chars.append(char_img)
|
| 90 |
+
total_width += char_img.width
|
| 91 |
+
|
| 92 |
+
# Gaps (Allowing negative numbers means strokes might overlap naturally!)
|
| 93 |
+
if i < len(chars_to_draw) - 1:
|
| 94 |
+
gap = random.randint(-6, 1)
|
| 95 |
+
gaps.append(gap)
|
| 96 |
+
total_width += gap
|
| 97 |
+
|
| 98 |
+
# Translation
|
| 99 |
+
max_start_x = BOX_WIDTH - total_width
|
| 100 |
+
start_x = random.randint(2, max(2, max_start_x - 2))
|
| 101 |
+
|
| 102 |
+
# Paste using the EMNIST mask
|
| 103 |
+
current_x = start_x
|
| 104 |
+
for i, char_mask in enumerate(prepared_chars):
|
| 105 |
+
# Create a solid block of our chosen ink color
|
| 106 |
+
ink_layer = Image.new('RGB', char_mask.size, color=ink_color)
|
| 107 |
+
|
| 108 |
+
max_y = BOX_HEIGHT - char_mask.height
|
| 109 |
+
paste_y = random.randint(1, max(1, max_y - 1))
|
| 110 |
+
|
| 111 |
+
# Paste the ink onto the paper, using the white EMNIST digit as the stencil
|
| 112 |
+
img.paste(ink_layer, (current_x, paste_y), mask=char_mask)
|
| 113 |
+
|
| 114 |
+
if i < len(gaps):
|
| 115 |
+
current_x += char_mask.width + gaps[i]
|
| 116 |
+
|
| 117 |
+
# --- PHILOSOPHY 4: THE "DIRTY" REALITY ---
|
| 118 |
+
# 1. Add noise
|
| 119 |
+
img_arr = np.array(img)
|
| 120 |
+
img_arr = add_camera_noise(img_arr)
|
| 121 |
+
final_img = Image.fromarray(img_arr)
|
| 122 |
+
|
| 123 |
+
# 2. Add random camera blur
|
| 124 |
+
blur_radius = random.uniform(0.1, 0.8)
|
| 125 |
+
final_img = final_img.filter(ImageFilter.GaussianBlur(radius=blur_radius))
|
| 126 |
+
|
| 127 |
+
return final_img
|
| 128 |
+
|
| 129 |
+
# --- EXECUTION ---
|
| 130 |
+
print(f"Generating {len(classes)} classes with Real-World Domain Shift...")
|
| 131 |
+
|
| 132 |
+
for cls in classes:
|
| 133 |
+
class_dir = os.path.join(OUTPUT_PATH, cls)
|
| 134 |
+
os.makedirs(class_dir, exist_ok=True)
|
| 135 |
+
|
| 136 |
+
for i in range(SAMPLES_PER_CLASS):
|
| 137 |
+
box_img = create_move_image(cls)
|
| 138 |
+
box_img.save(os.path.join(class_dir, f"{cls}_{i}.png"))
|
| 139 |
+
|
| 140 |
+
if (classes.index(cls) + 1) % 10 == 0:
|
| 141 |
+
print(f"Generated {classes.index(cls) + 1}/{len(classes)} classes...")
|
| 142 |
+
|
| 143 |
+
print(f"\nSuccess! Generated highly robust dataset in {OUTPUT_PATH}")
|
legacy/book.txt
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
1. 98 11 K1:10 K2:10 K3:20 T:00
|
| 2 |
+
1. 98 22 K1:10 K2:10 K3:20 T:00
|
| 3 |
+
1. 98 33 K1:10 K2:10 K3:20 T:00
|
| 4 |
+
1. 98 44 K1:10 K2:10 K3:20 T:00
|
| 5 |
+
1. 98 55 K1:10 K2:10 K3:20 T:00
|
| 6 |
+
1. 98 66 K1:10 K2:10 K3:20 T:00
|
| 7 |
+
1. 98 77 K1:10 K2:10 K3:20 T:00
|
| 8 |
+
1. 98 98 K1:10 K2:10 K3:20 T:00
|
legacy/c.py
ADDED
|
@@ -0,0 +1,108 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import torch.optim as optim
|
| 4 |
+
from torchvision import datasets, transforms, models
|
| 5 |
+
from torch.utils.data import DataLoader, random_split
|
| 6 |
+
import json
|
| 7 |
+
import os
|
| 8 |
+
import time
|
| 9 |
+
|
| 10 |
+
# 1. HYPERPARAMETERS
|
| 11 |
+
BATCH_SIZE = 256
|
| 12 |
+
EPOCHS = 6 # Lowered to 6 because you were already at 95% by epoch 2!
|
| 13 |
+
LEARNING_RATE = 0.001
|
| 14 |
+
DATA_DIR = "train_data3"
|
| 15 |
+
|
| 16 |
+
def main():
|
| 17 |
+
if torch.cuda.is_available():
|
| 18 |
+
device = torch.device("cuda")
|
| 19 |
+
use_pin_memory = True
|
| 20 |
+
elif torch.backends.mps.is_available():
|
| 21 |
+
device = torch.device("mps") # Apple Silicon Mac!
|
| 22 |
+
use_pin_memory = False # MPS doesn't support pin_memory yet
|
| 23 |
+
else:
|
| 24 |
+
device = torch.device("cpu")
|
| 25 |
+
use_pin_memory = False
|
| 26 |
+
|
| 27 |
+
print(f"Using device: {device}")
|
| 28 |
+
|
| 29 |
+
# 2. DATA TRANSFORMATIONS
|
| 30 |
+
# Change THIS line in c.py before you retrain!
|
| 31 |
+
transform = transforms.Compose([
|
| 32 |
+
transforms.Resize((40, 80)), # <--- Changed from 120 to 80
|
| 33 |
+
transforms.ColorJitter(brightness=0.2, contrast=0.2),
|
| 34 |
+
transforms.ToTensor(),
|
| 35 |
+
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
|
| 36 |
+
])
|
| 37 |
+
|
| 38 |
+
# 3. LOAD DATASET
|
| 39 |
+
print("Loading dataset...")
|
| 40 |
+
full_dataset = datasets.ImageFolder(root=DATA_DIR, transform=transform)
|
| 41 |
+
NUM_CLASSES = len(full_dataset.classes)
|
| 42 |
+
class_names = full_dataset.classes
|
| 43 |
+
|
| 44 |
+
print(f"Total classes detected: {NUM_CLASSES}")
|
| 45 |
+
|
| 46 |
+
with open("class_mapping.json", "w") as f:
|
| 47 |
+
json.dump(full_dataset.class_to_idx, f)
|
| 48 |
+
print("Class mapping saved to class_mapping.json")
|
| 49 |
+
|
| 50 |
+
# 4. TRAIN/VALIDATION SPLIT
|
| 51 |
+
train_size = int(0.8 * len(full_dataset))
|
| 52 |
+
val_size = len(full_dataset) - train_size
|
| 53 |
+
train_dataset, val_dataset = random_split(full_dataset, [train_size, val_size])
|
| 54 |
+
|
| 55 |
+
# The workers are safely spawned now!
|
| 56 |
+
train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True,
|
| 57 |
+
num_workers=4, pin_memory=use_pin_memory)
|
| 58 |
+
val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False,
|
| 59 |
+
num_workers=4, pin_memory=use_pin_memory)
|
| 60 |
+
|
| 61 |
+
# 5. INITIALIZE ResNet-18
|
| 62 |
+
model = models.resnet18(weights=None)
|
| 63 |
+
num_ftrs = model.fc.in_features
|
| 64 |
+
model.fc = nn.Linear(num_ftrs, NUM_CLASSES)
|
| 65 |
+
model = model.to(device)
|
| 66 |
+
|
| 67 |
+
# 6. LOSS AND OPTIMIZER
|
| 68 |
+
criterion = nn.CrossEntropyLoss()
|
| 69 |
+
optimizer = optim.Adam(model.parameters(), lr=LEARNING_RATE)
|
| 70 |
+
|
| 71 |
+
# 7. TRAINING LOOP
|
| 72 |
+
print("Starting training...")
|
| 73 |
+
start_time = time.time()
|
| 74 |
+
|
| 75 |
+
for epoch in range(EPOCHS):
|
| 76 |
+
model.train()
|
| 77 |
+
running_loss = 0.0
|
| 78 |
+
for images, labels in train_loader:
|
| 79 |
+
images, labels = images.to(device), labels.to(device)
|
| 80 |
+
optimizer.zero_grad()
|
| 81 |
+
outputs = model(images)
|
| 82 |
+
loss = criterion(outputs, labels)
|
| 83 |
+
loss.backward()
|
| 84 |
+
optimizer.step()
|
| 85 |
+
running_loss += loss.item()
|
| 86 |
+
|
| 87 |
+
# VALIDATION
|
| 88 |
+
model.eval()
|
| 89 |
+
correct = 0
|
| 90 |
+
total = 0
|
| 91 |
+
with torch.no_grad():
|
| 92 |
+
for images, labels in val_loader:
|
| 93 |
+
images, labels = images.to(device), labels.to(device)
|
| 94 |
+
outputs = model(images)
|
| 95 |
+
_, predicted = torch.max(outputs.data, 1)
|
| 96 |
+
total += labels.size(0)
|
| 97 |
+
correct += (predicted == labels).sum().item()
|
| 98 |
+
|
| 99 |
+
elapsed = time.time() - start_time
|
| 100 |
+
print(f"[{elapsed:.2f}s] Epoch [{epoch+1}/{EPOCHS}] Loss: {running_loss/len(train_loader):.4f} Val Acc: {100*correct/total:.2f}%")
|
| 101 |
+
|
| 102 |
+
# 8. SAVE MODEL
|
| 103 |
+
torch.save(model.state_dict(), "togyzkumalak_model.pth")
|
| 104 |
+
print("Model saved to togyzkumalak_model.pth. Done!")
|
| 105 |
+
|
| 106 |
+
# THIS IS THE MAGIC LINE THAT FIXES THE CRASH
|
| 107 |
+
if __name__ == '__main__':
|
| 108 |
+
main()
|
legacy/class_mapping.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{"11": 0, "11x": 1, "12": 2, "12x": 3, "13": 4, "13x": 5, "14": 6, "14x": 7, "15": 8, "15x": 9, "16": 10, "16x": 11, "17": 12, "17x": 13, "18": 14, "18x": 15, "19": 16, "19x": 17, "21": 18, "21x": 19, "22": 20, "22x": 21, "23": 22, "23x": 23, "24": 24, "24x": 25, "25": 26, "25x": 27, "26": 28, "26x": 29, "27": 30, "27x": 31, "28": 32, "28x": 33, "29": 34, "29x": 35, "31": 36, "31x": 37, "32": 38, "32x": 39, "33": 40, "33x": 41, "34": 42, "34x": 43, "35": 44, "35x": 45, "36": 46, "36x": 47, "37": 48, "37x": 49, "38": 50, "38x": 51, "39": 52, "39x": 53, "41": 54, "41x": 55, "42": 56, "42x": 57, "43": 58, "43x": 59, "44": 60, "44x": 61, "45": 62, "45x": 63, "46": 64, "46x": 65, "47": 66, "47x": 67, "48": 68, "48x": 69, "49": 70, "49x": 71, "51": 72, "51x": 73, "52": 74, "52x": 75, "53": 76, "53x": 77, "54": 78, "54x": 79, "55": 80, "55x": 81, "56": 82, "56x": 83, "57": 84, "57x": 85, "58": 86, "58x": 87, "59": 88, "59x": 89, "61": 90, "61x": 91, "62": 92, "62x": 93, "63": 94, "63x": 95, "64": 96, "64x": 97, "65": 98, "65x": 99, "66": 100, "66x": 101, "67": 102, "67x": 103, "68": 104, "68x": 105, "69": 106, "69x": 107, "71": 108, "71x": 109, "72": 110, "72x": 111, "73": 112, "73x": 113, "74": 114, "74x": 115, "75": 116, "75x": 117, "76": 118, "76x": 119, "77": 120, "77x": 121, "78": 122, "78x": 123, "79": 124, "79x": 125, "81": 126, "81x": 127, "82": 128, "82x": 129, "83": 130, "83x": 131, "84": 132, "84x": 133, "85": 134, "85x": 135, "86": 136, "86x": 137, "87": 138, "87x": 139, "88": 140, "88x": 141, "89": 142, "89x": 143, "91": 144, "91x": 145, "92": 146, "92x": 147, "93": 148, "93x": 149, "94": 150, "94x": 151, "95": 152, "95x": 153, "96": 154, "96x": 155, "97": 156, "97x": 157, "98": 158, "98x": 159, "99": 160, "99x": 161, "empty": 162}
|
legacy/d.py
ADDED
|
@@ -0,0 +1,64 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
from torchvision import transforms, models
|
| 4 |
+
from PIL import Image
|
| 5 |
+
import json
|
| 6 |
+
import sys
|
| 7 |
+
|
| 8 |
+
def main(image_path):
|
| 9 |
+
# 1. SETUP DEVICE
|
| 10 |
+
if torch.backends.mps.is_available():
|
| 11 |
+
device = torch.device("mps")
|
| 12 |
+
elif torch.cuda.is_available():
|
| 13 |
+
device = torch.device("cuda")
|
| 14 |
+
else:
|
| 15 |
+
device = torch.device("cpu")
|
| 16 |
+
|
| 17 |
+
# 2. LOAD CLASS MAPPING
|
| 18 |
+
# Reverse the mapping from { "19x": 5 } to { 5: "19x" }
|
| 19 |
+
with open("class_mapping.json", "r") as f:
|
| 20 |
+
class_to_idx = json.load(f)
|
| 21 |
+
idx_to_class = {v: k for k, v in class_to_idx.items()}
|
| 22 |
+
NUM_CLASSES = len(idx_to_class)
|
| 23 |
+
|
| 24 |
+
# 3. LOAD THE MODEL
|
| 25 |
+
model = models.resnet18(weights=None)
|
| 26 |
+
num_ftrs = model.fc.in_features
|
| 27 |
+
model.fc = nn.Linear(num_ftrs, NUM_CLASSES)
|
| 28 |
+
model.load_state_dict(torch.load("togyzkumalak_model.pth", map_location=device))
|
| 29 |
+
model = model.to(device)
|
| 30 |
+
model.eval() # Set to evaluation mode!
|
| 31 |
+
|
| 32 |
+
# 4. PREPARE THE IMAGE
|
| 33 |
+
# Notice we removed ColorJitter because we don't augment during inference!
|
| 34 |
+
transform = transforms.Compose([
|
| 35 |
+
transforms.Resize((40, 120)),
|
| 36 |
+
transforms.ToTensor(),
|
| 37 |
+
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
|
| 38 |
+
])
|
| 39 |
+
|
| 40 |
+
image = Image.open(image_path).convert('RGB') # Ensure it's 3-channel
|
| 41 |
+
image_tensor = transform(image).unsqueeze(0).to(device) # Add batch dimension
|
| 42 |
+
|
| 43 |
+
# 5. PREDICT
|
| 44 |
+
with torch.no_grad():
|
| 45 |
+
outputs = model(image_tensor)
|
| 46 |
+
# Apply Softmax to convert raw logits into percentages (0 to 1)
|
| 47 |
+
probabilities = torch.nn.functional.softmax(outputs[0], dim=0)
|
| 48 |
+
|
| 49 |
+
# 6. GET TOP 3 PREDICTIONS
|
| 50 |
+
top_probs, top_classes = torch.topk(probabilities, 3)
|
| 51 |
+
|
| 52 |
+
print(f"\n--- AI PREDICTIONS FOR '{image_path}' ---")
|
| 53 |
+
for i in range(3):
|
| 54 |
+
prob = top_probs[i].item() * 100
|
| 55 |
+
class_idx = top_classes[i].item()
|
| 56 |
+
class_name = idx_to_class[class_idx]
|
| 57 |
+
print(f"Choice {i+1}: Move '{class_name}' with {prob:.2f}% confidence")
|
| 58 |
+
print("------------------------------------------\n")
|
| 59 |
+
|
| 60 |
+
if __name__ == '__main__':
|
| 61 |
+
if len(sys.argv) < 2:
|
| 62 |
+
print("Usage: python predict.py <path_to_image>")
|
| 63 |
+
else:
|
| 64 |
+
main(sys.argv[1])
|
legacy/d2.py
ADDED
|
@@ -0,0 +1,118 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
from torchvision import transforms, models
|
| 4 |
+
from PIL import Image, ImageOps
|
| 5 |
+
import json
|
| 6 |
+
import sys
|
| 7 |
+
import cv2
|
| 8 |
+
import numpy as np
|
| 9 |
+
|
| 10 |
+
def preprocess_for_ai(image_path):
|
| 11 |
+
img = cv2.imread(image_path)
|
| 12 |
+
if img is None:
|
| 13 |
+
print(f"Error: Could not read image at {image_path}")
|
| 14 |
+
sys.exit(1)
|
| 15 |
+
|
| 16 |
+
# 1. Convert to Grayscale
|
| 17 |
+
gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
|
| 18 |
+
|
| 19 |
+
# 2. Gaussian Blur (Smooths out paper texture before thresholding)
|
| 20 |
+
blurred = cv2.GaussianBlur(gray, (5, 5), 0)
|
| 21 |
+
|
| 22 |
+
# 3. FIX HOLLOW LINES: Increase block size to 35, constant to 10.
|
| 23 |
+
# This forces it to see the whole pen stroke, not just the edges!
|
| 24 |
+
# 1. Lower the constant from 10 to 5 (This tells the math: "Don't be so aggressive at deleting faint ink!")
|
| 25 |
+
thresh = cv2.adaptiveThreshold(
|
| 26 |
+
blurred, 255,
|
| 27 |
+
cv2.ADAPTIVE_THRESH_GAUSSIAN_C,
|
| 28 |
+
cv2.THRESH_BINARY,
|
| 29 |
+
35, 5 # <-- Changed 10 to 5
|
| 30 |
+
)
|
| 31 |
+
|
| 32 |
+
ink_is_white = cv2.bitwise_not(thresh)
|
| 33 |
+
|
| 34 |
+
# 2. Bigger kernel to fill the white holes inside the '9'
|
| 35 |
+
kernel_close = np.ones((5,5), np.uint8)
|
| 36 |
+
closed_ink = cv2.morphologyEx(ink_is_white, cv2.MORPH_CLOSE, kernel_close)
|
| 37 |
+
|
| 38 |
+
# 3. DILATION (THE NEW FIX): Thicken the ink to bridge the gap in the '3'
|
| 39 |
+
kernel_dilate = np.ones((3,3), np.uint8)
|
| 40 |
+
thick_ink = cv2.dilate(closed_ink, kernel_dilate, iterations=1)
|
| 41 |
+
|
| 42 |
+
# 4. Remove noise dots (same as before, but on thick_ink)
|
| 43 |
+
contours, _ = cv2.findContours(thick_ink, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
|
| 44 |
+
for cnt in contours:
|
| 45 |
+
area = cv2.contourArea(cnt)
|
| 46 |
+
if area < 40:
|
| 47 |
+
cv2.drawContours(thick_ink, [cnt], -1, 0, -1)
|
| 48 |
+
|
| 49 |
+
# Invert back to black ink on white paper
|
| 50 |
+
final_img = cv2.bitwise_not(thick_ink)
|
| 51 |
+
|
| 52 |
+
# --- ALIGNMENT FIX (b2.py vs d2.py) ---
|
| 53 |
+
# b2.py creates un-stretched characters accurately placed in a 120x40 (3:1) bounding box.
|
| 54 |
+
# To prevent PyTorch from squashing arbitrary crops, we pad it to exactly 120x40 first!
|
| 55 |
+
pil_img = Image.fromarray(final_img).convert('RGB')
|
| 56 |
+
pil_img = ImageOps.pad(pil_img, (120, 40), color=(255, 255, 255))
|
| 57 |
+
# ---------------------------------------
|
| 58 |
+
|
| 59 |
+
# Save debug image strictly as the AI sees it
|
| 60 |
+
pil_img.save("debug_AI_eyes.png")
|
| 61 |
+
|
| 62 |
+
return pil_img
|
| 63 |
+
|
| 64 |
+
def main(image_path):
|
| 65 |
+
# SETUP DEVICE
|
| 66 |
+
if torch.backends.mps.is_available():
|
| 67 |
+
device = torch.device("mps")
|
| 68 |
+
elif torch.cuda.is_available():
|
| 69 |
+
device = torch.device("cuda")
|
| 70 |
+
else:
|
| 71 |
+
device = torch.device("cpu")
|
| 72 |
+
|
| 73 |
+
# LOAD CLASS MAPPING
|
| 74 |
+
with open("class_mapping.json", "r") as f:
|
| 75 |
+
class_to_idx = json.load(f)
|
| 76 |
+
idx_to_class = {int(v): k for k, v in class_to_idx.items()}
|
| 77 |
+
NUM_CLASSES = len(idx_to_class)
|
| 78 |
+
|
| 79 |
+
# LOAD MODEL
|
| 80 |
+
model = models.resnet18(weights=None)
|
| 81 |
+
num_ftrs = model.fc.in_features
|
| 82 |
+
model.fc = nn.Linear(num_ftrs, NUM_CLASSES)
|
| 83 |
+
model.load_state_dict(torch.load("togyzkumalak_model_OBED.pth", map_location=device))
|
| 84 |
+
model = model.to(device)
|
| 85 |
+
model.eval()
|
| 86 |
+
|
| 87 |
+
# PREPARE IMAGE USING OPENCV PREPROCESSING
|
| 88 |
+
pil_image = preprocess_for_ai(image_path)
|
| 89 |
+
|
| 90 |
+
transform = transforms.Compose([
|
| 91 |
+
transforms.Resize((40, 120)),
|
| 92 |
+
transforms.ToTensor(),
|
| 93 |
+
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
|
| 94 |
+
])
|
| 95 |
+
|
| 96 |
+
image_tensor = transform(pil_image).unsqueeze(0).to(device)
|
| 97 |
+
|
| 98 |
+
# PREDICT
|
| 99 |
+
with torch.no_grad():
|
| 100 |
+
outputs = model(image_tensor)
|
| 101 |
+
probabilities = torch.nn.functional.softmax(outputs[0], dim=0)
|
| 102 |
+
|
| 103 |
+
top_probs, top_classes = torch.topk(probabilities, 3)
|
| 104 |
+
|
| 105 |
+
print(f"\n--- AI PREDICTIONS FOR '{image_path}' ---")
|
| 106 |
+
for i in range(3):
|
| 107 |
+
prob = top_probs[i].item() * 100
|
| 108 |
+
class_idx = top_classes[i].item()
|
| 109 |
+
class_name = idx_to_class[class_idx]
|
| 110 |
+
print(f"Choice {i+1}: Move '{class_name}' with {prob:.2f}% confidence")
|
| 111 |
+
print("------------------------------------------\n")
|
| 112 |
+
print("Check 'debug_AI_eyes.png' in your folder to see the preprocessed image!")
|
| 113 |
+
|
| 114 |
+
if __name__ == '__main__':
|
| 115 |
+
if len(sys.argv) < 2:
|
| 116 |
+
print("Usage: python d2.py <path_to_image>")
|
| 117 |
+
else:
|
| 118 |
+
main(sys.argv[1])
|
legacy/d3.py
ADDED
|
@@ -0,0 +1,105 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
from torchvision import transforms, models
|
| 4 |
+
from PIL import Image
|
| 5 |
+
import json
|
| 6 |
+
import sys
|
| 7 |
+
import cv2
|
| 8 |
+
import numpy as np
|
| 9 |
+
|
| 10 |
+
# --- EXACT SAME FUNCTION FROM b2.py ---
|
| 11 |
+
def crop_and_center_ink(cv2_gray_img, target_w=80, target_h=40, margin=4):
|
| 12 |
+
ink_mask = cv2.bitwise_not(cv2_gray_img)
|
| 13 |
+
coords = cv2.findNonZero(ink_mask)
|
| 14 |
+
|
| 15 |
+
if coords is not None:
|
| 16 |
+
x, y, w, h = cv2.boundingRect(coords)
|
| 17 |
+
if w > 0 and h > 0:
|
| 18 |
+
crop = cv2_gray_img[y:y+h, x:x+w]
|
| 19 |
+
scale = min((target_w - 2*margin) / w, (target_h - 2*margin) / h)
|
| 20 |
+
new_w = max(1, int(w * scale))
|
| 21 |
+
new_h = max(1, int(h * scale))
|
| 22 |
+
|
| 23 |
+
resized_crop = cv2.resize(crop, (new_w, new_h), interpolation=cv2.INTER_AREA)
|
| 24 |
+
|
| 25 |
+
canvas = np.full((target_h, target_w), 255, dtype=np.uint8)
|
| 26 |
+
start_x = (target_w - new_w) // 2
|
| 27 |
+
start_y = (target_h - new_h) // 2
|
| 28 |
+
canvas[start_y:start_y+new_h, start_x:start_x+new_w] = resized_crop
|
| 29 |
+
return canvas
|
| 30 |
+
|
| 31 |
+
return np.full((target_h, target_w), 255, dtype=np.uint8)
|
| 32 |
+
|
| 33 |
+
# ... (imports and crop_and_center_ink stay the same) ...
|
| 34 |
+
|
| 35 |
+
def preprocess_for_ai(image_path):
|
| 36 |
+
img = cv2.imread(image_path)
|
| 37 |
+
gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
|
| 38 |
+
|
| 39 |
+
# Use a sharper threshold to keep lines crisp
|
| 40 |
+
blurred = cv2.GaussianBlur(gray, (3, 3), 0)
|
| 41 |
+
thresh = cv2.adaptiveThreshold(blurred, 255, cv2.ADAPTIVE_THRESH_GAUSSIAN_C,
|
| 42 |
+
cv2.THRESH_BINARY, 21, 7)
|
| 43 |
+
|
| 44 |
+
# Clean up small noise dots ONLY
|
| 45 |
+
ink_is_white = cv2.bitwise_not(thresh)
|
| 46 |
+
contours, _ = cv2.findContours(ink_is_white, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
|
| 47 |
+
for cnt in contours:
|
| 48 |
+
if cv2.contourArea(cnt) < 15: # Smaller threshold for noise
|
| 49 |
+
cv2.drawContours(ink_is_white, [cnt], -1, 0, -1)
|
| 50 |
+
|
| 51 |
+
# DO NOT dilate here. Let the centering handle the size.
|
| 52 |
+
final_img = cv2.bitwise_not(ink_is_white)
|
| 53 |
+
|
| 54 |
+
# Center it
|
| 55 |
+
centered_np = crop_and_center_ink(final_img, target_w=80, target_h=40)
|
| 56 |
+
|
| 57 |
+
# Optional: Apply a standard 2x2 dilation to match training "thick" mode
|
| 58 |
+
# if the handwriting is very thin.
|
| 59 |
+
# centered_np = cv2.erode(centered_np, np.ones((2,2), np.uint8))
|
| 60 |
+
|
| 61 |
+
pil_img = Image.fromarray(centered_np).convert('RGB')
|
| 62 |
+
pil_img.save("debug_AI_eyes.png")
|
| 63 |
+
return pil_img
|
| 64 |
+
|
| 65 |
+
# ... (main function stays same as previous) ...
|
| 66 |
+
|
| 67 |
+
def main(image_path):
|
| 68 |
+
device = torch.device("mps") if torch.backends.mps.is_available() else torch.device("cpu")
|
| 69 |
+
|
| 70 |
+
with open("class_mapping.json", "r") as f:
|
| 71 |
+
class_to_idx = json.load(f)
|
| 72 |
+
idx_to_class = {int(v): k for k, v in class_to_idx.items()}
|
| 73 |
+
NUM_CLASSES = len(idx_to_class)
|
| 74 |
+
|
| 75 |
+
model = models.resnet18(weights=None)
|
| 76 |
+
model.fc = nn.Linear(model.fc.in_features, NUM_CLASSES)
|
| 77 |
+
model.load_state_dict(torch.load("togyzkumalak_model.pth", map_location=device))
|
| 78 |
+
model = model.to(device)
|
| 79 |
+
model.eval()
|
| 80 |
+
|
| 81 |
+
pil_image = preprocess_for_ai(image_path)
|
| 82 |
+
|
| 83 |
+
transform = transforms.Compose([
|
| 84 |
+
transforms.Resize((40, 80)), # MUST match the 2:1 ratio
|
| 85 |
+
transforms.ToTensor(),
|
| 86 |
+
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])
|
| 87 |
+
])
|
| 88 |
+
|
| 89 |
+
image_tensor = transform(pil_image).unsqueeze(0).to(device)
|
| 90 |
+
|
| 91 |
+
with torch.no_grad():
|
| 92 |
+
outputs = model(image_tensor)
|
| 93 |
+
probabilities = torch.nn.functional.softmax(outputs[0], dim=0)
|
| 94 |
+
|
| 95 |
+
top_probs, top_classes = torch.topk(probabilities, 3)
|
| 96 |
+
|
| 97 |
+
print(f"\n--- AI PREDICTIONS FOR '{image_path}' ---")
|
| 98 |
+
for i in range(3):
|
| 99 |
+
print(f"Choice {i+1}: Move '{idx_to_class[top_classes[i].item()]}' with {top_probs[i].item() * 100:.2f}% conf")
|
| 100 |
+
|
| 101 |
+
if __name__ == '__main__':
|
| 102 |
+
if len(sys.argv) < 2:
|
| 103 |
+
print("Usage: python d2.py <path_to_image>")
|
| 104 |
+
else:
|
| 105 |
+
main(sys.argv[1])
|
legacy/d4.py
ADDED
|
@@ -0,0 +1,78 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
os.environ["OMP_NUM_THREADS"] = "1"
|
| 3 |
+
os.environ["MKL_NUM_THREADS"] = "1"
|
| 4 |
+
os.environ["VECLIB_MAXIMUM_THREADS"] = "1"
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
torch.set_num_threads(1)
|
| 8 |
+
|
| 9 |
+
import torchvision.transforms as T
|
| 10 |
+
from PIL import Image
|
| 11 |
+
from fastai.vision.all import load_learner
|
| 12 |
+
import pickle
|
| 13 |
+
import plum._resolver
|
| 14 |
+
|
| 15 |
+
# --- THE MODEL LOADING PATCH ---
|
| 16 |
+
class PatchedResolver(plum._resolver.Resolver):
|
| 17 |
+
def __setstate__(self, state):
|
| 18 |
+
if isinstance(state, dict):
|
| 19 |
+
for k, v in state.items():
|
| 20 |
+
try: setattr(self, k, v)
|
| 21 |
+
except AttributeError: pass
|
| 22 |
+
elif isinstance(state, tuple):
|
| 23 |
+
for item in state:
|
| 24 |
+
if isinstance(item, dict):
|
| 25 |
+
for k, v in item.items():
|
| 26 |
+
try: setattr(self, k, v)
|
| 27 |
+
except AttributeError: pass
|
| 28 |
+
|
| 29 |
+
class SafeUnpickler(pickle.Unpickler):
|
| 30 |
+
def find_class(self, module, name):
|
| 31 |
+
if name == 'Resolver' and 'plum' in module:
|
| 32 |
+
return PatchedResolver
|
| 33 |
+
return super().find_class(module, name)
|
| 34 |
+
|
| 35 |
+
class SafePickle:
|
| 36 |
+
Unpickler = SafeUnpickler
|
| 37 |
+
# -------------------------------
|
| 38 |
+
|
| 39 |
+
# 1. Load the model using our safe unpickler
|
| 40 |
+
print("Loading model...")
|
| 41 |
+
learn = load_learner('handwriting_classifier.pkl', pickle_module=SafePickle)
|
| 42 |
+
print("Model loaded successfully!")
|
| 43 |
+
|
| 44 |
+
# 2. Extract the raw PyTorch model and the classes (vocab)
|
| 45 |
+
pytorch_model = learn.model
|
| 46 |
+
pytorch_model.eval() # Set model to evaluation mode
|
| 47 |
+
vocab = list(learn.dls.vocab) # Get the list of your 163 classes
|
| 48 |
+
|
| 49 |
+
# 3. Standard PyTorch Image Preprocessing (Bypasses FastAI transforms entirely)
|
| 50 |
+
# This perfectly matches the Resize and standard ImageNet normalization used by ResNet
|
| 51 |
+
preprocess = T.Compose([
|
| 52 |
+
T.Resize((40, 120)),
|
| 53 |
+
T.ToTensor(),
|
| 54 |
+
T.Normalize(
|
| 55 |
+
mean=[0.485, 0.456, 0.406], # Standard ImageNet stats
|
| 56 |
+
std=[0.229, 0.224, 0.225]
|
| 57 |
+
)
|
| 58 |
+
])
|
| 59 |
+
|
| 60 |
+
# 4. Open the image using Pillow and preprocess it
|
| 61 |
+
img = Image.open('IMG_1521 (4) (1).jpg').convert('RGB')
|
| 62 |
+
input_tensor = preprocess(img).unsqueeze(0) # Add batch dimension [1, 3, 40, 120]
|
| 63 |
+
|
| 64 |
+
# 5. Raw PyTorch Forward Pass
|
| 65 |
+
print("Running raw PyTorch inference...")
|
| 66 |
+
with torch.no_grad():
|
| 67 |
+
outputs = pytorch_model(input_tensor)
|
| 68 |
+
# Apply softmax to convert raw outputs to probabilities
|
| 69 |
+
probabilities = torch.nn.functional.softmax(outputs[0], dim=0)
|
| 70 |
+
|
| 71 |
+
# 6. Extract the top 3 predictions
|
| 72 |
+
top_prob, top_catid = torch.topk(probabilities, 3)
|
| 73 |
+
|
| 74 |
+
print("\nPrediction Results:")
|
| 75 |
+
for i in range(top_prob.size(0)):
|
| 76 |
+
class_name = vocab[top_catid[i].item()]
|
| 77 |
+
confidence = top_prob[i].item() * 100
|
| 78 |
+
print(f"Class {class_name}: {confidence:.2f}%")
|
legacy/d5.py
ADDED
|
@@ -0,0 +1,80 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
os.environ["OMP_NUM_THREADS"] = "1"
|
| 3 |
+
os.environ["MKL_NUM_THREADS"] = "1"
|
| 4 |
+
os.environ["VECLIB_MAXIMUM_THREADS"] = "1"
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
torch.set_num_threads(1)
|
| 8 |
+
|
| 9 |
+
import torchvision.transforms as T
|
| 10 |
+
from PIL import Image, ImageOps
|
| 11 |
+
from fastai.vision.all import load_learner
|
| 12 |
+
import pickle
|
| 13 |
+
import plum._resolver
|
| 14 |
+
|
| 15 |
+
# --- THE MODEL LOADING PATCH ---
|
| 16 |
+
class PatchedResolver(plum._resolver.Resolver):
|
| 17 |
+
def __setstate__(self, state):
|
| 18 |
+
if isinstance(state, dict):
|
| 19 |
+
for k, v in state.items():
|
| 20 |
+
try: setattr(self, k, v)
|
| 21 |
+
except AttributeError: pass
|
| 22 |
+
elif isinstance(state, tuple):
|
| 23 |
+
for item in state:
|
| 24 |
+
if isinstance(item, dict):
|
| 25 |
+
for k, v in item.items():
|
| 26 |
+
try: setattr(self, k, v)
|
| 27 |
+
except AttributeError: pass
|
| 28 |
+
|
| 29 |
+
class SafeUnpickler(pickle.Unpickler):
|
| 30 |
+
def find_class(self, module, name):
|
| 31 |
+
if name == 'Resolver' and 'plum' in module:
|
| 32 |
+
return PatchedResolver
|
| 33 |
+
return super().find_class(module, name)
|
| 34 |
+
|
| 35 |
+
class SafePickle:
|
| 36 |
+
Unpickler = SafeUnpickler
|
| 37 |
+
# -------------------------------
|
| 38 |
+
|
| 39 |
+
# 1. Load the model using our safe unpickler
|
| 40 |
+
print("Loading model...")
|
| 41 |
+
learn = load_learner('handwriting_classifier_best (1).pkl', pickle_module=SafePickle)
|
| 42 |
+
print("Model loaded successfully!")
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
# 2. Extract the raw PyTorch model and the classes (vocab)
|
| 46 |
+
pytorch_model = learn.model
|
| 47 |
+
pytorch_model.eval() # Set model to evaluation mode
|
| 48 |
+
vocab = list(learn.dls.vocab) # Get the list of your 163 classes
|
| 49 |
+
|
| 50 |
+
# 3. Standard PyTorch Image Preprocessing (Bypasses FastAI transforms entirely)
|
| 51 |
+
# This perfectly matches the Resize and standard ImageNet normalization used by ResNet
|
| 52 |
+
preprocess = T.Compose([
|
| 53 |
+
T.Resize((40, 80)),
|
| 54 |
+
T.ToTensor(),
|
| 55 |
+
T.Normalize(
|
| 56 |
+
mean=[0.485, 0.456, 0.406], # Standard ImageNet stats
|
| 57 |
+
std=[0.229, 0.224, 0.225]
|
| 58 |
+
)
|
| 59 |
+
])
|
| 60 |
+
|
| 61 |
+
# 4. Open the image using Pillow and preprocess it
|
| 62 |
+
img = Image.open('IMG_1521 (1) (1) (1).jpg')
|
| 63 |
+
img = ImageOps.exif_transpose(img).convert('L').convert('RGB') # Fix iPhone rotation bug by applying EXIF orientation
|
| 64 |
+
input_tensor = preprocess(img).unsqueeze(0) # Add batch dimension [1, 3, 40, 80]
|
| 65 |
+
|
| 66 |
+
# 5. Raw PyTorch Forward Pass
|
| 67 |
+
print("Running raw PyTorch inference...")
|
| 68 |
+
with torch.no_grad():
|
| 69 |
+
outputs = pytorch_model(input_tensor)
|
| 70 |
+
# Apply softmax to convert raw outputs to probabilities
|
| 71 |
+
probabilities = torch.nn.functional.softmax(outputs[0], dim=0)
|
| 72 |
+
|
| 73 |
+
# 6. Extract the top 3 predictions
|
| 74 |
+
top_prob, top_catid = torch.topk(probabilities, 3)
|
| 75 |
+
|
| 76 |
+
print("\nPrediction Results:")
|
| 77 |
+
for i in range(top_prob.size(0)):
|
| 78 |
+
class_name = vocab[top_catid[i].item()]
|
| 79 |
+
confidence = top_prob[i].item() * 100
|
| 80 |
+
print(f"Class {class_name}: {confidence:.2f}%")
|
legacy/debug_AI_eyes.png
ADDED
|
legacy/e4.py
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from fastai.vision.all import load_learner, PILImage
|
| 2 |
+
from PIL import ImageOps
|
| 3 |
+
|
| 4 |
+
# 1. Load the model
|
| 5 |
+
print("Loading model...")
|
| 6 |
+
learn = load_learner('handwriting_classifier_best.pkl')
|
| 7 |
+
|
| 8 |
+
# 2. Open the image and FIX the iPhone rotation bug!
|
| 9 |
+
raw_img = PILImage.create('IMG_1521 (4) (1).jpg')
|
| 10 |
+
|
| 11 |
+
# FastAI's PILImage.create actually attempts to transpose EXIF under the hood,
|
| 12 |
+
# but let's double check by physically saving what the model sees:
|
| 13 |
+
raw_img.save("what_the_model_actually_sees.png")
|
| 14 |
+
print("Saved 'what_the_model_actually_sees.png'. Please open this file and check if it is upright!")
|
| 15 |
+
|
| 16 |
+
# 3. Predict using FastAI's official pipeline
|
| 17 |
+
pred, pred_idx, probs = learn.predict(raw_img)
|
| 18 |
+
|
| 19 |
+
print(f"\n--- Prediction ---")
|
| 20 |
+
print(f"Predicted Class: {pred}")
|
| 21 |
+
print(f"Confidence: {probs[pred_idx]*100:.2f}%")
|
| 22 |
+
|
| 23 |
+
# Show top 3
|
| 24 |
+
class_probabilities = dict(zip(learn.dls.vocab, probs.tolist()))
|
| 25 |
+
sorted_guesses = sorted(class_probabilities.items(), key=lambda x: x[1], reverse=True)
|
| 26 |
+
print("\nTop 3 Guesses:")
|
| 27 |
+
for cls, prob in sorted_guesses[:3]:
|
| 28 |
+
print(f"Class {cls}: {prob * 100:.2f}%")
|
legacy/image copy.png
ADDED
|
legacy/image.png
ADDED
|
legacy/what_the_model_actually_sees.png
ADDED
|
models/best.classes.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
["11", "11x", "12", "12x", "13", "13x", "14", "14x", "15", "15x", "16", "16x", "17", "17x", "18", "18x", "19", "19x", "21", "21x", "22", "22x", "23", "23x", "24", "24x", "25", "25x", "26", "26x", "27", "27x", "28", "28x", "29", "29x", "31", "31x", "32", "32x", "33", "33x", "34", "34x", "35", "35x", "36", "36x", "37", "37x", "38", "38x", "39", "39x", "41", "41x", "42", "42x", "43", "43x", "44", "44x", "45", "45x", "46", "46x", "47", "47x", "48", "48x", "49", "49x", "51", "51x", "52", "52x", "53", "53x", "54", "54x", "55", "55x", "56", "56x", "57", "57x", "58", "58x", "59", "59x", "61", "61x", "62", "62x", "63", "63x", "64", "64x", "65", "65x", "66", "66x", "67", "67x", "68", "68x", "69", "69x", "71", "71x", "72", "72x", "73", "73x", "74", "74x", "75", "75x", "76", "76x", "77", "77x", "78", "78x", "79", "79x", "81", "81x", "82", "82x", "83", "83x", "84", "84x", "85", "85x", "86", "86x", "87", "87x", "88", "88x", "89", "89x", "91", "91x", "92", "92x", "93", "93x", "94", "94x", "95", "95x", "96", "96x", "97", "97x", "98", "98x", "99", "99x", "empty"]
|
models/best.onnx
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:4549967e71873eba598076908be7c42b17926f75c978adf8b18abe899f98f386
|
| 3 |
+
size 45005946
|
models/kazan.classes.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
["00", "01", "02", "03", "04", "05", "06", "07", "08", "09", "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", "empty"]
|
models/kazan.onnx
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:889c53f835d12377b4da8455ca8ab282d0c9582b38a8d7b4a880f1a7572db1ea
|
| 3 |
+
size 44841783
|
notebooks/colab_train.ipynb
ADDED
|
@@ -0,0 +1,93 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"cells": [
|
| 3 |
+
{
|
| 4 |
+
"cell_type": "markdown",
|
| 5 |
+
"metadata": {},
|
| 6 |
+
"source": [
|
| 7 |
+
"# Togyzkumalak Move OCR — Colab training\n",
|
| 8 |
+
"\n",
|
| 9 |
+
"Runtime → Change runtime type → **GPU**, then run the cells in order.\n",
|
| 10 |
+
"\n",
|
| 11 |
+
"Get the project into Colab either by cloning your repo or by uploading a zip\n",
|
| 12 |
+
"of the project directory (without `venv311/`, `train_data5/`, `legacy/`)."
|
| 13 |
+
]
|
| 14 |
+
},
|
| 15 |
+
{
|
| 16 |
+
"cell_type": "code",
|
| 17 |
+
"execution_count": null,
|
| 18 |
+
"metadata": {},
|
| 19 |
+
"outputs": [],
|
| 20 |
+
"source": [
|
| 21 |
+
"# Option A: clone (edit URL)\n",
|
| 22 |
+
"# !git clone https://github.com/<you>/9OCR.git project\n",
|
| 23 |
+
"\n",
|
| 24 |
+
"# Option B: upload a zip named project.zip via the Files sidebar, then:\n",
|
| 25 |
+
"# !unzip -q project.zip -d project\n",
|
| 26 |
+
"\n",
|
| 27 |
+
"%cd project"
|
| 28 |
+
]
|
| 29 |
+
},
|
| 30 |
+
{
|
| 31 |
+
"cell_type": "code",
|
| 32 |
+
"execution_count": null,
|
| 33 |
+
"metadata": {},
|
| 34 |
+
"outputs": [],
|
| 35 |
+
"source": [
|
| 36 |
+
"!pip install -q -r requirements.txt"
|
| 37 |
+
]
|
| 38 |
+
},
|
| 39 |
+
{
|
| 40 |
+
"cell_type": "code",
|
| 41 |
+
"execution_count": null,
|
| 42 |
+
"metadata": {},
|
| 43 |
+
"outputs": [],
|
| 44 |
+
"source": [
|
| 45 |
+
"# Build glyph pools: ARDIS (European handwriting) + EMNIST via torchvision\n",
|
| 46 |
+
"# (--emnist is needed because ingredients/ is gitignored and absent here)\n",
|
| 47 |
+
"!python scripts/get_glyphs.py --emnist"
|
| 48 |
+
]
|
| 49 |
+
},
|
| 50 |
+
{
|
| 51 |
+
"cell_type": "code",
|
| 52 |
+
"execution_count": null,
|
| 53 |
+
"metadata": {},
|
| 54 |
+
"outputs": [],
|
| 55 |
+
"source": [
|
| 56 |
+
"# Sanity check: synthetic cells should resemble data/real_crops/*.jpg\n",
|
| 57 |
+
"!python -m togyz.synth --preview preview.png\n",
|
| 58 |
+
"from IPython.display import Image as IPImage\n",
|
| 59 |
+
"IPImage('preview.png')"
|
| 60 |
+
]
|
| 61 |
+
},
|
| 62 |
+
{
|
| 63 |
+
"cell_type": "code",
|
| 64 |
+
"execution_count": null,
|
| 65 |
+
"metadata": {},
|
| 66 |
+
"outputs": [],
|
| 67 |
+
"source": [
|
| 68 |
+
"!python train.py --epochs 30 --samples-per-epoch 50000 --batch-size 256 --workers 2"
|
| 69 |
+
]
|
| 70 |
+
},
|
| 71 |
+
{
|
| 72 |
+
"cell_type": "code",
|
| 73 |
+
"execution_count": null,
|
| 74 |
+
"metadata": {},
|
| 75 |
+
"outputs": [],
|
| 76 |
+
"source": [
|
| 77 |
+
"!python eval.py --ckpt checkpoints/best.pt\n",
|
| 78 |
+
"\n",
|
| 79 |
+
"# download the trained model\n",
|
| 80 |
+
"from google.colab import files\n",
|
| 81 |
+
"files.download('checkpoints/best.pt')"
|
| 82 |
+
]
|
| 83 |
+
}
|
| 84 |
+
],
|
| 85 |
+
"metadata": {
|
| 86 |
+
"accelerator": "GPU",
|
| 87 |
+
"colab": {"provenance": []},
|
| 88 |
+
"kernelspec": {"display_name": "Python 3", "name": "python3"},
|
| 89 |
+
"language_info": {"name": "python"}
|
| 90 |
+
},
|
| 91 |
+
"nbformat": 4,
|
| 92 |
+
"nbformat_minor": 4
|
| 93 |
+
}
|
predict.py
ADDED
|
@@ -0,0 +1,68 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Classify cropped cell images.
|
| 2 |
+
|
| 3 |
+
python predict.py data/real_crops/96_a.jpg
|
| 4 |
+
python predict.py "data/real_crops/*.jpg" --topk 3
|
| 5 |
+
python predict.py cell.jpg --allowed "12,34x,56" # restrict to legal moves
|
| 6 |
+
|
| 7 |
+
--allowed renormalizes probabilities over the given moves - at any game state
|
| 8 |
+
at most 9 moves are legal, so an upstream game tracker can pass them here.
|
| 9 |
+
"""
|
| 10 |
+
|
| 11 |
+
import argparse
|
| 12 |
+
import glob
|
| 13 |
+
|
| 14 |
+
import numpy as np
|
| 15 |
+
import torch
|
| 16 |
+
|
| 17 |
+
from togyz.classes import CLASS_TO_IDX
|
| 18 |
+
from togyz.model import auto_device, load_checkpoint
|
| 19 |
+
from togyz.preprocess import preprocess_file
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
def main() -> None:
|
| 23 |
+
parser = argparse.ArgumentParser(description=__doc__)
|
| 24 |
+
parser.add_argument("images", nargs="+", help="image paths or globs")
|
| 25 |
+
parser.add_argument("--ckpt", default="checkpoints/best.pt")
|
| 26 |
+
parser.add_argument("--topk", type=int, default=3)
|
| 27 |
+
parser.add_argument("--allowed", default=None,
|
| 28 |
+
help="comma-separated legal moves, e.g. '12,34x,56'")
|
| 29 |
+
parser.add_argument("--device", default=None)
|
| 30 |
+
args = parser.parse_args()
|
| 31 |
+
|
| 32 |
+
paths = []
|
| 33 |
+
for pattern in args.images:
|
| 34 |
+
matched = sorted(glob.glob(pattern))
|
| 35 |
+
paths.extend(matched if matched else [pattern])
|
| 36 |
+
|
| 37 |
+
device = torch.device(args.device) if args.device else auto_device()
|
| 38 |
+
model, ckpt = load_checkpoint(args.ckpt, device)
|
| 39 |
+
classes = ckpt["classes"]
|
| 40 |
+
|
| 41 |
+
allowed_idx = None
|
| 42 |
+
if args.allowed:
|
| 43 |
+
moves = [m.strip() for m in args.allowed.split(",") if m.strip()]
|
| 44 |
+
unknown = [m for m in moves if m not in CLASS_TO_IDX]
|
| 45 |
+
if unknown:
|
| 46 |
+
parser.error(f"unknown moves in --allowed: {unknown}")
|
| 47 |
+
allowed_idx = torch.tensor([CLASS_TO_IDX[m] for m in moves])
|
| 48 |
+
|
| 49 |
+
batch = torch.from_numpy(np.stack([preprocess_file(p) for p in paths]))
|
| 50 |
+
with torch.no_grad():
|
| 51 |
+
probs = torch.softmax(model(batch.to(device)).cpu(), dim=1)
|
| 52 |
+
|
| 53 |
+
for path, p in zip(paths, probs):
|
| 54 |
+
topk = p.topk(min(args.topk, len(classes)))
|
| 55 |
+
guesses = ", ".join(
|
| 56 |
+
f"{classes[i]} {v:.1%}" for i, v in zip(topk.indices.tolist(), topk.values.tolist())
|
| 57 |
+
)
|
| 58 |
+
line = f"{path}: {guesses}"
|
| 59 |
+
if allowed_idx is not None:
|
| 60 |
+
legal = p[allowed_idx]
|
| 61 |
+
legal = legal / legal.sum()
|
| 62 |
+
best = int(legal.argmax())
|
| 63 |
+
line += f" | legal pick: {classes[int(allowed_idx[best])]} {legal[best]:.1%}"
|
| 64 |
+
print(line)
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
if __name__ == "__main__":
|
| 68 |
+
main()
|
read_game.py
ADDED
|
@@ -0,0 +1,95 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Read an entire Togyzkumalak scoresheet photo into game records (ONNX).
|
| 2 |
+
|
| 3 |
+
python read_game.py "data/2026-07-06 00.00.20.jpg" --out out/sheet1 --result 0-1
|
| 4 |
+
|
| 5 |
+
Outputs in --out:
|
| 6 |
+
game.json per ply: bbox, probabilities for all 163 classes, top-k,
|
| 7 |
+
raw argmax, legal move set, legality flag
|
| 8 |
+
raw.pgn pure classifier argmax for every ply (even if illegal)
|
| 9 |
+
legal.pgn replayed under the rules; STOPS at the first illegal argmax,
|
| 10 |
+
the first empty cell, or when the game is over
|
| 11 |
+
beam.pgn best fully-legal reconstruction (beam search + kazan/result
|
| 12 |
+
evidence)
|
| 13 |
+
annotated.jpg sheet with cell boxes and the beam reconstruction labels
|
| 14 |
+
cells/ every scanned cell crop, e.g. 07_W.png
|
| 15 |
+
|
| 16 |
+
Inference runs on the exported ONNX models (torch-free). Regenerate them with
|
| 17 |
+
`python scripts/export_onnx.py` after training. PGN move annotations: '+' =
|
| 18 |
+
capture, 'x' = tuzdyk creation; strip '+' to feed the moves to the 9Q engine.
|
| 19 |
+
"""
|
| 20 |
+
|
| 21 |
+
import argparse
|
| 22 |
+
import json
|
| 23 |
+
from pathlib import Path
|
| 24 |
+
|
| 25 |
+
from togyz.pipeline import RESULT_CODES, load_classifier, run_pipeline
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def main() -> None:
|
| 29 |
+
parser = argparse.ArgumentParser(description=__doc__)
|
| 30 |
+
parser.add_argument("image", help="scoresheet photo")
|
| 31 |
+
parser.add_argument("--onnx", default="checkpoints/best.onnx",
|
| 32 |
+
help="move classifier ONNX (with a .classes.json sidecar)")
|
| 33 |
+
parser.add_argument("--kazan-onnx", default="checkpoints/kazan/best.onnx",
|
| 34 |
+
help="kazan-number ONNX; checkpoint matching is skipped "
|
| 35 |
+
"if the file does not exist")
|
| 36 |
+
parser.add_argument("--out", default=None, help="output dir (default: out/<image stem>)")
|
| 37 |
+
parser.add_argument("--topk", type=int, default=5)
|
| 38 |
+
parser.add_argument("--beam-width", type=int, default=1024,
|
| 39 |
+
help="hypotheses kept during beam decoding")
|
| 40 |
+
parser.add_argument("--beam-top", type=int, default=9,
|
| 41 |
+
help="legal continuations considered per ply (9 = all)")
|
| 42 |
+
parser.add_argument("--result", choices=sorted(RESULT_CODES),
|
| 43 |
+
help="known game result from the sheet footer "
|
| 44 |
+
"(1-0 = Bast./White won); re-ranks the beam pool")
|
| 45 |
+
parser.add_argument("--no-tta", action="store_true",
|
| 46 |
+
help="disable test-time augmentation (7 shifted views/cell)")
|
| 47 |
+
parser.add_argument("--temperature", type=float, default=1.0,
|
| 48 |
+
help="softmax temperature; >1 softens overconfident cells")
|
| 49 |
+
args = parser.parse_args()
|
| 50 |
+
|
| 51 |
+
out_dir = Path(args.out or Path("out") / Path(args.image).stem)
|
| 52 |
+
out_dir.mkdir(parents=True, exist_ok=True)
|
| 53 |
+
|
| 54 |
+
moves_clf = load_classifier(args.onnx)
|
| 55 |
+
kazan_clf = None
|
| 56 |
+
if Path(args.kazan_onnx).exists():
|
| 57 |
+
kazan_clf = load_classifier(args.kazan_onnx)
|
| 58 |
+
else:
|
| 59 |
+
print(f"No kazan classifier at {args.kazan_onnx} - checkpoint matching off.")
|
| 60 |
+
|
| 61 |
+
print(f"Reading {args.image} ...")
|
| 62 |
+
out = run_pipeline(
|
| 63 |
+
args.image, moves_clf, kazan_clf,
|
| 64 |
+
result=args.result, topk=args.topk,
|
| 65 |
+
beam_width=args.beam_width, per_ply=args.beam_top,
|
| 66 |
+
temperature=args.temperature, tta=not args.no_tta,
|
| 67 |
+
save_cells_dir=out_dir / "cells",
|
| 68 |
+
)
|
| 69 |
+
|
| 70 |
+
if out["low_resolution"]:
|
| 71 |
+
print(f"WARNING: median cell height is only {out['median_cell_height']}px - "
|
| 72 |
+
"accuracy suffers at this resolution; re-photograph at full camera "
|
| 73 |
+
"resolution if possible.")
|
| 74 |
+
if out["kazan_report"]:
|
| 75 |
+
print("Kazan checkpoints read: "
|
| 76 |
+
f"{[(r['move'], r['side'], r['read']) for r in out['kazan_report']]}")
|
| 77 |
+
|
| 78 |
+
result = {"image": args.image, "onnx": args.onnx, **out["game_json"]}
|
| 79 |
+
(out_dir / "game.json").write_text(json.dumps(result, indent=1))
|
| 80 |
+
(out_dir / "raw.pgn").write_text(out["raw_pgn"])
|
| 81 |
+
(out_dir / "legal.pgn").write_text(out["legal_pgn"])
|
| 82 |
+
(out_dir / "beam.pgn").write_text(out["beam_pgn"])
|
| 83 |
+
out["annotated_image"].save(out_dir / "annotated.jpg", quality=90)
|
| 84 |
+
|
| 85 |
+
beam = out["game_json"]["beam"]
|
| 86 |
+
agree = sum(d["agrees_with_raw"] for d in beam["moves"])
|
| 87 |
+
print(f"Scanned {out['plies_scanned']} plies; strict legal replay covers {out['legal_plies']}.")
|
| 88 |
+
print(f"Beam decode: {out['beam_plies']} fully legal plies "
|
| 89 |
+
f"(log-prob {beam['log_prob']:.1f}, agrees with raw argmax on {agree}/{out['beam_plies']}).")
|
| 90 |
+
print(f"Stopped: {out['stopped']}")
|
| 91 |
+
print(f"Outputs in {out_dir}/: game.json, raw.pgn, legal.pgn, beam.pgn, annotated.jpg, cells/")
|
| 92 |
+
|
| 93 |
+
|
| 94 |
+
if __name__ == "__main__":
|
| 95 |
+
main()
|
requirements-train.txt
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Training + ONNX export deps (local / Colab GPU). The Space itself does NOT
|
| 2 |
+
# install these - it only needs the torch-free serving set in requirements.txt.
|
| 3 |
+
torch>=2.1
|
| 4 |
+
torchvision>=0.16
|
| 5 |
+
onnx>=1.16
|
| 6 |
+
onnxscript>=0.1
|
| 7 |
+
numpy>=1.24
|
| 8 |
+
opencv-python-headless>=4.8
|
| 9 |
+
pillow>=10.0
|
requirements.txt
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# HuggingFace Space (serving) deps - torch-free, installed by the Space builder.
|
| 2 |
+
# Training/export uses requirements-train.txt instead.
|
| 3 |
+
gradio>=6.20
|
| 4 |
+
onnxruntime>=1.20
|
| 5 |
+
numpy>=1.24
|
| 6 |
+
# pin below OpenCV 5.x: the pipeline's cv2 calls are tested against 4.x
|
| 7 |
+
opencv-python-headless>=4.8,<5
|
| 8 |
+
pillow>=10.0
|
scripts/export_onnx.py
ADDED
|
@@ -0,0 +1,82 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Export the trained .pt checkpoints to ONNX for torch-free serving.
|
| 2 |
+
|
| 3 |
+
python scripts/export_onnx.py
|
| 4 |
+
|
| 5 |
+
Writes, next to each source checkpoint:
|
| 6 |
+
checkpoints/best.onnx + checkpoints/best.classes.json
|
| 7 |
+
checkpoints/kazan/best.onnx + checkpoints/kazan/best.classes.json
|
| 8 |
+
|
| 9 |
+
ONNX carries no Python metadata, so the class list (which maps output index
|
| 10 |
+
-> move/kazan string) is saved as a JSON sidecar the serving code reads. The
|
| 11 |
+
graph has a dynamic batch axis so one session handles any number of cells.
|
| 12 |
+
"""
|
| 13 |
+
|
| 14 |
+
import argparse
|
| 15 |
+
import json
|
| 16 |
+
from pathlib import Path
|
| 17 |
+
|
| 18 |
+
import onnx
|
| 19 |
+
import torch
|
| 20 |
+
|
| 21 |
+
from togyz.model import build_model
|
| 22 |
+
from togyz.preprocess import TARGET_H, TARGET_W
|
| 23 |
+
|
| 24 |
+
OPSET = 17
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
def export_one(ckpt_path: Path) -> Path:
|
| 28 |
+
ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=True)
|
| 29 |
+
classes = ckpt["classes"]
|
| 30 |
+
model = build_model(len(classes))
|
| 31 |
+
model.load_state_dict(ckpt["model_state"])
|
| 32 |
+
model.eval()
|
| 33 |
+
|
| 34 |
+
onnx_path = ckpt_path.with_suffix(".onnx")
|
| 35 |
+
classes_path = ckpt_path.with_suffix(".classes.json")
|
| 36 |
+
# Legacy TorchScript exporter (dynamo=False) embeds weights directly, so
|
| 37 |
+
# the result is a single self-contained .onnx (no sidecar .data file) -
|
| 38 |
+
# much cleaner to commit to the HuggingFace Space via Git LFS.
|
| 39 |
+
dummy = torch.zeros(1, 1, TARGET_H, TARGET_W)
|
| 40 |
+
torch.onnx.export(
|
| 41 |
+
model,
|
| 42 |
+
dummy,
|
| 43 |
+
str(onnx_path),
|
| 44 |
+
input_names=["input"],
|
| 45 |
+
output_names=["logits"],
|
| 46 |
+
dynamic_axes={"input": {0: "batch"}, "logits": {0: "batch"}},
|
| 47 |
+
opset_version=OPSET,
|
| 48 |
+
dynamo=False,
|
| 49 |
+
)
|
| 50 |
+
# Belt and braces: if any external data slipped out, fold it back in so the
|
| 51 |
+
# .onnx is guaranteed standalone, then drop the stray .data file.
|
| 52 |
+
model_proto = onnx.load(str(onnx_path))
|
| 53 |
+
onnx.save_model(model_proto, str(onnx_path), save_as_external_data=False)
|
| 54 |
+
data_file = onnx_path.with_suffix(".onnx.data")
|
| 55 |
+
if data_file.exists():
|
| 56 |
+
data_file.unlink()
|
| 57 |
+
|
| 58 |
+
classes_path.write_text(json.dumps(classes))
|
| 59 |
+
size_mb = onnx_path.stat().st_size / 1e6
|
| 60 |
+
print(f"{ckpt_path} -> {onnx_path} ({len(classes)} classes, {size_mb:.1f} MB)")
|
| 61 |
+
return onnx_path
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def main() -> None:
|
| 65 |
+
parser = argparse.ArgumentParser(description=__doc__)
|
| 66 |
+
parser.add_argument(
|
| 67 |
+
"checkpoints",
|
| 68 |
+
nargs="*",
|
| 69 |
+
default=["checkpoints/best.pt", "checkpoints/kazan/best.pt"],
|
| 70 |
+
help="checkpoint(s) to export (default: moves + kazan best.pt)",
|
| 71 |
+
)
|
| 72 |
+
args = parser.parse_args()
|
| 73 |
+
for path in args.checkpoints:
|
| 74 |
+
p = Path(path)
|
| 75 |
+
if p.exists():
|
| 76 |
+
export_one(p)
|
| 77 |
+
else:
|
| 78 |
+
print(f"skip: {p} not found")
|
| 79 |
+
|
| 80 |
+
|
| 81 |
+
if __name__ == "__main__":
|
| 82 |
+
main()
|
scripts/get_glyphs.py
ADDED
|
@@ -0,0 +1,182 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Build the glyph pools used by the synthesizer (idempotent).
|
| 2 |
+
|
| 3 |
+
python scripts/get_glyphs.py # download + prepare ARDIS
|
| 4 |
+
python scripts/get_glyphs.py --emnist # also rebuild the EMNIST pool
|
| 5 |
+
# (only needed if ingredients/ is absent,
|
| 6 |
+
# e.g. on Colab)
|
| 7 |
+
|
| 8 |
+
ARDIS (https://ardisdataset.github.io/ARDIS/) Dataset II contains ~10k digit
|
| 9 |
+
images cropped from 19th-century European church records - crossed 7s, serif
|
| 10 |
+
1s, cursive styles that match Kazakh/European handwriting far better than
|
| 11 |
+
EMNIST. Images are normalized here into white-on-black float masks under
|
| 12 |
+
glyph_data/ardis/<digit>/.
|
| 13 |
+
"""
|
| 14 |
+
|
| 15 |
+
import argparse
|
| 16 |
+
import shutil
|
| 17 |
+
import subprocess
|
| 18 |
+
import sys
|
| 19 |
+
from pathlib import Path
|
| 20 |
+
|
| 21 |
+
import numpy as np
|
| 22 |
+
from PIL import Image
|
| 23 |
+
|
| 24 |
+
PROJECT_ROOT = Path(__file__).resolve().parent.parent
|
| 25 |
+
GLYPH_DATA = PROJECT_ROOT / "glyph_data"
|
| 26 |
+
ARDIS_URL = (
|
| 27 |
+
"https://github.com/ardisdataset/ARDIS/raw/Updates-Date-String/ARDIS_DATASET_II.rar"
|
| 28 |
+
)
|
| 29 |
+
IMAGE_EXTS = {".png", ".jpg", ".jpeg", ".bmp"}
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def download(url: str, dest: Path) -> None:
|
| 33 |
+
if dest.exists() and dest.stat().st_size > 0:
|
| 34 |
+
print(f"Already downloaded: {dest}")
|
| 35 |
+
return
|
| 36 |
+
dest.parent.mkdir(parents=True, exist_ok=True)
|
| 37 |
+
print(f"Downloading {url} ...")
|
| 38 |
+
subprocess.run(["curl", "-L", "--fail", "-o", str(dest), url], check=True)
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
def extract_rar(archive: Path, dest: Path) -> None:
|
| 42 |
+
if dest.exists() and any(dest.iterdir()):
|
| 43 |
+
print(f"Already extracted: {dest}")
|
| 44 |
+
return
|
| 45 |
+
dest.mkdir(parents=True, exist_ok=True)
|
| 46 |
+
if shutil.which("bsdtar"):
|
| 47 |
+
subprocess.run(["bsdtar", "-xf", str(archive), "-C", str(dest)], check=True)
|
| 48 |
+
elif shutil.which("unar"):
|
| 49 |
+
subprocess.run(["unar", "-quiet", "-o", str(dest), str(archive)], check=True)
|
| 50 |
+
elif shutil.which("unrar"):
|
| 51 |
+
subprocess.run(["unrar", "x", "-inul", str(archive), str(dest)], check=True)
|
| 52 |
+
else:
|
| 53 |
+
sys.exit(
|
| 54 |
+
"No RAR extractor found (need bsdtar, unar, or unrar). "
|
| 55 |
+
"On Ubuntu/Colab: apt-get install -y unar"
|
| 56 |
+
)
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
def infer_digit_label(path: Path) -> str | None:
|
| 60 |
+
if path.parent.name.isdigit() and len(path.parent.name) == 1:
|
| 61 |
+
return path.parent.name
|
| 62 |
+
stem = path.stem
|
| 63 |
+
for sep in ("_", "-", " "):
|
| 64 |
+
token = stem.split(sep)[0]
|
| 65 |
+
if token.isdigit() and len(token) == 1:
|
| 66 |
+
return token
|
| 67 |
+
return None
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
def photo_to_mask(path: Path) -> np.ndarray | None:
|
| 71 |
+
"""Ink-on-paper photo -> white-on-black float mask, tight-cropped."""
|
| 72 |
+
with Image.open(path) as img:
|
| 73 |
+
gray = np.asarray(img.convert("L"), dtype=np.float32)
|
| 74 |
+
if gray.shape[0] < 10 or gray.shape[1] < 5:
|
| 75 |
+
return None
|
| 76 |
+
inverted = 255.0 - gray
|
| 77 |
+
background = np.percentile(inverted, 50)
|
| 78 |
+
signal = np.clip(inverted - background, 0, None)
|
| 79 |
+
peak = np.percentile(signal, 99.5)
|
| 80 |
+
if peak < 12: # effectively blank
|
| 81 |
+
return None
|
| 82 |
+
mask = np.clip(signal / peak, 0.0, 1.0)
|
| 83 |
+
mask[mask < 0.18] = 0.0 # suppress paper texture
|
| 84 |
+
ys, xs = np.where(mask > 0.18)
|
| 85 |
+
if len(ys) < 20:
|
| 86 |
+
return None
|
| 87 |
+
mask = mask[ys.min() : ys.max() + 1, xs.min() : xs.max() + 1]
|
| 88 |
+
if mask.shape[0] < 8 or mask.shape[1] < 3:
|
| 89 |
+
return None
|
| 90 |
+
return mask
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
def prepare_ardis() -> None:
|
| 94 |
+
archive = GLYPH_DATA / "_downloads" / "ARDIS_DATASET_II.rar"
|
| 95 |
+
raw_dir = GLYPH_DATA / "_raw" / "ardis2"
|
| 96 |
+
out_root = GLYPH_DATA / "ardis"
|
| 97 |
+
if out_root.exists() and any(out_root.glob("*/*.png")):
|
| 98 |
+
print(f"ARDIS pool already prepared at {out_root}")
|
| 99 |
+
return
|
| 100 |
+
|
| 101 |
+
download(ARDIS_URL, archive)
|
| 102 |
+
extract_rar(archive, raw_dir)
|
| 103 |
+
|
| 104 |
+
counts: dict[str, int] = {}
|
| 105 |
+
skipped = 0
|
| 106 |
+
for path in sorted(raw_dir.rglob("*")):
|
| 107 |
+
if path.suffix.lower() not in IMAGE_EXTS:
|
| 108 |
+
continue
|
| 109 |
+
label = infer_digit_label(path)
|
| 110 |
+
if label is None: # 0 is kept: kazan checkpoint numbers use it
|
| 111 |
+
skipped += 1
|
| 112 |
+
continue
|
| 113 |
+
mask = photo_to_mask(path)
|
| 114 |
+
if mask is None:
|
| 115 |
+
skipped += 1
|
| 116 |
+
continue
|
| 117 |
+
out_dir = out_root / label
|
| 118 |
+
out_dir.mkdir(parents=True, exist_ok=True)
|
| 119 |
+
index = counts.get(label, 0)
|
| 120 |
+
Image.fromarray((mask * 255).astype(np.uint8)).save(out_dir / f"{label}_{index}.png")
|
| 121 |
+
counts[label] = index + 1
|
| 122 |
+
|
| 123 |
+
print(f"ARDIS pool: {sum(counts.values())} glyphs "
|
| 124 |
+
f"({', '.join(f'{k}:{v}' for k, v in sorted(counts.items()))}), "
|
| 125 |
+
f"skipped {skipped}")
|
| 126 |
+
if not counts:
|
| 127 |
+
print("WARNING: no ARDIS glyphs extracted - the synthesizer will fall back "
|
| 128 |
+
"to EMNIST + procedural styles.")
|
| 129 |
+
|
| 130 |
+
|
| 131 |
+
def prepare_emnist(per_class: int) -> None:
|
| 132 |
+
from torchvision.datasets import EMNIST
|
| 133 |
+
|
| 134 |
+
out_root = GLYPH_DATA / "emnist"
|
| 135 |
+
if out_root.exists() and any(out_root.glob("*/*.png")):
|
| 136 |
+
print(f"EMNIST pool already prepared at {out_root}")
|
| 137 |
+
return
|
| 138 |
+
|
| 139 |
+
print("Downloading EMNIST (byclass) via torchvision ...")
|
| 140 |
+
dataset = EMNIST(str(GLYPH_DATA / "_downloads"), split="byclass", train=True, download=True)
|
| 141 |
+
# byclass labels: 0-9 digits, 10-35 A-Z, 36-61 a-z
|
| 142 |
+
wanted = {label: str(label) for label in range(0, 10)}
|
| 143 |
+
wanted[33] = "x" # 'X'
|
| 144 |
+
wanted[59] = "x" # 'x'
|
| 145 |
+
|
| 146 |
+
counts: dict[str, int] = {}
|
| 147 |
+
for img, label in zip(dataset.data, dataset.targets):
|
| 148 |
+
char = wanted.get(int(label))
|
| 149 |
+
if char is None:
|
| 150 |
+
continue
|
| 151 |
+
if counts.get(char, 0) >= per_class:
|
| 152 |
+
if all(counts.get(c, 0) >= per_class for c in wanted.values()):
|
| 153 |
+
break
|
| 154 |
+
continue
|
| 155 |
+
arr = img.numpy().T # EMNIST images are stored transposed
|
| 156 |
+
out_dir = out_root / char
|
| 157 |
+
out_dir.mkdir(parents=True, exist_ok=True)
|
| 158 |
+
index = counts.get(char, 0)
|
| 159 |
+
Image.fromarray(arr).save(out_dir / f"{char}_{index}.png")
|
| 160 |
+
counts[char] = index + 1
|
| 161 |
+
print(f"EMNIST pool: {sum(counts.values())} glyphs "
|
| 162 |
+
f"({', '.join(f'{k}:{v}' for k, v in sorted(counts.items()))})")
|
| 163 |
+
|
| 164 |
+
|
| 165 |
+
if __name__ == "__main__":
|
| 166 |
+
parser = argparse.ArgumentParser(description=__doc__)
|
| 167 |
+
parser.add_argument("--emnist", action="store_true",
|
| 168 |
+
help="also build EMNIST pool (needed when ingredients/ is absent)")
|
| 169 |
+
parser.add_argument("--per-class", type=int, default=2500,
|
| 170 |
+
help="max EMNIST glyphs per character")
|
| 171 |
+
args = parser.parse_args()
|
| 172 |
+
|
| 173 |
+
prepare_ardis()
|
| 174 |
+
if args.emnist or not (PROJECT_ROOT / "ingredients").is_dir():
|
| 175 |
+
prepare_emnist(args.per_class)
|
| 176 |
+
|
| 177 |
+
# final summary through the sampler itself
|
| 178 |
+
sys.path.insert(0, str(PROJECT_ROOT))
|
| 179 |
+
from togyz.glyphs import GlyphSampler
|
| 180 |
+
|
| 181 |
+
print("\nGlyph pools as seen by the synthesizer:")
|
| 182 |
+
print(GlyphSampler(PROJECT_ROOT).describe())
|
tests/fixtures/shortest_candidate_replay.tsv
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
shortest_candidate_halfmoves 11
|
| 2 |
+
initial_position p1_side=81 p2_side=81 p1_kazan=0 p2_kazan=0
|
| 3 |
+
halfmove player move p1_kazan p2_kazan p1_side p2_side p1_tuzdyk p2_tuzdyk winner
|
| 4 |
+
1 P1 98 10 0 73 79 -1 -1 ongoing
|
| 5 |
+
2 P2 22 10 10 65 77 -1 -1 ongoing
|
| 6 |
+
3 P1 87 22 10 58 72 -1 -1 ongoing
|
| 7 |
+
4 P2 12 22 10 60 70 -1 -1 ongoing
|
| 8 |
+
5 P1 76 36 10 54 62 -1 -1 ongoing
|
| 9 |
+
6 P2 25 36 10 54 62 -1 -1 ongoing
|
| 10 |
+
7 P1 65 52 10 49 51 -1 -1 ongoing
|
| 11 |
+
8 P2 13 52 10 49 51 -1 -1 ongoing
|
| 12 |
+
9 P1 93 70 10 46 36 -1 -1 ongoing
|
| 13 |
+
10 P2 25 70 10 46 36 -1 -1 ongoing
|
| 14 |
+
11 P1 54 88 10 42 22 -1 -1 P1
|
tests/fixtures/shortest_terminal_game.txt
ADDED
|
@@ -0,0 +1,6 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
1. 98 22
|
| 2 |
+
2. 87 12
|
| 3 |
+
3. 76 25
|
| 4 |
+
4. 65 13
|
| 5 |
+
5. 93 25
|
| 6 |
+
6. 54
|
tests/test_rules.py
ADDED
|
@@ -0,0 +1,110 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Validate togyz/rules.py against the 9Q engine's shortest-game fixture.
|
| 2 |
+
|
| 3 |
+
Replays tests/fixtures/shortest_candidate_replay.tsv move by move and asserts
|
| 4 |
+
that notation, kazans, side totals, tuzdyks, and the winner all match the
|
| 5 |
+
C++ engine's recorded state after every halfmove.
|
| 6 |
+
|
| 7 |
+
python tests/test_rules.py
|
| 8 |
+
"""
|
| 9 |
+
|
| 10 |
+
import csv
|
| 11 |
+
import sys
|
| 12 |
+
from pathlib import Path
|
| 13 |
+
|
| 14 |
+
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
| 15 |
+
|
| 16 |
+
from togyz.rules import Game
|
| 17 |
+
|
| 18 |
+
FIXTURE = Path(__file__).parent / "fixtures" / "shortest_candidate_replay.tsv"
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def main() -> None:
|
| 22 |
+
with open(FIXTURE, newline="") as f:
|
| 23 |
+
lines = f.read().splitlines()
|
| 24 |
+
rows = list(csv.DictReader(lines[2:], delimiter="\t")) # skip 2 meta lines
|
| 25 |
+
|
| 26 |
+
game = Game()
|
| 27 |
+
for row in rows:
|
| 28 |
+
expected = row["move"]
|
| 29 |
+
player = 0 if row["player"] == "P1" else 1
|
| 30 |
+
assert game.to_play == player, f"halfmove {row['halfmove']}: wrong side to move"
|
| 31 |
+
|
| 32 |
+
legal = {m.notation for m in game.legal_moves()}
|
| 33 |
+
assert expected in legal, f"halfmove {row['halfmove']}: {expected} not in {legal}"
|
| 34 |
+
|
| 35 |
+
played = game.play_notation(expected)
|
| 36 |
+
assert played.notation == expected
|
| 37 |
+
|
| 38 |
+
state = (game.kazans[0], game.kazans[1], *game.side_totals(),
|
| 39 |
+
game.tuzdyks[0], game.tuzdyks[1])
|
| 40 |
+
want = tuple(int(row[k]) for k in
|
| 41 |
+
("p1_kazan", "p2_kazan", "p1_side", "p2_side", "p1_tuzdyk", "p2_tuzdyk"))
|
| 42 |
+
assert state == want, (
|
| 43 |
+
f"halfmove {row['halfmove']} ({expected}): state {state} != fixture {want}"
|
| 44 |
+
)
|
| 45 |
+
|
| 46 |
+
if row["winner"] != "ongoing":
|
| 47 |
+
assert game.is_over, f"halfmove {row['halfmove']}: game should be over"
|
| 48 |
+
expected_winner = {"P1": 0, "P2": 1, "draw": -1}[row["winner"]]
|
| 49 |
+
assert game.winner() == expected_winner
|
| 50 |
+
|
| 51 |
+
print(f"OK: replayed {len(rows)} halfmoves, all states match the 9Q fixture")
|
| 52 |
+
|
| 53 |
+
# extra sanity: fresh game has exactly 9 legal moves, all distinct notations
|
| 54 |
+
fresh = Game()
|
| 55 |
+
moves = fresh.legal_moves()
|
| 56 |
+
assert len(moves) == 9 and len({m.notation for m in moves}) == 9
|
| 57 |
+
print("OK: initial position has 9 distinct legal moves")
|
| 58 |
+
|
| 59 |
+
# tuzdyk scenario (not covered by the fixture): P1 sows 3 stones from
|
| 60 |
+
# hole 9 -> last stone lands in P2's hole 2 (pit 10) making it 3 -> "92x"
|
| 61 |
+
g = Game()
|
| 62 |
+
g.pits = [9, 9, 9, 9, 9, 9, 9, 9, 3] + [2, 2, 9, 9, 9, 9, 9, 9, 9]
|
| 63 |
+
move = g.play_notation("92x")
|
| 64 |
+
assert move.makes_tuzdyk and g.tuzdyks[0] == 1 and g.kazans[0] == 3
|
| 65 |
+
assert g.pits[10] == 0
|
| 66 |
+
# P2 may not move from their own hole 2 (it is P1's tuzdyk now)
|
| 67 |
+
assert 1 not in g.legal_actions()
|
| 68 |
+
# P2 sows 3 stones from their hole 1 (pit 9): one falls into the tuzdyk
|
| 69 |
+
# and must be collected by its owner P1
|
| 70 |
+
g.play_notation("13")
|
| 71 |
+
assert g.kazans[0] == 3 + 1, "stone sown into tuzdyk must go to owner"
|
| 72 |
+
assert g.pits[10] == 0
|
| 73 |
+
print("OK: tuzdyk creation, blocking, and collection behave correctly")
|
| 74 |
+
|
| 75 |
+
# BOTH players may create one tuzdyk each (as on the real scoresheet:
|
| 76 |
+
# 47x by White then 91x by Black) - but never a second one, never in
|
| 77 |
+
# hole 9, and never mirroring the opponent's tuzdyk column.
|
| 78 |
+
g = Game()
|
| 79 |
+
g.tuzdyks = [4, -1] # White already owns a tuzdyk (Black's hole 5)
|
| 80 |
+
g.to_play = 1
|
| 81 |
+
g.pits = [9, 9, 2, 9, 9, 9, 9, 9, 9] + [9] * 9
|
| 82 |
+
# Black hole 4 (pit 12), 9 stones: lands (12+9-1)%18 = pit 2, making 3
|
| 83 |
+
move = g.play_notation("43x")
|
| 84 |
+
assert move.makes_tuzdyk and g.tuzdyks == [4, 2]
|
| 85 |
+
# a second tuzdyk for the same player must be impossible anywhere
|
| 86 |
+
g.to_play = 1
|
| 87 |
+
assert not any(m.makes_tuzdyk for m in g.legal_moves())
|
| 88 |
+
print("OK: both players can own one tuzdyk; no second tuzdyk possible")
|
| 89 |
+
|
| 90 |
+
# mirror restriction: creation blocked where the opponent's tuzdyk sits
|
| 91 |
+
g = Game()
|
| 92 |
+
g.tuzdyks = [-1, 2] # Black owns White's hole 3 (column index 2)
|
| 93 |
+
g.pits = [9] * 18
|
| 94 |
+
g.pits[8] = 4 # White hole 9: lands (8+4-1)%18 = pit 11 (Black col 2)
|
| 95 |
+
g.pits[11] = 2 # becomes 3 on landing - but mirrored, so no tuzdyk
|
| 96 |
+
move = g.play_notation("93")
|
| 97 |
+
assert not move.makes_tuzdyk and g.tuzdyks[0] == -1
|
| 98 |
+
|
| 99 |
+
# hole 9 restriction: landing makes 3 in opponent hole 9 -> no tuzdyk
|
| 100 |
+
g = Game()
|
| 101 |
+
g.pits = [9] * 18
|
| 102 |
+
g.pits[8] = 10 # White hole 9: lands (8+10-1)%18 = pit 17 (Black hole 9)
|
| 103 |
+
g.pits[17] = 2 # becomes 3 on landing
|
| 104 |
+
move = g.play_notation("99")
|
| 105 |
+
assert not move.makes_tuzdyk and g.tuzdyks[0] == -1, "hole 9 must never be tuzdyk"
|
| 106 |
+
print("OK: mirror and hole-9 tuzdyk restrictions enforced")
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
if __name__ == "__main__":
|
| 110 |
+
main()
|
togyz/__init__.py
ADDED
|
File without changes
|
togyz/classes.py
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Canonical class list for Togyzkumalak move OCR.
|
| 2 |
+
|
| 3 |
+
162 moves (start hole 1-9, end hole 1-9, with or without a capture 'x')
|
| 4 |
+
plus 'empty'. The order matches the alphabetical order used by
|
| 5 |
+
torchvision ImageFolder in the legacy pipeline (legacy/class_mapping.json),
|
| 6 |
+
so old and new indices agree.
|
| 7 |
+
"""
|
| 8 |
+
|
| 9 |
+
MOVES = [
|
| 10 |
+
f"{start}{end}{suffix}"
|
| 11 |
+
for start in range(1, 10)
|
| 12 |
+
for end in range(1, 10)
|
| 13 |
+
for suffix in ("", "x")
|
| 14 |
+
]
|
| 15 |
+
|
| 16 |
+
EMPTY = "empty"
|
| 17 |
+
CLASSES = MOVES + [EMPTY]
|
| 18 |
+
CLASS_TO_IDX = {name: i for i, name in enumerate(CLASSES)}
|
| 19 |
+
NUM_CLASSES = len(CLASSES) # 163
|
| 20 |
+
|
| 21 |
+
# Kazan checkpoint numbers on the scoresheet summary strips: always exactly
|
| 22 |
+
# two digits (with leading zero), values 00-81, never an 'x'.
|
| 23 |
+
KAZAN_CLASSES = [f"{v:02d}" for v in range(82)] + [EMPTY]
|
togyz/dataset.py
ADDED
|
@@ -0,0 +1,71 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Datasets: on-the-fly synthetic cells and the labeled real crops."""
|
| 2 |
+
|
| 3 |
+
import csv
|
| 4 |
+
import random
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
from torch.utils.data import Dataset
|
| 9 |
+
|
| 10 |
+
from .classes import CLASSES, CLASS_TO_IDX
|
| 11 |
+
from .glyphs import GlyphSampler
|
| 12 |
+
from .preprocess import preprocess_file, preprocess_pil
|
| 13 |
+
from .synth import synthesize_cell
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
class SyntheticCellDataset(Dataset):
|
| 17 |
+
"""Generates cells on the fly - no dataset folders on disk.
|
| 18 |
+
|
| 19 |
+
With seed=None every access is fresh random data (infinite training
|
| 20 |
+
variety; DataLoader workers are seeded independently by torch). With a
|
| 21 |
+
seed, sample i is always the same image: a fixed validation set.
|
| 22 |
+
"""
|
| 23 |
+
|
| 24 |
+
def __init__(self, num_samples: int, seed: int | None = None, project_root=".",
|
| 25 |
+
classes: list[str] | None = None):
|
| 26 |
+
self.sampler = GlyphSampler(project_root)
|
| 27 |
+
self.num_samples = num_samples
|
| 28 |
+
self.seed = seed
|
| 29 |
+
self.classes = classes or CLASSES
|
| 30 |
+
self.class_to_idx = {c: i for i, c in enumerate(self.classes)}
|
| 31 |
+
|
| 32 |
+
def __len__(self) -> int:
|
| 33 |
+
return self.num_samples
|
| 34 |
+
|
| 35 |
+
def __getitem__(self, index: int):
|
| 36 |
+
if self.seed is not None:
|
| 37 |
+
rng = random.Random(f"{self.seed}:{index}")
|
| 38 |
+
else:
|
| 39 |
+
rng = random.Random(random.getrandbits(64))
|
| 40 |
+
class_name = self.classes[rng.randrange(len(self.classes))]
|
| 41 |
+
img = synthesize_cell(class_name, self.sampler, rng)
|
| 42 |
+
return torch.from_numpy(preprocess_pil(img)), self.class_to_idx[class_name]
|
| 43 |
+
|
| 44 |
+
|
| 45 |
+
class RealCropDataset(Dataset):
|
| 46 |
+
"""The labeled real photos in data/real_crops (tiny; used for eval only)."""
|
| 47 |
+
|
| 48 |
+
def __init__(self, root: str = "data/real_crops"):
|
| 49 |
+
self.root = Path(root)
|
| 50 |
+
self.items: list[tuple[str, str]] = []
|
| 51 |
+
labels_file = self.root / "labels.csv"
|
| 52 |
+
if labels_file.exists():
|
| 53 |
+
with open(labels_file, newline="") as f:
|
| 54 |
+
for row in csv.DictReader(f):
|
| 55 |
+
label = row["label"].strip()
|
| 56 |
+
if label in CLASS_TO_IDX:
|
| 57 |
+
self.items.append((row["filename"].strip(), label))
|
| 58 |
+
|
| 59 |
+
def __len__(self) -> int:
|
| 60 |
+
return len(self.items)
|
| 61 |
+
|
| 62 |
+
def __getitem__(self, index: int):
|
| 63 |
+
filename, label = self.items[index]
|
| 64 |
+
tensor = torch.from_numpy(preprocess_file(str(self.root / filename)))
|
| 65 |
+
return tensor, CLASS_TO_IDX[label]
|
| 66 |
+
|
| 67 |
+
def batch(self) -> tuple[torch.Tensor, torch.Tensor, list[str]]:
|
| 68 |
+
"""All crops as one batch plus their filenames (for reporting)."""
|
| 69 |
+
tensors, labels = zip(*(self[i] for i in range(len(self))))
|
| 70 |
+
names = [name for name, _ in self.items]
|
| 71 |
+
return torch.stack(tensors), torch.tensor(labels), names
|
togyz/glyphs.py
ADDED
|
@@ -0,0 +1,167 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Glyph pools: per-character handwriting samples used by the synthesizer.
|
| 2 |
+
|
| 3 |
+
Sources (all optional, at least one digit source must exist):
|
| 4 |
+
* ingredients/<char>/*.png EMNIST-style masks, white ink on black
|
| 5 |
+
* glyph_data/emnist/<char>/*.png same, regenerated via scripts/get_glyphs.py --emnist
|
| 6 |
+
* glyph_data/ardis/<digit>/*.png European handwriting (ARDIS), pre-normalized
|
| 7 |
+
to white-on-black masks by scripts/get_glyphs.py
|
| 8 |
+
|
| 9 |
+
EMNIST glyphs are American-style; real Togyzkumalak scoresheets use
|
| 10 |
+
European/Kazakh styles (crossed 7, serif 1, cursive 9). ARDIS provides real
|
| 11 |
+
European glyphs; on top of that, procedural edits (crossbar on 7, flag on 1)
|
| 12 |
+
convert a share of EMNIST glyphs to European style, and 'x' marks are also
|
| 13 |
+
drawn procedurally for extra variety.
|
| 14 |
+
"""
|
| 15 |
+
|
| 16 |
+
import math
|
| 17 |
+
import os
|
| 18 |
+
import random
|
| 19 |
+
from pathlib import Path
|
| 20 |
+
|
| 21 |
+
import cv2
|
| 22 |
+
import numpy as np
|
| 23 |
+
from PIL import Image
|
| 24 |
+
|
| 25 |
+
DIGITS = "123456789" # move notation digits; '0' appears only in kazan numbers
|
| 26 |
+
CHARS = "0" + DIGITS + "x"
|
| 27 |
+
INK_THRESHOLD = 0.15 # mask values above this count as ink
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def _list_images(directory: Path) -> list[str]:
|
| 31 |
+
if not directory.is_dir():
|
| 32 |
+
return []
|
| 33 |
+
exts = {".png", ".jpg", ".jpeg", ".bmp"}
|
| 34 |
+
return sorted(
|
| 35 |
+
str(p) for p in directory.iterdir() if p.suffix.lower() in exts
|
| 36 |
+
)
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def _tight_crop(mask: np.ndarray) -> np.ndarray:
|
| 40 |
+
ys, xs = np.where(mask > INK_THRESHOLD)
|
| 41 |
+
if len(ys) == 0:
|
| 42 |
+
return mask
|
| 43 |
+
return mask[ys.min() : ys.max() + 1, xs.min() : xs.max() + 1]
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
def _load_mask(path: str) -> np.ndarray:
|
| 47 |
+
"""Load a white-on-black glyph image as float mask in [0, 1]."""
|
| 48 |
+
with Image.open(path) as img:
|
| 49 |
+
arr = np.asarray(img.convert("L"), dtype=np.float32) / 255.0
|
| 50 |
+
return _tight_crop(arr)
|
| 51 |
+
|
| 52 |
+
|
| 53 |
+
def add_crossbar_to_7(mask: np.ndarray, rng: random.Random) -> np.ndarray:
|
| 54 |
+
"""European 7: horizontal bar through the stem."""
|
| 55 |
+
h, w = mask.shape
|
| 56 |
+
if h < 8 or w < 4:
|
| 57 |
+
return mask
|
| 58 |
+
y = int(h * rng.uniform(0.40, 0.60))
|
| 59 |
+
x0 = int(w * rng.uniform(-0.05, 0.15))
|
| 60 |
+
x1 = int(w * rng.uniform(0.85, 1.05))
|
| 61 |
+
dy = int(w * rng.uniform(-0.10, 0.10))
|
| 62 |
+
thickness = max(2, int(round(h * 0.05 * rng.uniform(0.8, 1.6))))
|
| 63 |
+
out = mask.copy()
|
| 64 |
+
cv2.line(out, (x0, y - dy // 2), (x1, y + dy // 2), 1.0, thickness, cv2.LINE_AA)
|
| 65 |
+
return np.clip(out, 0.0, 1.0)
|
| 66 |
+
|
| 67 |
+
|
| 68 |
+
def add_flag_to_1(mask: np.ndarray, rng: random.Random) -> np.ndarray:
|
| 69 |
+
"""European 1: diagonal flag from the top of the stem down-left."""
|
| 70 |
+
h, w = mask.shape
|
| 71 |
+
if h < 8:
|
| 72 |
+
return mask
|
| 73 |
+
ink_rows = np.where(mask.max(axis=1) > INK_THRESHOLD)[0]
|
| 74 |
+
if len(ink_rows) == 0:
|
| 75 |
+
return mask
|
| 76 |
+
y_top = int(ink_rows[0])
|
| 77 |
+
row = mask[min(y_top + 1, h - 1)]
|
| 78 |
+
x_top = int(np.argmax(row))
|
| 79 |
+
length = h * rng.uniform(0.22, 0.42)
|
| 80 |
+
angle = math.radians(rng.uniform(25, 55))
|
| 81 |
+
x_end = int(x_top - length * math.cos(angle))
|
| 82 |
+
y_end = int(y_top + length * math.sin(angle))
|
| 83 |
+
thickness = max(2, int(round(h * 0.05 * rng.uniform(0.8, 1.4))))
|
| 84 |
+
# give the flag room on the left if it would leave the canvas
|
| 85 |
+
pad = max(0, -x_end + 2)
|
| 86 |
+
out = np.pad(mask, ((0, 0), (pad, 0)))
|
| 87 |
+
cv2.line(out, (x_top + pad, y_top), (x_end + pad, y_end), 1.0, thickness, cv2.LINE_AA)
|
| 88 |
+
return np.clip(out, 0.0, 1.0)
|
| 89 |
+
|
| 90 |
+
|
| 91 |
+
def _bezier_points(p0, p1, p2, n=24):
|
| 92 |
+
t = np.linspace(0.0, 1.0, n)[:, None]
|
| 93 |
+
pts = (1 - t) ** 2 * p0 + 2 * (1 - t) * t * p1 + t**2 * p2
|
| 94 |
+
return pts.astype(np.int32)
|
| 95 |
+
|
| 96 |
+
|
| 97 |
+
def procedural_x(rng: random.Random) -> np.ndarray:
|
| 98 |
+
"""Two crossing, slightly curved pen strokes."""
|
| 99 |
+
s = 64
|
| 100 |
+
canvas = np.zeros((s, s), np.float32)
|
| 101 |
+
for (ax, ay), (bx, by) in [((0.15, 0.12), (0.85, 0.88)), ((0.85, 0.10), (0.15, 0.90))]:
|
| 102 |
+
p0 = np.array([ax + rng.uniform(-0.08, 0.08), ay + rng.uniform(-0.08, 0.08)]) * s
|
| 103 |
+
p2 = np.array([bx + rng.uniform(-0.08, 0.08), by + rng.uniform(-0.08, 0.08)]) * s
|
| 104 |
+
mid = (p0 + p2) / 2 + np.array([rng.uniform(-0.15, 0.15), rng.uniform(-0.15, 0.15)]) * s
|
| 105 |
+
pts = _bezier_points(p0, mid, p2)
|
| 106 |
+
thickness = rng.randint(3, 6)
|
| 107 |
+
cv2.polylines(canvas, [pts], False, 1.0, thickness, cv2.LINE_AA)
|
| 108 |
+
canvas = cv2.GaussianBlur(canvas, (3, 3), 0.6)
|
| 109 |
+
return _tight_crop(np.clip(canvas, 0.0, 1.0))
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
class GlyphSampler:
|
| 113 |
+
"""Samples a glyph mask (float [0,1], white ink, tight-cropped) per character."""
|
| 114 |
+
|
| 115 |
+
def __init__(self, project_root: str | os.PathLike = "."):
|
| 116 |
+
root = Path(project_root)
|
| 117 |
+
self.pools: dict[str, dict[str, list[str]]] = {c: {} for c in CHARS}
|
| 118 |
+
for char in CHARS:
|
| 119 |
+
emnist = _list_images(root / "ingredients" / char) + _list_images(
|
| 120 |
+
root / "glyph_data" / "emnist" / char
|
| 121 |
+
)
|
| 122 |
+
if emnist:
|
| 123 |
+
self.pools[char]["emnist"] = emnist
|
| 124 |
+
if char != "x":
|
| 125 |
+
ardis = _list_images(root / "glyph_data" / "ardis" / char)
|
| 126 |
+
if ardis:
|
| 127 |
+
self.pools[char]["ardis"] = ardis
|
| 128 |
+
|
| 129 |
+
missing = [c for c in DIGITS if not self.pools[c]]
|
| 130 |
+
if missing:
|
| 131 |
+
raise FileNotFoundError(
|
| 132 |
+
f"No glyph images found for digits {missing}. "
|
| 133 |
+
"Run `python scripts/get_glyphs.py` (optionally with --emnist) first."
|
| 134 |
+
)
|
| 135 |
+
|
| 136 |
+
def describe(self) -> str:
|
| 137 |
+
lines = []
|
| 138 |
+
for char in CHARS:
|
| 139 |
+
sources = ", ".join(
|
| 140 |
+
f"{name}:{len(paths)}" for name, paths in self.pools[char].items()
|
| 141 |
+
) or "procedural only"
|
| 142 |
+
lines.append(f" '{char}': {sources}")
|
| 143 |
+
return "\n".join(lines)
|
| 144 |
+
|
| 145 |
+
def sample(self, char: str, rng: random.Random) -> np.ndarray:
|
| 146 |
+
if char == "x" and (not self.pools["x"] or rng.random() < 0.35):
|
| 147 |
+
return procedural_x(rng)
|
| 148 |
+
|
| 149 |
+
sources = self.pools[char]
|
| 150 |
+
if not sources:
|
| 151 |
+
raise FileNotFoundError(
|
| 152 |
+
f"No glyphs for {char!r} - re-run scripts/get_glyphs.py "
|
| 153 |
+
"(digit 0 needs the ARDIS pool or --emnist)"
|
| 154 |
+
)
|
| 155 |
+
if "ardis" in sources and "emnist" in sources:
|
| 156 |
+
source = "ardis" if rng.random() < 0.55 else "emnist"
|
| 157 |
+
else:
|
| 158 |
+
source = next(iter(sources))
|
| 159 |
+
mask = _load_mask(rng.choice(sources[source]))
|
| 160 |
+
|
| 161 |
+
# Convert a share of EMNIST glyphs to European style; ARDIS already is.
|
| 162 |
+
euro_prob = 0.15 if source == "ardis" else 0.45
|
| 163 |
+
if char == "7" and rng.random() < euro_prob:
|
| 164 |
+
mask = add_crossbar_to_7(mask, rng)
|
| 165 |
+
if char == "1" and source == "emnist" and rng.random() < 0.35:
|
| 166 |
+
mask = add_flag_to_1(mask, rng)
|
| 167 |
+
return mask
|
togyz/model.py
ADDED
|
@@ -0,0 +1,46 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Model factory and checkpoint helpers (plain torch.save, no fastai pickles)."""
|
| 2 |
+
|
| 3 |
+
import torch
|
| 4 |
+
import torch.nn as nn
|
| 5 |
+
from torchvision import models
|
| 6 |
+
|
| 7 |
+
from .classes import CLASSES
|
| 8 |
+
from .preprocess import TARGET_H, TARGET_W
|
| 9 |
+
|
| 10 |
+
|
| 11 |
+
def build_model(num_classes: int = len(CLASSES)) -> nn.Module:
|
| 12 |
+
model = models.resnet18(weights=None)
|
| 13 |
+
model.conv1 = nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False)
|
| 14 |
+
model.fc = nn.Linear(model.fc.in_features, num_classes)
|
| 15 |
+
return model
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def auto_device() -> torch.device:
|
| 19 |
+
if torch.cuda.is_available():
|
| 20 |
+
return torch.device("cuda")
|
| 21 |
+
if torch.backends.mps.is_available():
|
| 22 |
+
return torch.device("mps")
|
| 23 |
+
return torch.device("cpu")
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def save_checkpoint(path, model, epoch, val_acc, optimizer=None, classes=None):
|
| 27 |
+
torch.save(
|
| 28 |
+
{
|
| 29 |
+
"model_state": model.state_dict(),
|
| 30 |
+
"classes": classes or CLASSES,
|
| 31 |
+
"input_size": [TARGET_H, TARGET_W],
|
| 32 |
+
"epoch": epoch,
|
| 33 |
+
"val_acc": val_acc,
|
| 34 |
+
"optimizer_state": optimizer.state_dict() if optimizer else None,
|
| 35 |
+
},
|
| 36 |
+
path,
|
| 37 |
+
)
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def load_checkpoint(path, device=None) -> tuple[nn.Module, dict]:
|
| 41 |
+
device = device or auto_device()
|
| 42 |
+
ckpt = torch.load(path, map_location=device, weights_only=True)
|
| 43 |
+
model = build_model(len(ckpt["classes"]))
|
| 44 |
+
model.load_state_dict(ckpt["model_state"])
|
| 45 |
+
model.to(device).eval()
|
| 46 |
+
return model, ckpt
|
togyz/pipeline.py
ADDED
|
@@ -0,0 +1,355 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Torch-free scoresheet -> game reconstruction (onnxruntime + numpy).
|
| 2 |
+
|
| 3 |
+
This is the single inference core shared by the CLI (`read_game.py`) and the
|
| 4 |
+
Gradio demo (`app.py`). It deliberately imports **no torch**: the move and
|
| 5 |
+
kazan classifiers run as ONNX sessions, and every array op is numpy. The
|
| 6 |
+
heavy lifting of cell extraction (`togyz.sheet`) and the rules engine
|
| 7 |
+
(`togyz.rules`) are already torch-free and imported as-is.
|
| 8 |
+
|
| 9 |
+
The scoring/beam logic here is a straight port of the torch version whose
|
| 10 |
+
weights and variants are benchmarked against hand-transcribed truth in the
|
| 11 |
+
git history (sheet 1 = 36/39 exact plies). Keep it numerically equivalent.
|
| 12 |
+
"""
|
| 13 |
+
|
| 14 |
+
from __future__ import annotations
|
| 15 |
+
|
| 16 |
+
import json
|
| 17 |
+
import math
|
| 18 |
+
from dataclasses import dataclass, field
|
| 19 |
+
from pathlib import Path
|
| 20 |
+
|
| 21 |
+
import numpy as np
|
| 22 |
+
import onnxruntime as ort
|
| 23 |
+
from PIL import Image
|
| 24 |
+
|
| 25 |
+
from .classes import EMPTY
|
| 26 |
+
from .preprocess import preprocess_pil
|
| 27 |
+
from .rules import Game
|
| 28 |
+
from .sheet import clean_cell, extract_cells, extract_kazan_cells, render_overlay
|
| 29 |
+
|
| 30 |
+
MIN_CELL_HEIGHT = 35 # px; below this, resolution is low (surfaced as a warning)
|
| 31 |
+
MIN_INK_RATIO = 0.02 # cells with less handwriting ink than this are empty
|
| 32 |
+
MIN_KAZAN_INK = 0.028 # kazan boxes are small; box-edge remnants score higher
|
| 33 |
+
|
| 34 |
+
TTA_SHIFTS = [(0, 0), (-2, 0), (2, 0), (0, -2), (0, 2), (-2, -2), (2, 2)]
|
| 35 |
+
|
| 36 |
+
EXACT_WEIGHT = 0.5 # weight of the exact class-string probability
|
| 37 |
+
FACTOR_WEIGHT = 0.5 # weight of the factorized digit-marginal probability
|
| 38 |
+
KAZAN_FLOOR = 1e-4 # a misread checkpoint must not single-handedly kill truth
|
| 39 |
+
RESULT_CODES = {"1-0": 0, "0-1": 1, "draw": -1}
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
# --------------------------------------------------------------------------- #
|
| 43 |
+
# ONNX session loading
|
| 44 |
+
# --------------------------------------------------------------------------- #
|
| 45 |
+
@dataclass
|
| 46 |
+
class Classifier:
|
| 47 |
+
"""An ONNX classifier plus the class list its output indices map to."""
|
| 48 |
+
|
| 49 |
+
session: ort.InferenceSession
|
| 50 |
+
classes: list[str]
|
| 51 |
+
input_name: str = field(default="")
|
| 52 |
+
|
| 53 |
+
def __post_init__(self):
|
| 54 |
+
self.input_name = self.session.get_inputs()[0].name
|
| 55 |
+
|
| 56 |
+
def run(self, batch: np.ndarray) -> np.ndarray:
|
| 57 |
+
"""batch [N,1,H,W] float32 -> logits [N, num_classes] float32."""
|
| 58 |
+
return self.session.run(None, {self.input_name: batch})[0]
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def load_classifier(onnx_path: str | Path) -> Classifier:
|
| 62 |
+
"""Load a .onnx and its `<stem>.classes.json` sidecar (written by export)."""
|
| 63 |
+
onnx_path = Path(onnx_path)
|
| 64 |
+
classes_path = onnx_path.with_suffix(".classes.json")
|
| 65 |
+
classes = json.loads(classes_path.read_text())
|
| 66 |
+
session = ort.InferenceSession(str(onnx_path), providers=["CPUExecutionProvider"])
|
| 67 |
+
return Classifier(session, classes)
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
# --------------------------------------------------------------------------- #
|
| 71 |
+
# Cell classification (numpy + ONNX, with test-time augmentation)
|
| 72 |
+
# --------------------------------------------------------------------------- #
|
| 73 |
+
def _softmax(logits: np.ndarray, temperature: float = 1.0) -> np.ndarray:
|
| 74 |
+
z = logits / temperature
|
| 75 |
+
z = z - z.max(axis=1, keepdims=True)
|
| 76 |
+
e = np.exp(z)
|
| 77 |
+
return e / e.sum(axis=1, keepdims=True)
|
| 78 |
+
|
| 79 |
+
|
| 80 |
+
def _tta_views(img: Image.Image):
|
| 81 |
+
"""Slightly shifted crops of one cell - averaging their predictions
|
| 82 |
+
smooths out crop-alignment luck (measured: sheet 1 29 -> 36/39 exact)."""
|
| 83 |
+
w, h = img.size
|
| 84 |
+
views = []
|
| 85 |
+
for dx, dy in TTA_SHIFTS:
|
| 86 |
+
views.append(img.crop((max(0, dx), max(0, dy), w + min(0, dx), h + min(0, dy))))
|
| 87 |
+
return views
|
| 88 |
+
|
| 89 |
+
|
| 90 |
+
def classify_cells(images, clf: Classifier, temperature=1.0, tta=True, batch_size=64):
|
| 91 |
+
"""Class probabilities per cell, averaged over TTA views."""
|
| 92 |
+
views_per_cell = len(TTA_SHIFTS) if tta else 1
|
| 93 |
+
tensors = []
|
| 94 |
+
for img in images:
|
| 95 |
+
for view in _tta_views(img) if tta else [img]:
|
| 96 |
+
tensors.append(preprocess_pil(view))
|
| 97 |
+
if not tensors:
|
| 98 |
+
return np.empty((0, len(clf.classes)), dtype=np.float32)
|
| 99 |
+
|
| 100 |
+
probs_chunks = []
|
| 101 |
+
for i in range(0, len(tensors), batch_size):
|
| 102 |
+
batch = np.stack(tensors[i : i + batch_size]).astype(np.float32)
|
| 103 |
+
probs_chunks.append(_softmax(clf.run(batch), temperature))
|
| 104 |
+
stacked = np.concatenate(probs_chunks, axis=0)
|
| 105 |
+
return stacked.reshape(len(images), views_per_cell, -1).mean(axis=1)
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
# --------------------------------------------------------------------------- #
|
| 109 |
+
# Legal-move scoring, kazan evidence, beam search (numpy port)
|
| 110 |
+
# --------------------------------------------------------------------------- #
|
| 111 |
+
def _legal_scores(probs, legal_moves, classes, class_idx):
|
| 112 |
+
"""Redistribute the 162-class distribution over the current legal moves.
|
| 113 |
+
|
| 114 |
+
A legal move earns credit from (a) its exact class string and (b) a
|
| 115 |
+
factorized term P(source digit) * P(landing digit), times agreement with
|
| 116 |
+
the written x-mark probability. Scores stay UNNORMALIZED for the beam so a
|
| 117 |
+
hypothesis whose legal set explains the observation poorly accumulates a
|
| 118 |
+
genuinely low joint likelihood (measured: 24/39 vs 29/39 when normalized).
|
| 119 |
+
"""
|
| 120 |
+
source_marginal = [0.0] * 9
|
| 121 |
+
landing_marginal = [0.0] * 9
|
| 122 |
+
x_marginal = 0.0
|
| 123 |
+
for c, p in zip(classes, probs.tolist()):
|
| 124 |
+
if c == "empty":
|
| 125 |
+
continue
|
| 126 |
+
source_marginal[int(c[0]) - 1] += p
|
| 127 |
+
landing_marginal[int(c[1]) - 1] += p
|
| 128 |
+
if c.endswith("x"):
|
| 129 |
+
x_marginal += p
|
| 130 |
+
|
| 131 |
+
scored = []
|
| 132 |
+
for move in legal_moves:
|
| 133 |
+
exact = float(probs[class_idx[move.notation]])
|
| 134 |
+
factorized = source_marginal[move.action] * landing_marginal[int(move.notation[1]) - 1] * 9
|
| 135 |
+
x_agreement = x_marginal if move.makes_tuzdyk else 1.0 - x_marginal
|
| 136 |
+
scored.append(
|
| 137 |
+
((EXACT_WEIGHT * exact + FACTOR_WEIGHT * factorized) * x_agreement, move)
|
| 138 |
+
)
|
| 139 |
+
total = sum(s for s, _ in scored) or 1.0
|
| 140 |
+
return sorted(((s, s / total, m) for s, m in scored), key=lambda t: -t[0])
|
| 141 |
+
|
| 142 |
+
|
| 143 |
+
def _result_consistency(game: Game, result: str | None) -> int:
|
| 144 |
+
"""How well a final hypothesis state matches how the game really ended."""
|
| 145 |
+
if result is None:
|
| 146 |
+
return 1 if game.is_over else 0
|
| 147 |
+
want = RESULT_CODES[result]
|
| 148 |
+
if game.is_over and game.winner() == want:
|
| 149 |
+
return 2
|
| 150 |
+
kazans = game.kazans
|
| 151 |
+
leader = -1 if kazans[0] == kazans[1] else (0 if kazans[0] > kazans[1] else 1)
|
| 152 |
+
return 1 if leader == want else 0
|
| 153 |
+
|
| 154 |
+
|
| 155 |
+
def _kazan_logp(game: Game, checkpoint_probs) -> float:
|
| 156 |
+
"""Log-likelihood of a hypothesis' kazan counts under the checkpoint reads."""
|
| 157 |
+
logp = 0.0
|
| 158 |
+
for side, player in (("W", 0), ("B", 1)):
|
| 159 |
+
probs = checkpoint_probs.get(side)
|
| 160 |
+
if probs is None:
|
| 161 |
+
continue
|
| 162 |
+
value = game.kazans[player]
|
| 163 |
+
p = float(probs[value]) if value < len(probs) else 0.0
|
| 164 |
+
logp += math.log(max(p, KAZAN_FLOOR))
|
| 165 |
+
return logp
|
| 166 |
+
|
| 167 |
+
|
| 168 |
+
def beam_decode(ply_probs, classes, beam_width=1024, per_ply=9, result=None,
|
| 169 |
+
checkpoints=None):
|
| 170 |
+
"""Longest fully-legal move sequence with the highest joint probability,
|
| 171 |
+
exploring all `per_ply` legal continuations per ply; the final pool is
|
| 172 |
+
re-ranked by kazan-checkpoint and known-result consistency.
|
| 173 |
+
|
| 174 |
+
Returns (annotated_moves, log_prob, per_ply_share_list).
|
| 175 |
+
"""
|
| 176 |
+
checkpoints = checkpoints or {}
|
| 177 |
+
class_idx = {c: i for i, c in enumerate(classes)}
|
| 178 |
+
active = [(Game(), [], [], 0.0)] # (game, annotated_moves, chosen_shares, logp)
|
| 179 |
+
finished = []
|
| 180 |
+
|
| 181 |
+
for probs in ply_probs:
|
| 182 |
+
expanded = []
|
| 183 |
+
for game, moves, chosen, logp in active:
|
| 184 |
+
candidates = _legal_scores(probs, game.legal_moves(), classes, class_idx)
|
| 185 |
+
for raw, share, move in candidates[:per_ply]:
|
| 186 |
+
nxt = game.copy()
|
| 187 |
+
played = nxt.play(move.action)
|
| 188 |
+
annotated = played.notation + ("+" if played.captured else "")
|
| 189 |
+
new_logp = logp + math.log(max(raw, 1e-12))
|
| 190 |
+
if len(moves) + 1 in checkpoints:
|
| 191 |
+
new_logp += _kazan_logp(nxt, checkpoints[len(moves) + 1])
|
| 192 |
+
hyp = (nxt, moves + [annotated], chosen + [share], new_logp)
|
| 193 |
+
(finished if nxt.is_over else expanded).append(hyp)
|
| 194 |
+
if not expanded:
|
| 195 |
+
break
|
| 196 |
+
active = sorted(expanded, key=lambda h: -h[3])[:beam_width]
|
| 197 |
+
|
| 198 |
+
pool = active + finished
|
| 199 |
+
best = max(pool, key=lambda h: (len(h[1]), _result_consistency(h[0], result), h[3]))
|
| 200 |
+
return best[1], best[3], best[2]
|
| 201 |
+
|
| 202 |
+
|
| 203 |
+
def format_pgn(plies: list[str]) -> str:
|
| 204 |
+
"""plies -> '1. 54+ 42x' lines (one White/Black pair per line)."""
|
| 205 |
+
lines = []
|
| 206 |
+
for i in range(0, len(plies), 2):
|
| 207 |
+
pair = " ".join(plies[i : i + 2])
|
| 208 |
+
lines.append(f"{i // 2 + 1}. {pair}")
|
| 209 |
+
return "\n".join(lines) + "\n"
|
| 210 |
+
|
| 211 |
+
|
| 212 |
+
# --------------------------------------------------------------------------- #
|
| 213 |
+
# The single entry point
|
| 214 |
+
# --------------------------------------------------------------------------- #
|
| 215 |
+
def run_pipeline(image, moves_clf: Classifier, kazan_clf: Classifier | None = None,
|
| 216 |
+
result=None, topk=5, beam_width=1024, per_ply=9,
|
| 217 |
+
temperature=1.0, tta=True, save_cells_dir=None):
|
| 218 |
+
"""Read one scoresheet image into game records.
|
| 219 |
+
|
| 220 |
+
`image` is a path or a PIL.Image. Returns a dict with the reconstructed
|
| 221 |
+
PGNs and diagnostics; no files are written unless `save_cells_dir` is set
|
| 222 |
+
(the CLI uses it, the Gradio app does not).
|
| 223 |
+
|
| 224 |
+
Keys: beam_pgn, raw_pgn, legal_pgn, game_json (dict), stopped,
|
| 225 |
+
plies_scanned, legal_plies, beam_plies, kazan_report, annotated_image (PIL),
|
| 226 |
+
median_cell_height, low_resolution (bool).
|
| 227 |
+
"""
|
| 228 |
+
classes = moves_clf.classes
|
| 229 |
+
empty_idx = classes.index(EMPTY)
|
| 230 |
+
|
| 231 |
+
cells, sheet = extract_cells(image)
|
| 232 |
+
cells.sort(key=lambda c: (c.move_no, c.side != "W")) # game order: 1W 1B 2W ...
|
| 233 |
+
median_h = sorted(c.bbox[3] for c in cells)[len(cells) // 2]
|
| 234 |
+
|
| 235 |
+
cleaned = [clean_cell(c.image) for c in cells] # (image, ink_ratio) pairs
|
| 236 |
+
all_probs = classify_cells(
|
| 237 |
+
[img for img, _ in cleaned], moves_clf, temperature=temperature, tta=tta
|
| 238 |
+
)
|
| 239 |
+
|
| 240 |
+
game = Game()
|
| 241 |
+
plies, raw_pgn, legal_pgn = [], [], []
|
| 242 |
+
ply_prob_arrays = []
|
| 243 |
+
labels = {}
|
| 244 |
+
stopped = None
|
| 245 |
+
in_sync = True
|
| 246 |
+
|
| 247 |
+
cells_dir = Path(save_cells_dir) if save_cells_dir else None
|
| 248 |
+
if cells_dir:
|
| 249 |
+
cells_dir.mkdir(parents=True, exist_ok=True)
|
| 250 |
+
|
| 251 |
+
for cell, (clean_img, ink_ratio), probs in zip(cells, cleaned, all_probs):
|
| 252 |
+
if ink_ratio < MIN_INK_RATIO or int(np.argmax(probs)) == empty_idx:
|
| 253 |
+
stopped = {"move": cell.move_no, "side": cell.side, "reason": "empty cell"}
|
| 254 |
+
break
|
| 255 |
+
|
| 256 |
+
move_probs = probs.copy()
|
| 257 |
+
move_probs[empty_idx] = 0.0
|
| 258 |
+
raw = classes[int(np.argmax(move_probs))]
|
| 259 |
+
|
| 260 |
+
legal = {m.notation: m for m in game.legal_moves()} if in_sync else {}
|
| 261 |
+
topk_idx = np.argsort(-probs)[:topk]
|
| 262 |
+
entry = {
|
| 263 |
+
"move": cell.move_no,
|
| 264 |
+
"side": cell.side,
|
| 265 |
+
"bbox": list(cell.bbox),
|
| 266 |
+
"raw": raw,
|
| 267 |
+
"raw_is_legal": raw in legal if in_sync else None,
|
| 268 |
+
"legal": sorted(legal),
|
| 269 |
+
"top5": [[classes[int(i)], round(float(probs[int(i)]), 4)] for i in topk_idx],
|
| 270 |
+
"probs": {c: round(float(p), 6) for c, p in zip(classes, probs)},
|
| 271 |
+
}
|
| 272 |
+
plies.append(entry)
|
| 273 |
+
ply_prob_arrays.append(probs)
|
| 274 |
+
if cells_dir:
|
| 275 |
+
clean_img.save(cells_dir / f"{cell.move_no:02d}_{cell.side}.png")
|
| 276 |
+
|
| 277 |
+
if in_sync and raw in legal:
|
| 278 |
+
played = game.play_notation(raw)
|
| 279 |
+
annotated = played.notation + ("+" if played.captured else "")
|
| 280 |
+
raw_pgn.append(annotated)
|
| 281 |
+
legal_pgn.append(annotated)
|
| 282 |
+
if game.is_over:
|
| 283 |
+
stopped = {"move": cell.move_no, "side": cell.side, "reason": "game over",
|
| 284 |
+
"winner": ["P1", "P2", "draw"][game.winner()]}
|
| 285 |
+
in_sync = False
|
| 286 |
+
else:
|
| 287 |
+
if in_sync:
|
| 288 |
+
entry["legal_alternatives"] = sorted(
|
| 289 |
+
((n, round(float(probs[classes.index(n)]), 4)) for n in legal),
|
| 290 |
+
key=lambda t: -t[1],
|
| 291 |
+
)
|
| 292 |
+
stopped = stopped or {"move": cell.move_no, "side": cell.side,
|
| 293 |
+
"reason": f"illegal move {raw!r}"}
|
| 294 |
+
in_sync = False
|
| 295 |
+
raw_pgn.append(raw)
|
| 296 |
+
|
| 297 |
+
# kazan checkpoint numbers from the summary strips (every 10 moves)
|
| 298 |
+
checkpoints, checkpoint_report = {}, []
|
| 299 |
+
if kazan_clf is not None:
|
| 300 |
+
kz_cells = [(c, *clean_cell(c.image)) for c in extract_kazan_cells(sheet)]
|
| 301 |
+
kz_cells = [(c, img) for c, img, ink in kz_cells if ink >= MIN_KAZAN_INK]
|
| 302 |
+
if kz_cells:
|
| 303 |
+
kz_probs = classify_cells(
|
| 304 |
+
[img for _, img in kz_cells], kazan_clf,
|
| 305 |
+
temperature=temperature, tta=tta,
|
| 306 |
+
)
|
| 307 |
+
for (cell, img), probs in zip(kz_cells, kz_probs):
|
| 308 |
+
if kazan_clf.classes[int(np.argmax(probs))] == EMPTY:
|
| 309 |
+
continue
|
| 310 |
+
values = probs[:82] / max(float(probs[:82].sum()), 1e-9)
|
| 311 |
+
checkpoints.setdefault(cell.move_no * 2, {})[cell.side] = values
|
| 312 |
+
top = int(np.argmax(values))
|
| 313 |
+
checkpoint_report.append(
|
| 314 |
+
{"move": cell.move_no, "side": cell.side,
|
| 315 |
+
"read": f"{top:02d}", "prob": round(float(values[top]), 4)}
|
| 316 |
+
)
|
| 317 |
+
if cells_dir:
|
| 318 |
+
img.save(cells_dir / f"kazan_{cell.move_no:02d}_{cell.side}.png")
|
| 319 |
+
|
| 320 |
+
beam_moves, beam_logp, beam_shares = beam_decode(
|
| 321 |
+
ply_prob_arrays, classes, beam_width, per_ply,
|
| 322 |
+
result=result, checkpoints=checkpoints,
|
| 323 |
+
)
|
| 324 |
+
beam_detail = [
|
| 325 |
+
{"move": p["move"], "side": p["side"], "chosen": m,
|
| 326 |
+
"prob": round(q, 4), "agrees_with_raw": m.rstrip("+") == p["raw"]}
|
| 327 |
+
for p, m, q in zip(plies, beam_moves, beam_shares)
|
| 328 |
+
]
|
| 329 |
+
for p, m in zip(plies, beam_moves):
|
| 330 |
+
labels[(p["move"], p["side"])] = m
|
| 331 |
+
|
| 332 |
+
annotated = render_overlay(sheet, cells[: len(plies)], labels)
|
| 333 |
+
|
| 334 |
+
game_json = {
|
| 335 |
+
"plies_scanned": len(plies),
|
| 336 |
+
"legal_plies": len(legal_pgn),
|
| 337 |
+
"beam": {"plies": len(beam_moves), "log_prob": round(beam_logp, 3), "moves": beam_detail},
|
| 338 |
+
"kazan_checkpoints": checkpoint_report,
|
| 339 |
+
"stopped": stopped or {"reason": "sheet exhausted"},
|
| 340 |
+
"plies": plies,
|
| 341 |
+
}
|
| 342 |
+
return {
|
| 343 |
+
"beam_pgn": format_pgn(beam_moves),
|
| 344 |
+
"raw_pgn": format_pgn(raw_pgn),
|
| 345 |
+
"legal_pgn": format_pgn(legal_pgn),
|
| 346 |
+
"game_json": game_json,
|
| 347 |
+
"stopped": stopped or {"reason": "sheet exhausted"},
|
| 348 |
+
"plies_scanned": len(plies),
|
| 349 |
+
"legal_plies": len(legal_pgn),
|
| 350 |
+
"beam_plies": len(beam_moves),
|
| 351 |
+
"kazan_report": checkpoint_report,
|
| 352 |
+
"annotated_image": annotated,
|
| 353 |
+
"median_cell_height": int(median_h),
|
| 354 |
+
"low_resolution": median_h < MIN_CELL_HEIGHT,
|
| 355 |
+
}
|
togyz/preprocess.py
ADDED
|
@@ -0,0 +1,52 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""The single image -> array preprocessing used by training AND inference.
|
| 2 |
+
|
| 3 |
+
The legacy pipeline failed partly because train and inference normalized
|
| 4 |
+
differently. Every consumer (dataset, eval, predict, the ONNX pipeline) must
|
| 5 |
+
import from here.
|
| 6 |
+
|
| 7 |
+
This module is deliberately **torch-free** (numpy only) so the serving path
|
| 8 |
+
can import it without pulling in PyTorch. Torch consumers wrap the returned
|
| 9 |
+
array with ``torch.from_numpy(...)``.
|
| 10 |
+
"""
|
| 11 |
+
|
| 12 |
+
import numpy as np
|
| 13 |
+
from PIL import Image, ImageOps
|
| 14 |
+
|
| 15 |
+
TARGET_H = 64
|
| 16 |
+
TARGET_W = 128
|
| 17 |
+
MEAN = 0.5
|
| 18 |
+
STD = 0.5
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def _estimate_background(gray: Image.Image) -> int:
|
| 22 |
+
"""Median of the image border, so padding blends with the paper."""
|
| 23 |
+
arr = np.asarray(gray)
|
| 24 |
+
border = np.concatenate([arr[0, :], arr[-1, :], arr[:, 0], arr[:, -1]])
|
| 25 |
+
return int(np.median(border))
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
def preprocess_pil(img: Image.Image) -> np.ndarray:
|
| 29 |
+
"""PIL image (any mode/size) -> float32 array [1, TARGET_H, TARGET_W]."""
|
| 30 |
+
img = ImageOps.exif_transpose(img).convert("L")
|
| 31 |
+
w, h = img.size
|
| 32 |
+
target_ratio = TARGET_W / TARGET_H
|
| 33 |
+
|
| 34 |
+
# Pad (never crop) to the target aspect ratio, centered on paper-colored canvas.
|
| 35 |
+
if w / h < target_ratio:
|
| 36 |
+
new_w, new_h = int(round(h * target_ratio)), h
|
| 37 |
+
else:
|
| 38 |
+
new_w, new_h = w, int(round(w / target_ratio))
|
| 39 |
+
if (new_w, new_h) != (w, h):
|
| 40 |
+
canvas = Image.new("L", (new_w, new_h), color=_estimate_background(img))
|
| 41 |
+
canvas.paste(img, ((new_w - w) // 2, (new_h - h) // 2))
|
| 42 |
+
img = canvas
|
| 43 |
+
|
| 44 |
+
img = img.resize((TARGET_W, TARGET_H), Image.BILINEAR)
|
| 45 |
+
arr = np.asarray(img, dtype=np.float32) / 255.0
|
| 46 |
+
arr = arr[None, :, :] # add channel axis -> [1, H, W]
|
| 47 |
+
return (arr - MEAN) / STD
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def preprocess_file(path: str) -> np.ndarray:
|
| 51 |
+
with Image.open(path) as img:
|
| 52 |
+
return preprocess_pil(img)
|
togyz/rules.py
ADDED
|
@@ -0,0 +1,143 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Togyzkumalak rules engine, ported from the 9Q C++ engine
|
| 2 |
+
(~/Desktop/9Q/src/togyzkumalak_rules.cpp) and validated against its
|
| 3 |
+
shortest-game replay fixture (tests/test_rules.py).
|
| 4 |
+
|
| 5 |
+
Notation matches 9Q and the scoresheet classifier classes:
|
| 6 |
+
"{source_hole}{landing_column}" with an "x" suffix when the move creates a
|
| 7 |
+
tuzdyk, e.g. "98", "13x". Source hole and landing column are 1-9; the landing
|
| 8 |
+
column is the pit's index within its own side, regardless of side.
|
| 9 |
+
"""
|
| 10 |
+
|
| 11 |
+
from dataclasses import dataclass, field
|
| 12 |
+
|
| 13 |
+
WIN_THRESHOLD = 82
|
| 14 |
+
PITS_PER_SIDE = 9
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
@dataclass
|
| 18 |
+
class Move:
|
| 19 |
+
action: int # 0-8 source hole index on the mover's side
|
| 20 |
+
notation: str # e.g. "98" or "13x"
|
| 21 |
+
captured: int # stones captured by the landing rule (excl. tuzdyk income)
|
| 22 |
+
makes_tuzdyk: bool
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
@dataclass
|
| 26 |
+
class Game:
|
| 27 |
+
pits: list[int] = field(default_factory=lambda: [9] * 18) # P1: 0-8, P2: 9-17
|
| 28 |
+
kazans: list[int] = field(default_factory=lambda: [0, 0])
|
| 29 |
+
# tuzdyks[p] = column (0-8) OWNED BY p on the opponent's side, -1 = none
|
| 30 |
+
tuzdyks: list[int] = field(default_factory=lambda: [-1, -1])
|
| 31 |
+
to_play: int = 0
|
| 32 |
+
|
| 33 |
+
def copy(self) -> "Game":
|
| 34 |
+
return Game(self.pits[:], self.kazans[:], self.tuzdyks[:], self.to_play)
|
| 35 |
+
|
| 36 |
+
# --- core mechanics -------------------------------------------------
|
| 37 |
+
|
| 38 |
+
def _side_of(self, pit: int) -> int:
|
| 39 |
+
return pit // PITS_PER_SIDE
|
| 40 |
+
|
| 41 |
+
def _apply(self, action: int) -> Move:
|
| 42 |
+
"""Apply a (legal) move for `to_play`; returns the Move played."""
|
| 43 |
+
p = self.to_play
|
| 44 |
+
opp = 1 - p
|
| 45 |
+
start = p * PITS_PER_SIDE + action
|
| 46 |
+
stones = self.pits[start]
|
| 47 |
+
|
| 48 |
+
if stones == 1:
|
| 49 |
+
self.pits[start] = 0
|
| 50 |
+
last = (start + 1) % 18
|
| 51 |
+
self.pits[last] += 1
|
| 52 |
+
else:
|
| 53 |
+
self.pits[start] = 1
|
| 54 |
+
pit = start
|
| 55 |
+
for _ in range(stones - 1):
|
| 56 |
+
pit = (pit + 1) % 18
|
| 57 |
+
self.pits[pit] += 1
|
| 58 |
+
last = pit
|
| 59 |
+
|
| 60 |
+
# stones sown into a tuzdyk go to its owner's kazan
|
| 61 |
+
for owner in (0, 1):
|
| 62 |
+
col = self.tuzdyks[owner]
|
| 63 |
+
if col >= 0:
|
| 64 |
+
tuz_pit = (1 - owner) * PITS_PER_SIDE + col
|
| 65 |
+
if self.pits[tuz_pit]:
|
| 66 |
+
self.kazans[owner] += self.pits[tuz_pit]
|
| 67 |
+
self.pits[tuz_pit] = 0
|
| 68 |
+
|
| 69 |
+
captured = 0
|
| 70 |
+
makes_tuzdyk = False
|
| 71 |
+
if self._side_of(last) == opp:
|
| 72 |
+
col = last % PITS_PER_SIDE
|
| 73 |
+
count = self.pits[last]
|
| 74 |
+
if (
|
| 75 |
+
count == 3
|
| 76 |
+
and self.tuzdyks[p] == -1
|
| 77 |
+
and col != PITS_PER_SIDE - 1 # hole 9 can never become a tuzdyk
|
| 78 |
+
and self.tuzdyks[opp] != col # no mirrored tuzdyks
|
| 79 |
+
):
|
| 80 |
+
self.tuzdyks[p] = col
|
| 81 |
+
self.kazans[p] += 3
|
| 82 |
+
self.pits[last] = 0
|
| 83 |
+
makes_tuzdyk = True
|
| 84 |
+
elif count % 2 == 0:
|
| 85 |
+
captured = count
|
| 86 |
+
self.kazans[p] += count
|
| 87 |
+
self.pits[last] = 0
|
| 88 |
+
|
| 89 |
+
notation = f"{action + 1}{(last % PITS_PER_SIDE) + 1}"
|
| 90 |
+
if makes_tuzdyk:
|
| 91 |
+
notation += "x"
|
| 92 |
+
self.to_play = opp
|
| 93 |
+
return Move(action, notation, captured, makes_tuzdyk)
|
| 94 |
+
|
| 95 |
+
# --- public API -----------------------------------------------------
|
| 96 |
+
|
| 97 |
+
def legal_actions(self) -> list[int]:
|
| 98 |
+
p = self.to_play
|
| 99 |
+
opp = 1 - p
|
| 100 |
+
return [
|
| 101 |
+
a
|
| 102 |
+
for a in range(PITS_PER_SIDE)
|
| 103 |
+
if self.pits[p * PITS_PER_SIDE + a] > 0 and self.tuzdyks[opp] != a
|
| 104 |
+
]
|
| 105 |
+
|
| 106 |
+
def legal_moves(self) -> list[Move]:
|
| 107 |
+
"""The <=9 legal moves with their exact notation (via simulation)."""
|
| 108 |
+
return [self.copy()._apply(a) for a in self.legal_actions()]
|
| 109 |
+
|
| 110 |
+
def play(self, action: int) -> Move:
|
| 111 |
+
if action not in self.legal_actions():
|
| 112 |
+
raise ValueError(f"illegal action {action + 1} for player {self.to_play + 1}")
|
| 113 |
+
return self._apply(action)
|
| 114 |
+
|
| 115 |
+
def play_notation(self, notation: str) -> Move:
|
| 116 |
+
"""Play by full notation string; raises ValueError if not legal."""
|
| 117 |
+
for move in self.legal_moves():
|
| 118 |
+
if move.notation == notation:
|
| 119 |
+
return self._apply(move.action)
|
| 120 |
+
raise ValueError(f"move {notation!r} is not legal here")
|
| 121 |
+
|
| 122 |
+
@property
|
| 123 |
+
def is_over(self) -> bool:
|
| 124 |
+
return (
|
| 125 |
+
max(self.kazans) >= WIN_THRESHOLD
|
| 126 |
+
or not self.legal_actions()
|
| 127 |
+
)
|
| 128 |
+
|
| 129 |
+
def winner(self) -> int:
|
| 130 |
+
"""0/1 = player, -1 = draw. If the side to move is stuck (atsyrau),
|
| 131 |
+
the opponent collects the stones remaining on their own side."""
|
| 132 |
+
kazans = self.kazans[:]
|
| 133 |
+
if max(kazans) < WIN_THRESHOLD and not self.legal_actions():
|
| 134 |
+
opp = 1 - self.to_play
|
| 135 |
+
kazans[opp] += sum(
|
| 136 |
+
self.pits[opp * PITS_PER_SIDE : (opp + 1) * PITS_PER_SIDE]
|
| 137 |
+
)
|
| 138 |
+
if kazans[0] == kazans[1]:
|
| 139 |
+
return -1
|
| 140 |
+
return 0 if kazans[0] > kazans[1] else 1
|
| 141 |
+
|
| 142 |
+
def side_totals(self) -> tuple[int, int]:
|
| 143 |
+
return sum(self.pits[:PITS_PER_SIDE]), sum(self.pits[PITS_PER_SIDE:])
|
togyz/sheet.py
ADDED
|
@@ -0,0 +1,298 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Locate the 8 printed move tables on a scoresheet photo and crop move cells.
|
| 2 |
+
|
| 3 |
+
Federation scoresheet layout: two bands of 4 tables (moves 1-40 top,
|
| 4 |
+
41-80 bottom), each table = header row + 10 move rows with columns
|
| 5 |
+
[No | Bast. (White) | Kost. (Black)]. Only these tables are read; the
|
| 6 |
+
header, summary strips, and footer are ignored.
|
| 7 |
+
|
| 8 |
+
Table boxes are found via printed grid lines (morphology + connected
|
| 9 |
+
components). Row/column separators inside a table are refined from detected
|
| 10 |
+
lines when they are complete, otherwise the known uniform structure is used
|
| 11 |
+
(photos are often too low-res for reliable thin-line detection).
|
| 12 |
+
"""
|
| 13 |
+
|
| 14 |
+
from dataclasses import dataclass
|
| 15 |
+
|
| 16 |
+
import cv2
|
| 17 |
+
import numpy as np
|
| 18 |
+
from PIL import Image, ImageOps
|
| 19 |
+
|
| 20 |
+
TABLES = 8
|
| 21 |
+
ROWS_PER_TABLE = 10 # plus one header row
|
| 22 |
+
# column boundaries as fractions of table width: No | Bast | Kost
|
| 23 |
+
COLUMN_FRACTIONS = (0.0, 0.20, 0.60, 1.0)
|
| 24 |
+
CELL_MARGIN = 0.15 # expand crops; handwriting overflows the printed cells
|
| 25 |
+
|
| 26 |
+
|
| 27 |
+
@dataclass
|
| 28 |
+
class Cell:
|
| 29 |
+
move_no: int # 1-80
|
| 30 |
+
side: str # "W" (Bast.) or "B" (Kost.)
|
| 31 |
+
bbox: tuple[int, int, int, int] # x, y, w, h in original image coords
|
| 32 |
+
image: Image.Image
|
| 33 |
+
quad: np.ndarray | None = None # 4x2 original-image corners (tilt-aware)
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def _grid_mask(gray: np.ndarray) -> np.ndarray:
|
| 37 |
+
h, w = gray.shape
|
| 38 |
+
thr = cv2.adaptiveThreshold(
|
| 39 |
+
gray, 255, cv2.ADAPTIVE_THRESH_MEAN_C, cv2.THRESH_BINARY_INV, 25, 12
|
| 40 |
+
)
|
| 41 |
+
horiz = cv2.morphologyEx(
|
| 42 |
+
thr, cv2.MORPH_OPEN, cv2.getStructuringElement(cv2.MORPH_RECT, (w // 30, 1))
|
| 43 |
+
)
|
| 44 |
+
vert = cv2.morphologyEx(
|
| 45 |
+
thr, cv2.MORPH_OPEN, cv2.getStructuringElement(cv2.MORPH_RECT, (1, h // 40))
|
| 46 |
+
)
|
| 47 |
+
return cv2.dilate(cv2.add(horiz, vert), np.ones((3, 3), np.uint8))
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def _order_corners(points: np.ndarray) -> np.ndarray:
|
| 51 |
+
"""Order 4 points as top-left, top-right, bottom-right, bottom-left."""
|
| 52 |
+
points = points.reshape(4, 2).astype(np.float32)
|
| 53 |
+
sums, diffs = points.sum(axis=1), np.diff(points, axis=1).ravel()
|
| 54 |
+
return np.array(
|
| 55 |
+
[points[sums.argmin()], points[diffs.argmin()],
|
| 56 |
+
points[sums.argmax()], points[diffs.argmax()]],
|
| 57 |
+
np.float32,
|
| 58 |
+
)
|
| 59 |
+
|
| 60 |
+
|
| 61 |
+
def _find_table_quads(gray: np.ndarray) -> list[np.ndarray]:
|
| 62 |
+
"""Corner quads of the 8 move tables (phone photos are tilted, so tables
|
| 63 |
+
are general quadrilaterals, not axis-aligned rectangles)."""
|
| 64 |
+
h, w = gray.shape
|
| 65 |
+
mask = _grid_mask(gray)
|
| 66 |
+
n, labels, stats, _ = cv2.connectedComponentsWithStats(mask)
|
| 67 |
+
quads = []
|
| 68 |
+
for i in range(1, n):
|
| 69 |
+
if stats[i][2] <= w * 0.1 or stats[i][3] <= h * 0.1:
|
| 70 |
+
continue
|
| 71 |
+
component = (labels == i).astype(np.uint8)
|
| 72 |
+
contours, _ = cv2.findContours(component, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
|
| 73 |
+
hull = cv2.convexHull(max(contours, key=cv2.contourArea))
|
| 74 |
+
approx = cv2.approxPolyDP(hull, 0.02 * cv2.arcLength(hull, True), True)
|
| 75 |
+
if len(approx) == 4:
|
| 76 |
+
quad = _order_corners(approx)
|
| 77 |
+
else: # fall back to the min-area rectangle corners
|
| 78 |
+
quad = _order_corners(cv2.boxPoints(cv2.minAreaRect(hull)))
|
| 79 |
+
quads.append((cv2.contourArea(quad), quad))
|
| 80 |
+
if len(quads) < TABLES:
|
| 81 |
+
raise RuntimeError(
|
| 82 |
+
f"Found only {len(quads)} move tables (need {TABLES}). "
|
| 83 |
+
"Check photo quality/framing."
|
| 84 |
+
)
|
| 85 |
+
quads = [q for _, q in sorted(quads, key=lambda t: -t[0])[:TABLES]]
|
| 86 |
+
# split into top/bottom bands by y, order each band left to right
|
| 87 |
+
quads.sort(key=lambda q: q[:, 1].mean())
|
| 88 |
+
top = sorted(quads[:4], key=lambda q: q[:, 0].mean())
|
| 89 |
+
bottom = sorted(quads[4:], key=lambda q: q[:, 0].mean())
|
| 90 |
+
return top + bottom
|
| 91 |
+
|
| 92 |
+
|
| 93 |
+
def _detect_lines(mask: np.ndarray, axis: int, min_frac: float) -> list[int]:
|
| 94 |
+
"""Positions of long line clusters along `axis` (0=rows, 1=cols)."""
|
| 95 |
+
profile = mask.sum(axis=1 - axis)
|
| 96 |
+
limit = 255 * mask.shape[1 - axis] * min_frac
|
| 97 |
+
positions = np.where(profile > limit)[0]
|
| 98 |
+
clusters: list[list[int]] = []
|
| 99 |
+
for pos in positions:
|
| 100 |
+
if clusters and pos - clusters[-1][-1] <= 2:
|
| 101 |
+
clusters[-1].append(int(pos))
|
| 102 |
+
else:
|
| 103 |
+
clusters.append([int(pos)])
|
| 104 |
+
return [int(np.mean(c)) for c in clusters]
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
UPSCALE = 2 # rectified tables are rendered at 2x for a little more pixel room
|
| 108 |
+
|
| 109 |
+
|
| 110 |
+
def _rectify_table(gray: np.ndarray, quad: np.ndarray):
|
| 111 |
+
"""Warp a (possibly tilted) table quad to a flat rectangle.
|
| 112 |
+
|
| 113 |
+
Returns (warped gray image, inverse homography back to sheet coords).
|
| 114 |
+
"""
|
| 115 |
+
top = np.linalg.norm(quad[1] - quad[0])
|
| 116 |
+
bottom = np.linalg.norm(quad[2] - quad[3])
|
| 117 |
+
left = np.linalg.norm(quad[3] - quad[0])
|
| 118 |
+
right = np.linalg.norm(quad[2] - quad[1])
|
| 119 |
+
tw = int(round((top + bottom) / 2)) * UPSCALE
|
| 120 |
+
th = int(round((left + right) / 2)) * UPSCALE
|
| 121 |
+
target = np.array([[0, 0], [tw, 0], [tw, th], [0, th]], np.float32)
|
| 122 |
+
matrix = cv2.getPerspectiveTransform(quad, target)
|
| 123 |
+
warped = cv2.warpPerspective(gray, matrix, (tw, th), flags=cv2.INTER_CUBIC)
|
| 124 |
+
return warped, np.linalg.inv(matrix)
|
| 125 |
+
|
| 126 |
+
|
| 127 |
+
def _row_lines(warped: np.ndarray) -> list[int]:
|
| 128 |
+
"""Row separators in a rectified table; uniform fallback."""
|
| 129 |
+
h, w = warped.shape
|
| 130 |
+
thr = cv2.adaptiveThreshold(
|
| 131 |
+
warped, 255, cv2.ADAPTIVE_THRESH_MEAN_C, cv2.THRESH_BINARY_INV, 25, 12
|
| 132 |
+
)
|
| 133 |
+
horiz = cv2.morphologyEx(
|
| 134 |
+
thr, cv2.MORPH_OPEN, cv2.getStructuringElement(cv2.MORPH_RECT, (max(w // 3, 3), 1))
|
| 135 |
+
)
|
| 136 |
+
rows = _detect_lines(horiz, axis=0, min_frac=0.4)
|
| 137 |
+
if len(rows) != ROWS_PER_TABLE + 2: # header + 10 rows needs 12 lines
|
| 138 |
+
rows = [round(i * h / (ROWS_PER_TABLE + 1)) for i in range(ROWS_PER_TABLE + 2)]
|
| 139 |
+
return rows
|
| 140 |
+
|
| 141 |
+
|
| 142 |
+
def extract_cells(image) -> tuple[list[Cell], Image.Image]:
|
| 143 |
+
"""All 160 move cells (80 moves x W/B) plus the upright sheet image.
|
| 144 |
+
|
| 145 |
+
`image` is a path (str/Path) or a PIL.Image. Each table is
|
| 146 |
+
perspective-rectified before being split, so cell crops stay aligned even
|
| 147 |
+
on tilted phone photos.
|
| 148 |
+
"""
|
| 149 |
+
if isinstance(image, Image.Image):
|
| 150 |
+
pil = ImageOps.exif_transpose(image).convert("L")
|
| 151 |
+
else:
|
| 152 |
+
with Image.open(image) as img:
|
| 153 |
+
pil = ImageOps.exif_transpose(img).convert("L")
|
| 154 |
+
gray = np.asarray(pil)
|
| 155 |
+
|
| 156 |
+
cells: list[Cell] = []
|
| 157 |
+
for t, quad in enumerate(_find_table_quads(gray)):
|
| 158 |
+
warped, inverse = _rectify_table(gray, quad)
|
| 159 |
+
th, tw = warped.shape
|
| 160 |
+
rows = _row_lines(warped)
|
| 161 |
+
cols = [round(f * tw) for f in COLUMN_FRACTIONS]
|
| 162 |
+
for r in range(ROWS_PER_TABLE):
|
| 163 |
+
y0, y1 = rows[r + 1], rows[r + 2] # rows[0..1] is the header
|
| 164 |
+
for col_index, side in ((1, "W"), (2, "B")):
|
| 165 |
+
x0, x1 = cols[col_index], cols[col_index + 1]
|
| 166 |
+
mx = round((x1 - x0) * CELL_MARGIN)
|
| 167 |
+
my = round((y1 - y0) * CELL_MARGIN)
|
| 168 |
+
cx0, cy0 = max(0, x0 - mx), max(0, y0 - my)
|
| 169 |
+
cx1, cy1 = min(tw, x1 + mx), min(th, y1 + my)
|
| 170 |
+
crop = Image.fromarray(warped[cy0:cy1, cx0:cx1])
|
| 171 |
+
|
| 172 |
+
# map the cell corners back onto the original sheet
|
| 173 |
+
corners = np.array(
|
| 174 |
+
[[cx0, cy0], [cx1, cy0], [cx1, cy1], [cx0, cy1]], np.float32
|
| 175 |
+
).reshape(-1, 1, 2)
|
| 176 |
+
sheet_quad = cv2.perspectiveTransform(corners, inverse).reshape(4, 2)
|
| 177 |
+
x_min, y_min = sheet_quad.min(axis=0)
|
| 178 |
+
x_max, y_max = sheet_quad.max(axis=0)
|
| 179 |
+
bbox = (int(x_min), int(y_min), int(x_max - x_min), int(y_max - y_min))
|
| 180 |
+
cells.append(
|
| 181 |
+
Cell(t * ROWS_PER_TABLE + r + 1, side, bbox, crop, sheet_quad)
|
| 182 |
+
)
|
| 183 |
+
return cells, pil
|
| 184 |
+
|
| 185 |
+
|
| 186 |
+
def extract_kazan_cells(sheet: Image.Image) -> list[Cell]:
|
| 187 |
+
"""Kazan checkpoint boxes from the 8 summary strips between table bands.
|
| 188 |
+
|
| 189 |
+
After every 10 moves the scorer records both kazan counts next to a small
|
| 190 |
+
board diagram: the box protruding ABOVE the strip is the Kost. (Black)
|
| 191 |
+
kazan, the box BELOW is the Bast. (White) kazan. Returned as Cell objects
|
| 192 |
+
with move_no = the checkpoint move (10, 20, ... 80) and side "B"/"W".
|
| 193 |
+
The 2x9 board-pit grid itself is ignored (too messy to read reliably).
|
| 194 |
+
"""
|
| 195 |
+
gray = np.asarray(sheet)
|
| 196 |
+
h, w = gray.shape
|
| 197 |
+
mask = _grid_mask(gray)
|
| 198 |
+
table_area = float(np.median([cv2.contourArea(q) for q in _find_table_quads(gray)]))
|
| 199 |
+
|
| 200 |
+
n, _, stats, _ = cv2.connectedComponentsWithStats(mask)
|
| 201 |
+
strips = [
|
| 202 |
+
tuple(int(v) for v in stats[i][:4])
|
| 203 |
+
for i in range(1, n)
|
| 204 |
+
if stats[i][2] > w * 0.1
|
| 205 |
+
and h * 0.02 < stats[i][3] < h * 0.12
|
| 206 |
+
and stats[i][2] > 1.5 * stats[i][3]
|
| 207 |
+
and stats[i][2] * stats[i][3] < 0.9 * table_area
|
| 208 |
+
]
|
| 209 |
+
strips.sort(key=lambda b: (b[1] > h / 2, b[0])) # band, then left-to-right
|
| 210 |
+
|
| 211 |
+
cells: list[Cell] = []
|
| 212 |
+
for k, (x, y, bw, bh) in enumerate(strips[:8]):
|
| 213 |
+
checkpoint = (k + 1) * 10
|
| 214 |
+
sub = mask[y : y + bh, x : x + bw]
|
| 215 |
+
long_rows = np.where((sub > 0).sum(axis=1) > bw * 0.55)[0]
|
| 216 |
+
if len(long_rows) == 0:
|
| 217 |
+
continue
|
| 218 |
+
# kazan boxes are ~1 row tall; a tight zone avoids swallowing the
|
| 219 |
+
# neighboring table's header text below the strip
|
| 220 |
+
zones = {
|
| 221 |
+
"B": (max(0, y + int(long_rows.min()) - 30), y + int(long_rows.min()) - 1, False),
|
| 222 |
+
"W": (y + int(long_rows.max()) + 3, min(h, y + int(long_rows.max()) + 31), True),
|
| 223 |
+
}
|
| 224 |
+
for side, (gy0, gy1, truncate_below) in zones.items():
|
| 225 |
+
if gy1 - gy0 < 8:
|
| 226 |
+
continue
|
| 227 |
+
zone = mask[gy0:gy1, x : x + bw]
|
| 228 |
+
# stop at any full-width line (a neighboring table's grid)
|
| 229 |
+
full = np.where((zone > 0).sum(axis=1) > bw * 0.7)[0]
|
| 230 |
+
if len(full):
|
| 231 |
+
if truncate_below:
|
| 232 |
+
zone = zone[: full.min()]
|
| 233 |
+
else:
|
| 234 |
+
zone = zone[full.max() + 1 :]
|
| 235 |
+
gy0 += int(full.max()) + 1
|
| 236 |
+
ys, xs = np.where(zone > 0)
|
| 237 |
+
if len(xs) < 15:
|
| 238 |
+
continue
|
| 239 |
+
bx0, bx1 = int(xs.min()), int(xs.max())
|
| 240 |
+
by0, by1 = gy0 + int(ys.min()), gy0 + int(ys.max())
|
| 241 |
+
if bx1 - bx0 < 10 or by1 - by0 < 8:
|
| 242 |
+
continue
|
| 243 |
+
cx0, cy0 = max(0, x + bx0 - 3), max(0, by0 - 2)
|
| 244 |
+
cx1, cy1 = min(w, x + bx1 + 4), min(h, by1 + 3)
|
| 245 |
+
cells.append(Cell(
|
| 246 |
+
checkpoint, side, (cx0, cy0, cx1 - cx0, cy1 - cy0),
|
| 247 |
+
Image.fromarray(gray[cy0:cy1, cx0:cx1]),
|
| 248 |
+
))
|
| 249 |
+
return cells
|
| 250 |
+
|
| 251 |
+
|
| 252 |
+
def clean_cell(image: Image.Image) -> tuple[Image.Image, float]:
|
| 253 |
+
"""Remove printed grid-line fragments from a cell crop.
|
| 254 |
+
|
| 255 |
+
Returns the cleaned crop and the fraction of remaining ink pixels
|
| 256 |
+
(handwriting). Blank cells score near zero even when grid lines cross
|
| 257 |
+
the crop, so this is a classifier-independent emptiness signal.
|
| 258 |
+
"""
|
| 259 |
+
gray = np.asarray(image)
|
| 260 |
+
h, w = gray.shape
|
| 261 |
+
thr = cv2.adaptiveThreshold(
|
| 262 |
+
gray, 255, cv2.ADAPTIVE_THRESH_MEAN_C, cv2.THRESH_BINARY_INV, 25, 12
|
| 263 |
+
)
|
| 264 |
+
# printed lines: long straight runs spanning most of the crop
|
| 265 |
+
horiz = cv2.morphologyEx(
|
| 266 |
+
thr, cv2.MORPH_OPEN, cv2.getStructuringElement(cv2.MORPH_RECT, (max(3, int(w * 0.6)), 1))
|
| 267 |
+
)
|
| 268 |
+
vert = cv2.morphologyEx(
|
| 269 |
+
thr, cv2.MORPH_OPEN, cv2.getStructuringElement(cv2.MORPH_RECT, (1, max(3, int(h * 0.6))))
|
| 270 |
+
)
|
| 271 |
+
lines = cv2.dilate(cv2.add(horiz, vert), np.ones((3, 3), np.uint8))
|
| 272 |
+
|
| 273 |
+
background = int(np.median(gray[thr == 0])) if (thr == 0).any() else 255
|
| 274 |
+
cleaned = gray.copy()
|
| 275 |
+
cleaned[lines > 0] = background
|
| 276 |
+
|
| 277 |
+
ink = (thr > 0) & (lines == 0)
|
| 278 |
+
ink[: h // 8, :] = ink[-h // 8 :, :] = False # ignore crop edges: neighbor
|
| 279 |
+
ink[:, : w // 12] = ink[:, -w // 12 :] = False # rows/cells bleeding in
|
| 280 |
+
ink_ratio = float(ink.sum()) / (h * w)
|
| 281 |
+
return Image.fromarray(cleaned), ink_ratio
|
| 282 |
+
|
| 283 |
+
|
| 284 |
+
def render_overlay(sheet: Image.Image, cells: list[Cell], labels: dict | None = None) -> Image.Image:
|
| 285 |
+
"""Debug image: cell boxes (green) with optional predicted labels (red)."""
|
| 286 |
+
vis = cv2.cvtColor(np.asarray(sheet), cv2.COLOR_GRAY2BGR)
|
| 287 |
+
for cell in cells:
|
| 288 |
+
x, y, w, h = cell.bbox
|
| 289 |
+
if cell.quad is not None:
|
| 290 |
+
cv2.polylines(vis, [cell.quad.astype(np.int32)], True, (0, 180, 0), 1)
|
| 291 |
+
else:
|
| 292 |
+
cv2.rectangle(vis, (x, y), (x + w, y + h), (0, 180, 0), 1)
|
| 293 |
+
if labels:
|
| 294 |
+
text = labels.get((cell.move_no, cell.side))
|
| 295 |
+
if text:
|
| 296 |
+
cv2.putText(vis, text, (x + 2, y + h - 3),
|
| 297 |
+
cv2.FONT_HERSHEY_SIMPLEX, 0.45, (0, 0, 255), 1, cv2.LINE_AA)
|
| 298 |
+
return Image.fromarray(cv2.cvtColor(vis, cv2.COLOR_BGR2RGB))
|
togyz/synth.py
ADDED
|
@@ -0,0 +1,238 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Synthesizes realistic scoresheet-cell photos for any of the 163 classes.
|
| 2 |
+
|
| 3 |
+
Design targets, derived from the real crops in data/real_crops/:
|
| 4 |
+
* digits are large and fill most of the crop height (tight framing)
|
| 5 |
+
* cursive right slant, overlapping kerning, variable stroke thickness
|
| 6 |
+
* mid-gray paper with shading, blue-ink-turned-gray strokes (not black)
|
| 7 |
+
* phone-camera artifacts: blur, sensor noise, JPEG compression
|
| 8 |
+
* crop aspect ratio anywhere between ~1:1 and ~3.3:1
|
| 9 |
+
|
| 10 |
+
Everything is driven by an explicit `random.Random` so the validation set
|
| 11 |
+
can be fully deterministic while training data stays infinite.
|
| 12 |
+
"""
|
| 13 |
+
|
| 14 |
+
import argparse
|
| 15 |
+
import io
|
| 16 |
+
import math
|
| 17 |
+
import random
|
| 18 |
+
|
| 19 |
+
import cv2
|
| 20 |
+
import numpy as np
|
| 21 |
+
from PIL import Image
|
| 22 |
+
|
| 23 |
+
from .classes import CLASSES, EMPTY
|
| 24 |
+
from .glyphs import GlyphSampler
|
| 25 |
+
|
| 26 |
+
# working resolution: glyphs are rendered at this digit height, the finished
|
| 27 |
+
# cell is later downscaled to a random "photo" size (anti-aliased strokes)
|
| 28 |
+
RENDER_DIGIT_H = 96
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def _transform_mask(mask: np.ndarray, rot_deg: float, shear: float) -> np.ndarray:
|
| 32 |
+
"""Rotate + italic-shear a glyph mask on an expanded canvas."""
|
| 33 |
+
h, w = mask.shape
|
| 34 |
+
pad = int(max(h, w) * 0.4) + 2
|
| 35 |
+
mask = np.pad(mask, pad)
|
| 36 |
+
ph, pw = mask.shape
|
| 37 |
+
center = (pw / 2, ph / 2)
|
| 38 |
+
m = cv2.getRotationMatrix2D(center, rot_deg, 1.0)
|
| 39 |
+
# italic shear: top of the glyph shifts right relative to the bottom
|
| 40 |
+
shear_m = np.array([[1.0, -shear, shear * ph / 2], [0.0, 1.0, 0.0]])
|
| 41 |
+
full = np.vstack([shear_m, [0, 0, 1]]) @ np.vstack([m, [0, 0, 1]])
|
| 42 |
+
out = cv2.warpAffine(mask, full[:2], (pw, ph), flags=cv2.INTER_LINEAR)
|
| 43 |
+
ys, xs = np.where(out > 0.1)
|
| 44 |
+
if len(ys) == 0:
|
| 45 |
+
return mask
|
| 46 |
+
return out[ys.min() : ys.max() + 1, xs.min() : xs.max() + 1]
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
def _vary_thickness(mask: np.ndarray, rng: random.Random) -> np.ndarray:
|
| 50 |
+
op = rng.choice(["dilate", "none", "dilate", "erode", "none"])
|
| 51 |
+
if op == "none":
|
| 52 |
+
return mask
|
| 53 |
+
kernel = np.ones((2, 2), np.uint8)
|
| 54 |
+
fn = cv2.dilate if op == "dilate" else cv2.erode
|
| 55 |
+
return fn(mask, kernel, iterations=1)
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
def _paper(rng: random.Random, h: int, w: int) -> np.ndarray:
|
| 59 |
+
base = rng.uniform(160, 235)
|
| 60 |
+
img = np.full((h, w), base, np.float32)
|
| 61 |
+
ys, xs = np.mgrid[0:h, 0:w].astype(np.float32)
|
| 62 |
+
# linear lighting gradient
|
| 63 |
+
gx, gy = rng.uniform(-1, 1), rng.uniform(-1, 1)
|
| 64 |
+
plane = gx * xs / max(w, 1) + gy * ys / max(h, 1)
|
| 65 |
+
ptp = plane.max() - plane.min()
|
| 66 |
+
if ptp > 1e-6:
|
| 67 |
+
img += rng.uniform(0, 22) * ((plane - plane.min()) / ptp - 0.5)
|
| 68 |
+
# occasional soft shadow blob
|
| 69 |
+
if rng.random() < 0.4:
|
| 70 |
+
cx, cy = rng.uniform(0, w), rng.uniform(0, h)
|
| 71 |
+
r = rng.uniform(0.5, 1.3) * max(w, h)
|
| 72 |
+
d2 = ((xs - cx) ** 2 + (ys - cy) ** 2) / (r * r)
|
| 73 |
+
img -= rng.uniform(5, 20) * np.exp(-d2)
|
| 74 |
+
return img
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def _compose_ink(paper: np.ndarray, alpha: np.ndarray, rng: random.Random) -> np.ndarray:
|
| 78 |
+
"""Blend an ink alpha mask onto paper with per-pixel ink tone variation."""
|
| 79 |
+
ink_base = rng.uniform(30, 110)
|
| 80 |
+
ink = ink_base + np.random.default_rng(rng.getrandbits(63)).normal(
|
| 81 |
+
0, 8, paper.shape
|
| 82 |
+
).astype(np.float32)
|
| 83 |
+
alpha = np.clip(alpha, 0.0, 1.0) * rng.uniform(0.85, 1.0)
|
| 84 |
+
return paper * (1 - alpha) + ink * alpha
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
def _add_border_lines(img: np.ndarray, rng: random.Random) -> np.ndarray:
|
| 88 |
+
"""Fragment of the printed cell border caught by an imperfect crop."""
|
| 89 |
+
h, w = img.shape
|
| 90 |
+
edge = rng.choice(["top", "bottom", "left", "right"])
|
| 91 |
+
offset = rng.randint(0, max(1, int(0.06 * min(h, w))))
|
| 92 |
+
darkness = rng.uniform(60, 140)
|
| 93 |
+
thickness = rng.randint(1, 3)
|
| 94 |
+
if edge in ("top", "bottom"):
|
| 95 |
+
y = offset if edge == "top" else h - 1 - offset
|
| 96 |
+
x0, x1 = sorted((rng.randint(0, w // 2), rng.randint(w // 2, w - 1)))
|
| 97 |
+
cv2.line(img, (x0, y), (x1, y), darkness, thickness, cv2.LINE_AA)
|
| 98 |
+
else:
|
| 99 |
+
x = offset if edge == "left" else w - 1 - offset
|
| 100 |
+
y0, y1 = sorted((rng.randint(0, h // 2), rng.randint(h // 2, h - 1)))
|
| 101 |
+
cv2.line(img, (x, y0), (x, y1), darkness, thickness, cv2.LINE_AA)
|
| 102 |
+
return img
|
| 103 |
+
|
| 104 |
+
|
| 105 |
+
def _camera_effects(img: np.ndarray, rng: random.Random) -> np.ndarray:
|
| 106 |
+
# optical blur
|
| 107 |
+
img = cv2.GaussianBlur(img, (0, 0), rng.uniform(0.4, 1.4))
|
| 108 |
+
# occasional slight motion blur
|
| 109 |
+
if rng.random() < 0.15:
|
| 110 |
+
k = rng.choice([3, 5])
|
| 111 |
+
kernel = np.zeros((k, k), np.float32)
|
| 112 |
+
angle = rng.uniform(0, math.pi)
|
| 113 |
+
cv2.line(
|
| 114 |
+
kernel,
|
| 115 |
+
(0, int((k - 1) / 2 * (1 - math.sin(angle)))),
|
| 116 |
+
(k - 1, int((k - 1) / 2 * (1 + math.sin(angle)))),
|
| 117 |
+
1.0,
|
| 118 |
+
1,
|
| 119 |
+
)
|
| 120 |
+
kernel /= max(kernel.sum(), 1e-6)
|
| 121 |
+
img = cv2.filter2D(img, -1, kernel)
|
| 122 |
+
# sensor noise
|
| 123 |
+
noise_rng = np.random.default_rng(rng.getrandbits(63))
|
| 124 |
+
img = img + noise_rng.normal(0, rng.uniform(1.5, 7.0), img.shape).astype(np.float32)
|
| 125 |
+
# brightness / contrast jitter
|
| 126 |
+
img = (img - 128.0) * rng.uniform(0.82, 1.15) + 128.0 + rng.uniform(-15, 15)
|
| 127 |
+
img = np.clip(img, 0, 255).astype(np.uint8)
|
| 128 |
+
# JPEG artifacts
|
| 129 |
+
if rng.random() < 0.6:
|
| 130 |
+
ok, buf = cv2.imencode(".jpg", img, [cv2.IMWRITE_JPEG_QUALITY, rng.randint(25, 85)])
|
| 131 |
+
if ok:
|
| 132 |
+
img = cv2.imdecode(buf, cv2.IMREAD_GRAYSCALE)
|
| 133 |
+
return img
|
| 134 |
+
|
| 135 |
+
|
| 136 |
+
def synthesize_cell(
|
| 137 |
+
class_name: str, sampler: GlyphSampler, rng: random.Random
|
| 138 |
+
) -> Image.Image:
|
| 139 |
+
"""Render one cell photo for `class_name` ('11'..'99x' or 'empty')."""
|
| 140 |
+
if class_name == EMPTY:
|
| 141 |
+
h = rng.randint(60, 180)
|
| 142 |
+
w = int(h * rng.uniform(1.2, 3.3))
|
| 143 |
+
img = _paper(rng, h, w)
|
| 144 |
+
if rng.random() < 0.25:
|
| 145 |
+
img = _add_border_lines(img, rng)
|
| 146 |
+
return Image.fromarray(_camera_effects(img, rng))
|
| 147 |
+
|
| 148 |
+
# --- prepare glyph masks ---
|
| 149 |
+
slant = rng.uniform(-0.10, 0.40) # shared cursive slant, biased rightward
|
| 150 |
+
masks = []
|
| 151 |
+
for char in class_name:
|
| 152 |
+
mask = sampler.sample(char, rng)
|
| 153 |
+
rel_h = rng.uniform(0.5, 0.8) if char == "x" else rng.uniform(0.85, 1.15)
|
| 154 |
+
target_h = max(8, int(RENDER_DIGIT_H * rel_h))
|
| 155 |
+
target_w = max(4, int(mask.shape[1] * target_h / mask.shape[0]))
|
| 156 |
+
mask = cv2.resize(mask, (target_w, target_h), interpolation=cv2.INTER_LINEAR)
|
| 157 |
+
mask = _transform_mask(mask, rng.uniform(-7, 7), slant + rng.uniform(-0.05, 0.05))
|
| 158 |
+
mask = _vary_thickness(mask, rng)
|
| 159 |
+
masks.append(mask)
|
| 160 |
+
|
| 161 |
+
gaps = [
|
| 162 |
+
int(rng.uniform(-0.18, 0.15) * masks[i].shape[1])
|
| 163 |
+
for i in range(len(masks) - 1)
|
| 164 |
+
]
|
| 165 |
+
total_w = sum(m.shape[1] for m in masks) + sum(gaps)
|
| 166 |
+
max_h = max(m.shape[0] for m in masks)
|
| 167 |
+
|
| 168 |
+
# --- cell geometry: digits fill `fill` of the crop height ---
|
| 169 |
+
fill = rng.uniform(0.45, 0.92)
|
| 170 |
+
cell_h = int(max_h / fill)
|
| 171 |
+
aspect = rng.uniform(0.95, 3.3)
|
| 172 |
+
cell_w = max(int(cell_h * aspect), int(total_w * rng.uniform(1.02, 1.2)))
|
| 173 |
+
|
| 174 |
+
# horizontal placement: real crops are sometimes left-aligned with empty space
|
| 175 |
+
slack = cell_w - total_w
|
| 176 |
+
align = rng.random()
|
| 177 |
+
if align < 0.35: # left
|
| 178 |
+
x = int(slack * rng.uniform(0.0, 0.15))
|
| 179 |
+
elif align < 0.85: # center-ish
|
| 180 |
+
x = int(slack * rng.uniform(0.25, 0.6))
|
| 181 |
+
else: # right
|
| 182 |
+
x = int(slack * rng.uniform(0.7, 0.95))
|
| 183 |
+
|
| 184 |
+
# --- compose ink alpha on full-cell canvas ---
|
| 185 |
+
alpha = np.zeros((cell_h, cell_w), np.float32)
|
| 186 |
+
y_center = (cell_h - max_h) / 2
|
| 187 |
+
for i, mask in enumerate(masks):
|
| 188 |
+
mh, mw = mask.shape
|
| 189 |
+
y = int(y_center + (max_h - mh) * rng.uniform(0.2, 0.8) + rng.uniform(-0.04, 0.04) * cell_h)
|
| 190 |
+
y = min(max(y, 0), cell_h - mh)
|
| 191 |
+
x = min(max(x, 0), cell_w - mw)
|
| 192 |
+
region = alpha[y : y + mh, x : x + mw]
|
| 193 |
+
np.maximum(region, mask, out=region)
|
| 194 |
+
if i < len(gaps):
|
| 195 |
+
x += mw + gaps[i]
|
| 196 |
+
|
| 197 |
+
img = _compose_ink(_paper(rng, cell_h, cell_w), alpha, rng)
|
| 198 |
+
if rng.random() < 0.12:
|
| 199 |
+
img = _add_border_lines(img, rng)
|
| 200 |
+
|
| 201 |
+
# downscale to a random photo resolution (real crops are 115-260 px tall)
|
| 202 |
+
out_h = rng.randint(48, 190)
|
| 203 |
+
out_w = max(16, int(cell_w * out_h / cell_h))
|
| 204 |
+
img = cv2.resize(img, (out_w, out_h), interpolation=cv2.INTER_AREA)
|
| 205 |
+
|
| 206 |
+
return Image.fromarray(_camera_effects(img, rng))
|
| 207 |
+
|
| 208 |
+
|
| 209 |
+
def _preview(path: str, seed: int, count: int) -> None:
|
| 210 |
+
"""Render a comparison grid: the real-crop classes first, then random ones."""
|
| 211 |
+
sampler = GlyphSampler()
|
| 212 |
+
rng = random.Random(seed)
|
| 213 |
+
fixed = ["96", "77", "13x", "37", "85", "42"] * 2
|
| 214 |
+
names = fixed + [rng.choice(CLASSES) for _ in range(max(0, count - len(fixed)))]
|
| 215 |
+
|
| 216 |
+
tile_w, tile_h, caption = 200, 100, 14
|
| 217 |
+
cols = 6
|
| 218 |
+
rows = math.ceil(len(names) / cols)
|
| 219 |
+
sheet = Image.new("L", (cols * tile_w, rows * (tile_h + caption)), 255)
|
| 220 |
+
from PIL import ImageDraw
|
| 221 |
+
|
| 222 |
+
draw = ImageDraw.Draw(sheet)
|
| 223 |
+
for i, name in enumerate(names):
|
| 224 |
+
cell = synthesize_cell(name, sampler, rng).resize((tile_w, tile_h))
|
| 225 |
+
cx, cy = (i % cols) * tile_w, (i // cols) * (tile_h + caption)
|
| 226 |
+
sheet.paste(cell, (cx, cy + caption))
|
| 227 |
+
draw.text((cx + 4, cy + 1), name, fill=0)
|
| 228 |
+
sheet.save(path)
|
| 229 |
+
print(f"Saved {len(names)}-cell preview to {path}")
|
| 230 |
+
|
| 231 |
+
|
| 232 |
+
if __name__ == "__main__":
|
| 233 |
+
parser = argparse.ArgumentParser(description=__doc__)
|
| 234 |
+
parser.add_argument("--preview", default="preview.png", help="output image path")
|
| 235 |
+
parser.add_argument("--seed", type=int, default=0)
|
| 236 |
+
parser.add_argument("--n", type=int, default=48, help="number of cells")
|
| 237 |
+
args = parser.parse_args()
|
| 238 |
+
_preview(args.preview, args.seed, args.n)
|