ansarzeinulla commited on
Commit
20d7fde
·
0 Parent(s):

Initial commit

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
.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)