ansarzeinulla commited on
Commit
1b69166
·
1 Parent(s): 9a0144f

Read full board diagrams, dynamic beam width, CSV metadata batch UI

Browse files

- sheet.py: extract_diagram_cells reads the 2x9 pit grid + kazan boxes per
summary strip, table-guided strip location (splits merged components,
positional checkpoint numbering), median-area filter in _find_table_quads,
gridline export, richer render_overlay (kazans blue, pits orange, gridlines)
- classes/synth/dataset: unified 85-class diagram task (0-81, x, -, empty)
with overflow/neighbor-bleed synthesis; --task diagram everywhere
- pipeline: _diagram_logp scores kazans + pits per checkpoint, beam width
user-selectable with memory guard (safe_beam_width) + heapq.nlargest,
checkpoint_report/warnings in output
- app.py: batch of 5 images with shared Tournament,Location,Date,Round CSV
field, per-game White,Black,Result,WhiteTime,BlackTime lines, beam-width
dropdown 2^10..2^30, PGN tags in all outputs
- Colab notebook rewritten: clone repo, fetch glyphs at runtime, train both
tasks, export ONNX, download artifacts.zip

Files changed (11) hide show
  1. README.md +26 -15
  2. app.py +168 -45
  3. notebooks/colab_train.ipynb +70 -21
  4. read_game.py +16 -11
  5. scripts/export_onnx.py +3 -3
  6. togyz/classes.py +7 -3
  7. togyz/dataset.py +4 -3
  8. togyz/pipeline.py +127 -34
  9. togyz/sheet.py +163 -20
  10. togyz/synth.py +111 -6
  11. train.py +14 -9
README.md CHANGED
@@ -19,16 +19,20 @@ records, and a **Gradio demo** (`app.py`) packages it for a HuggingFace Space.
19
 
20
  ## Live demo (HuggingFace Space)
21
 
22
- `app.py` is a Gradio app: upload up to **5** scoresheet photos, optionally tag
23
- each game's known result, and download the reconstructed PGNs (beam / raw /
24
- legal) as a zip. Inference runs on the exported **ONNX** models (torch-free),
25
- so the Space needs only the light serving deps in `requirements.txt`.
 
 
 
 
26
 
27
  ```bash
28
  python scripts/export_onnx.py # checkpoints/*.pt -> single-file .onnx (+ .classes.json)
29
  cp checkpoints/best.onnx checkpoints/best.classes.json models/
30
- cp checkpoints/kazan/best.onnx models/kazan.onnx
31
- cp checkpoints/kazan/best.classes.json models/kazan.classes.json
32
  python app.py # serve locally at http://127.0.0.1:7860
33
  ```
34
 
@@ -85,12 +89,15 @@ A real training run needs a GPU — see Colab below. Defaults
85
  Open `notebooks/colab_train.ipynb` in Colab, or manually:
86
 
87
  ```bash
88
- git clone <this-repo> && cd 9OCR # or upload a zip of the project
89
- pip install -r requirements.txt
90
  python scripts/get_glyphs.py --emnist # ingredients/ is gitignored, so also
91
  # rebuild the EMNIST pool from torchvision
92
- python train.py --epochs 30
93
- # download checkpoints/best.pt when done
 
 
 
94
  ```
95
 
96
  `predict.py`/`eval.py` run anywhere (CPU is fine) with the downloaded
@@ -117,11 +124,15 @@ Pass `--result` (1-0 / 0-1 / draw, from the sheet footer) when known: the
117
  beam's final pool is re-ranked to prefer reconstructions whose end state is
118
  consistent with how the game actually ended.
119
 
120
- The summary strips between the tables record both kazan counts every 10
121
- moves (Black's box above the strip, White's below, always two digits 00-81).
122
- A second classifier reads them (`python train.py --task kazan`, saved to
123
- `checkpoints/kazan/best.pt`) and the beam gains likelihood when a
124
- hypothesis' computed kazans match these written checkpoints.
 
 
 
 
125
 
126
  Finds the 8 printed move tables, classifies every cell, and writes:
127
 
 
19
 
20
  ## Live demo (HuggingFace Space)
21
 
22
+ `app.py` is a Gradio app: upload up to **5** scoresheet photos of one round as
23
+ a single batch, describe the tournament (`Tournament,Location,Date,Round`) and
24
+ each game (`WhiteName,BlackName,Result,WhiteTime,BlackTime`, one CSV line per
25
+ image) in two text fields, pick a beam width (1024 … 2^30; oversized requests
26
+ are trimmed to fit memory), and download the reconstructed PGNs (beam / raw /
27
+ legal, with PGN tags) as a zip. Inference runs on the exported **ONNX** models
28
+ (torch-free), so the Space needs only the light serving deps in
29
+ `requirements.txt`.
30
 
31
  ```bash
32
  python scripts/export_onnx.py # checkpoints/*.pt -> single-file .onnx (+ .classes.json)
33
  cp checkpoints/best.onnx checkpoints/best.classes.json models/
34
+ cp checkpoints/diagram/best.onnx models/diagram.onnx
35
+ cp checkpoints/diagram/best.classes.json models/diagram.classes.json
36
  python app.py # serve locally at http://127.0.0.1:7860
37
  ```
38
 
 
89
  Open `notebooks/colab_train.ipynb` in Colab, or manually:
90
 
91
  ```bash
92
+ git clone https://github.com/ansarzeinulla/9OCR.git && cd 9OCR
93
+ pip install -r requirements-train.txt
94
  python scripts/get_glyphs.py --emnist # ingredients/ is gitignored, so also
95
  # rebuild the EMNIST pool from torchvision
96
+ python train.py --task moves --epochs 30
97
+ python train.py --task diagram --epochs 30
98
+ python scripts/export_onnx.py
99
+ # download checkpoints/best.onnx (+.classes.json) and
100
+ # checkpoints/diagram/best.onnx (+.classes.json) when done
101
  ```
102
 
103
  `predict.py`/`eval.py` run anywhere (CPU is fine) with the downloaded
 
124
  beam's final pool is re-ranked to prefer reconstructions whose end state is
125
  consistent with how the game actually ended.
126
 
127
+ The summary strips between the tables hold a full board diagram every 10
128
+ moves: both kazan counts (Black's box above the strip, White's below, two
129
+ digits 10-81) **and** the 2×9 pit grid (upper row = Black, pits 9…1
130
+ left-to-right; lower row = White, pits 1…9; each cell is `x` = tuzdyk,
131
+ `-` = 0, or a 1-2 digit count). A second, unified classifier reads all of it
132
+ (`python train.py --task diagram`, saved to `checkpoints/diagram/best.pt`;
133
+ classes 0-81 + `x` + `-` + `empty`, filtered per context at inference), and
134
+ the beam gains likelihood when a hypothesis' computed kazans and pit counts
135
+ match these written checkpoints.
136
 
137
  Finds the 8 printed move tables, classifies every cell, and writes:
138
 
app.py CHANGED
@@ -1,8 +1,9 @@
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
@@ -26,6 +27,8 @@ except Exception as e:
26
  print(f"[app] Failed to apply monkey-patch: {e}", flush=True)
27
  # --------------------------------------
28
 
 
 
29
  import os
30
  import tempfile
31
  import zipfile
@@ -52,16 +55,29 @@ from togyz.pipeline import load_classifier, run_pipeline
52
 
53
  MAX_IMAGES = 5
54
  MODEL_DIR = Path(__file__).parent / "models"
55
- RESULT_CHOICES = ["unknown", "1-0", "0-1", "draw"]
 
 
 
 
 
 
 
 
 
 
56
 
57
  # Load the ONNX sessions once at import - warm for the whole process lifetime.
58
  _log(f"loading move model from {MODEL_DIR / 'best.onnx'} ...")
59
  _MOVES = load_classifier(MODEL_DIR / "best.onnx")
60
- _KAZAN = None
61
- _kazan_path = MODEL_DIR / "kazan.onnx"
62
- if _kazan_path.exists():
63
- _log("loading kazan model ...")
64
- _KAZAN = load_classifier(_kazan_path)
 
 
 
65
  _log("models loaded")
66
 
67
 
@@ -70,10 +86,94 @@ def _safe_slug(text: str) -> str:
70
  return keep.strip("_")
71
 
72
 
73
- def _process_one(image_path, result_choice, base_name, out_dir: Path):
74
- """Run the pipeline on one image; return (row, gallery_item, pgn_files)."""
75
- result = None if result_choice in (None, "unknown") else result_choice
76
- out = run_pipeline(image_path, _MOVES, _KAZAN, result=result)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
77
 
78
  stop = out["stopped"]
79
  stop_txt = stop.get("reason", "")
@@ -81,53 +181,62 @@ def _process_one(image_path, result_choice, base_name, out_dir: Path):
81
  stop_txt += f" ({stop['winner']})"
82
  note = " ⚠ low-res" if out["low_resolution"] else ""
83
 
 
84
  files = []
85
  for kind in ("beam", "raw", "legal"):
86
  f = out_dir / f"{base_name}_{kind}.pgn"
87
- f.write_text(out[f"{kind}_pgn"])
88
  files.append(str(f))
89
 
90
  row = [base_name, out["beam_plies"], stop_txt + note, out["beam_pgn"].strip()]
91
  caption = f"{base_name}: {out['beam_plies']} plies"
92
- return row, (out["annotated_image"], caption), files
93
-
94
 
95
- def convert(round_label, *slot_values):
96
- """slot_values = [img1, res1, img2, res2, ...] for the 5 fixed slots."""
97
- images = slot_values[0::2]
98
- results = slot_values[1::2]
99
- provided = [(img, res) for img, res in zip(images, results) if img]
100
 
101
- if not provided:
 
 
 
102
  raise gr.Error("Please upload at least one scoresheet image.")
103
- if len(provided) > MAX_IMAGES: # defensive; the UI only exposes 5 slots
104
- raise gr.Error(f"This demo handles at most {MAX_IMAGES} images per run.")
 
 
 
 
 
 
105
 
106
  out_dir = Path(tempfile.mkdtemp(prefix="togyz_"))
107
- round_slug = _safe_slug(round_label)
108
- rows, gallery, all_files = [], [], []
109
 
110
- for i, (img, res) in enumerate(provided, start=1):
111
- prefix = f"{round_slug}_table{i}" if round_slug else f"table{i}"
112
  try:
113
- row, gal, files = _process_one(img, res, prefix, out_dir)
 
 
114
  except Exception as exc: # one bad image must not kill the batch
115
  rows.append([prefix, 0, f"error: {exc}", ""])
116
  continue
117
  rows.append(row)
118
  gallery.append(gal)
119
  all_files.extend(files)
 
120
 
 
121
  if not all_files:
122
  # every image errored - still return the table so the user sees why
123
- return rows, gallery, None
124
 
125
- zip_path = out_dir / (f"{round_slug}_pgns.zip" if round_slug else "pgns.zip")
126
  with zipfile.ZipFile(zip_path, "w") as zf:
127
  for f in all_files:
128
  zf.write(f, arcname=Path(f).name)
129
 
130
- return rows, gallery, str(zip_path)
131
 
132
 
133
  def _busy_wrapper(*args):
@@ -146,22 +255,35 @@ def _busy_wrapper(*args):
146
  with gr.Blocks(title="Togyzkumalak Scoresheet Reader") as demo:
147
  gr.Markdown(
148
  "# Togyzkumalak Scoresheet Reader\n"
149
- "Upload up to **5** scoresheet photos, optionally tag each game's known "
150
- "result (improves accuracy), then **Convert** to download the PGNs.\n\n"
 
151
  "Outputs per game: `beam` (best legal reconstruction), `raw` (pure OCR), "
152
  "`legal` (strict replay). Free demo — a first run may wake the Space, and "
153
  "images are processed one at a time."
154
  )
155
- round_label = gr.Textbox(label="Round (optional)", placeholder="e.g. 3",
156
- scale=1, max_lines=1)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
157
 
158
- slots = []
159
- for i in range(MAX_IMAGES):
160
- with gr.Row():
161
- img = gr.Image(label=f"Table {i + 1}", type="filepath", height=150)
162
- res = gr.Dropdown(RESULT_CHOICES, value="unknown",
163
- label="Result (1-0 = White won)", scale=1)
164
- slots.extend([img, res])
165
 
166
  convert_btn = gr.Button("Convert", variant="primary")
167
 
@@ -169,13 +291,14 @@ with gr.Blocks(title="Togyzkumalak Scoresheet Reader") as demo:
169
  headers=["game", "legal plies", "stopped", "beam PGN"],
170
  label="Results", wrap=True, interactive=False,
171
  )
 
172
  gallery = gr.Gallery(label="Annotated reconstruction", columns=2, height="auto")
173
  zip_out = gr.File(label="Download all PGNs (zip)")
174
 
175
  convert_btn.click(
176
  _busy_wrapper,
177
- inputs=[round_label, *slots],
178
- outputs=[results_table, gallery, zip_out],
179
  )
180
 
181
  # Serialize CPU-heavy runs: callers wait in a bounded queue instead of
 
1
  """Gradio demo for the Togyzkumalak scoresheet reader (HuggingFace Space).
2
 
3
+ Upload up to 5 scoresheet photos of one round as a single batch, describe the
4
+ tournament and the games in two CSV text fields, and download the
5
+ reconstructed PGNs (with proper PGN tags). Inference runs on the exported
6
+ ONNX models (torch-free) via `togyz.pipeline.run_pipeline`.
7
 
8
  This is a demo, not production: state is per-session, work is capped at 5
9
  images per run, and requests are serialized through Gradio's queue so a shared
 
27
  print(f"[app] Failed to apply monkey-patch: {e}", flush=True)
28
  # --------------------------------------
29
 
30
+ import csv
31
+ import io
32
  import os
33
  import tempfile
34
  import zipfile
 
55
 
56
  MAX_IMAGES = 5
57
  MODEL_DIR = Path(__file__).parent / "models"
58
+ BEAM_CHOICES = [str(2**k) for k in range(10, 31)] # 1024 ... 2^30
59
+ DEFAULT_BEAM = "8192"
60
+
61
+ # accepted spellings of a game result -> pipeline result code
62
+ RESULT_MAP = {
63
+ "": None, "*": None,
64
+ "1": "1-0", "1-0": "1-0",
65
+ "0": "0-1", "0-1": "0-1",
66
+ "5": "draw", "0.5-0.5": "draw", "1/2-1/2": "draw", "draw": "draw",
67
+ }
68
+ RESULT_TAGS = {"1-0": "1-0", "0-1": "0-1", "draw": "1/2-1/2", None: "*"}
69
 
70
  # Load the ONNX sessions once at import - warm for the whole process lifetime.
71
  _log(f"loading move model from {MODEL_DIR / 'best.onnx'} ...")
72
  _MOVES = load_classifier(MODEL_DIR / "best.onnx")
73
+ _DIAGRAM = None
74
+ _diagram_path = MODEL_DIR / "diagram.onnx"
75
+ if _diagram_path.exists():
76
+ _log("loading diagram model ...")
77
+ _DIAGRAM = load_classifier(_diagram_path)
78
+ else:
79
+ # the old kazan.onnx has incompatible classes - do not fall back to it
80
+ _log("no models/diagram.onnx - checkpoint evidence disabled")
81
  _log("models loaded")
82
 
83
 
 
86
  return keep.strip("_")
87
 
88
 
89
+ def _csv_fields(line: str) -> list[str]:
90
+ """One CSV line -> stripped fields (handles quoted commas)."""
91
+ rows = list(csv.reader(io.StringIO(line)))
92
+ return [f.strip() for f in rows[0]] if rows else []
93
+
94
+
95
+ def _parse_meta(text: str) -> dict:
96
+ """Shared metadata line: Tournament,Location,Date,Round (all optional)."""
97
+ fields = _csv_fields((text or "").strip())
98
+ if len(fields) > 4:
99
+ raise gr.Error(
100
+ "Metadata must be one CSV line: Tournament,Location,Date,Round "
101
+ f"(got {len(fields)} fields). Quote fields that contain commas."
102
+ )
103
+ fields += [""] * (4 - len(fields))
104
+ meta = {"event": fields[0], "site": fields[1], "date": fields[2],
105
+ "round": fields[3]}
106
+ if meta["round"]:
107
+ try:
108
+ rnd = int(meta["round"])
109
+ except ValueError:
110
+ raise gr.Error(f"Round must be a number 1-20, got {meta['round']!r}.")
111
+ if not 1 <= rnd <= 20:
112
+ raise gr.Error(f"Round must be between 1 and 20, got {rnd}.")
113
+ meta["round"] = str(rnd)
114
+ return meta
115
+
116
+
117
+ def _parse_games(text: str, n_images: int) -> list[dict]:
118
+ """Per-game lines: WhiteName,BlackName,Result,WhiteTime,BlackTime.
119
+
120
+ One line per uploaded image, in order. Fewer lines than images is fine
121
+ (missing games get empty metadata); more lines is an error.
122
+ """
123
+ lines = [ln for ln in (text or "").splitlines() if ln.strip()]
124
+ if len(lines) > n_images:
125
+ raise gr.Error(
126
+ f"{len(lines)} game lines for {n_images} image(s). "
127
+ "Provide at most one line per uploaded image, in order."
128
+ )
129
+ games = []
130
+ for lineno, line in enumerate(lines, start=1):
131
+ fields = _csv_fields(line)
132
+ if len(fields) > 5:
133
+ raise gr.Error(
134
+ f"Game line {lineno}: expected at most 5 CSV fields "
135
+ "(White,Black,Result,WhiteTime,BlackTime), got "
136
+ f"{len(fields)}. Quote fields that contain commas."
137
+ )
138
+ fields += [""] * (5 - len(fields))
139
+ raw_result = fields[2]
140
+ if raw_result not in RESULT_MAP:
141
+ raise gr.Error(
142
+ f"Game line {lineno}: unknown result {raw_result!r}. Accepted: "
143
+ "1 or 1-0 (White won), 0 or 0-1 (Black won), "
144
+ "5 / 0.5-0.5 / 1/2-1/2 (draw), or empty."
145
+ )
146
+ games.append({"white": fields[0], "black": fields[1],
147
+ "result": RESULT_MAP[raw_result],
148
+ "white_time": fields[3], "black_time": fields[4]})
149
+ games += [{"white": "", "black": "", "result": None,
150
+ "white_time": "", "black_time": ""}] * (n_images - len(games))
151
+ return games
152
+
153
+
154
+ def _pgn_tags(meta: dict, game: dict) -> str:
155
+ """Standard PGN tag section from the shared + per-game metadata."""
156
+ tags = [
157
+ ("Event", meta["event"] or "?"),
158
+ ("Site", meta["site"] or "?"),
159
+ ("Date", meta["date"] or "?"),
160
+ ("Round", meta["round"] or "?"),
161
+ ("White", game["white"] or "?"),
162
+ ("Black", game["black"] or "?"),
163
+ ("Result", RESULT_TAGS[game["result"]]),
164
+ ]
165
+ if game["white_time"]:
166
+ tags.append(("WhiteClock", game["white_time"]))
167
+ if game["black_time"]:
168
+ tags.append(("BlackClock", game["black_time"]))
169
+ return "".join(f'[{k} "{v}"]\n' for k, v in tags) + "\n"
170
+
171
+
172
+ def _process_one(image_path, game: dict, meta: dict, beam_width: int,
173
+ base_name: str, out_dir: Path):
174
+ """Run the pipeline on one image; return (row, gallery_item, files, warnings)."""
175
+ out = run_pipeline(image_path, _MOVES, _DIAGRAM,
176
+ result=game["result"], beam_width=beam_width)
177
 
178
  stop = out["stopped"]
179
  stop_txt = stop.get("reason", "")
 
181
  stop_txt += f" ({stop['winner']})"
182
  note = " ⚠ low-res" if out["low_resolution"] else ""
183
 
184
+ tags = _pgn_tags(meta, game)
185
  files = []
186
  for kind in ("beam", "raw", "legal"):
187
  f = out_dir / f"{base_name}_{kind}.pgn"
188
+ f.write_text(tags + out[f"{kind}_pgn"])
189
  files.append(str(f))
190
 
191
  row = [base_name, out["beam_plies"], stop_txt + note, out["beam_pgn"].strip()]
192
  caption = f"{base_name}: {out['beam_plies']} plies"
193
+ warnings = [f"{base_name}: {w}" for w in out["warnings"]]
194
+ return row, (out["annotated_image"], caption), files, warnings
195
 
 
 
 
 
 
196
 
197
+ def convert(meta_text, games_text, beam_choice, img1, img2, img3, img4, img5):
198
+ """One batch: up to 5 images sharing tournament metadata."""
199
+ images = [img for img in (img1, img2, img3, img4, img5) if img]
200
+ if not images:
201
  raise gr.Error("Please upload at least one scoresheet image.")
202
+
203
+ # parse everything up front - all input errors surface before any heavy work
204
+ meta = _parse_meta(meta_text)
205
+ games = _parse_games(games_text, len(images))
206
+ try:
207
+ beam_width = int(beam_choice)
208
+ except (TypeError, ValueError):
209
+ beam_width = int(DEFAULT_BEAM)
210
 
211
  out_dir = Path(tempfile.mkdtemp(prefix="togyz_"))
212
+ round_slug = _safe_slug(meta["round"])
213
+ rows, gallery, all_files, all_warnings = [], [], [], []
214
 
215
+ for i, (img, game) in enumerate(zip(images, games), start=1):
216
+ prefix = f"round{round_slug}_game{i}" if round_slug else f"game{i}"
217
  try:
218
+ row, gal, files, warns = _process_one(
219
+ img, game, meta, beam_width, prefix, out_dir
220
+ )
221
  except Exception as exc: # one bad image must not kill the batch
222
  rows.append([prefix, 0, f"error: {exc}", ""])
223
  continue
224
  rows.append(row)
225
  gallery.append(gal)
226
  all_files.extend(files)
227
+ all_warnings.extend(warns)
228
 
229
+ warnings_md = "\n".join(f"⚠ {w}" for w in dict.fromkeys(all_warnings))
230
  if not all_files:
231
  # every image errored - still return the table so the user sees why
232
+ return rows, gallery, None, warnings_md
233
 
234
+ zip_path = out_dir / (f"round{round_slug}_pgns.zip" if round_slug else "pgns.zip")
235
  with zipfile.ZipFile(zip_path, "w") as zf:
236
  for f in all_files:
237
  zf.write(f, arcname=Path(f).name)
238
 
239
+ return rows, gallery, str(zip_path), warnings_md
240
 
241
 
242
  def _busy_wrapper(*args):
 
255
  with gr.Blocks(title="Togyzkumalak Scoresheet Reader") as demo:
256
  gr.Markdown(
257
  "# Togyzkumalak Scoresheet Reader\n"
258
+ "Upload up to **5** scoresheet photos of one round as a single batch, "
259
+ "describe the round and the games in the two text fields, then "
260
+ "**Convert** to download the PGNs.\n\n"
261
  "Outputs per game: `beam` (best legal reconstruction), `raw` (pure OCR), "
262
  "`legal` (strict replay). Free demo — a first run may wake the Space, and "
263
  "images are processed one at a time."
264
  )
265
+ meta_box = gr.Textbox(
266
+ label="Tournament metadata (CSV): Tournament,Location,Date,Round — all optional",
267
+ placeholder="World Championship among boys, Astana city, 2026 7 July, 7",
268
+ max_lines=1,
269
+ )
270
+ games_box = gr.Textbox(
271
+ label="Games (one CSV line per image, in order): "
272
+ "WhiteName,BlackName,Result,WhiteTime,BlackTime — all optional; "
273
+ "Result: 1 / 1-0, 0 / 0-1, 5 / 0.5-0.5 / 1/2-1/2",
274
+ placeholder="Zhanabay Korkem, Kubzhasar Marhabat, 0-1, 0:07:55, 0:24:05",
275
+ lines=MAX_IMAGES,
276
+ )
277
+ beam_dd = gr.Dropdown(
278
+ BEAM_CHOICES, value=DEFAULT_BEAM,
279
+ label="Beam width (game hypotheses kept; larger = slower but more "
280
+ "thorough; very large values are trimmed to fit memory)",
281
+ )
282
 
283
+ image_slots = [
284
+ gr.Image(label=f"Game {i + 1}", type="filepath", height=150)
285
+ for i in range(MAX_IMAGES)
286
+ ]
 
 
 
287
 
288
  convert_btn = gr.Button("Convert", variant="primary")
289
 
 
291
  headers=["game", "legal plies", "stopped", "beam PGN"],
292
  label="Results", wrap=True, interactive=False,
293
  )
294
+ warnings_md = gr.Markdown()
295
  gallery = gr.Gallery(label="Annotated reconstruction", columns=2, height="auto")
296
  zip_out = gr.File(label="Download all PGNs (zip)")
297
 
298
  convert_btn.click(
299
  _busy_wrapper,
300
+ inputs=[meta_box, games_box, beam_dd, *image_slots],
301
+ outputs=[results_table, gallery, zip_out, warnings_md],
302
  )
303
 
304
  # Serialize CPU-heavy runs: callers wait in a bounded queue instead of
notebooks/colab_train.ipynb CHANGED
@@ -4,12 +4,15 @@
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
  {
@@ -18,13 +21,23 @@
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
  {
@@ -33,7 +46,7 @@
33
  "metadata": {},
34
  "outputs": [],
35
  "source": [
36
- "!pip install -q -r requirements.txt"
37
  ]
38
  },
39
  {
@@ -42,8 +55,9 @@
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
  },
@@ -53,10 +67,12 @@
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
  {
@@ -65,7 +81,8 @@
65
  "metadata": {},
66
  "outputs": [],
67
  "source": [
68
- "!python train.py --epochs 30 --samples-per-epoch 50000 --batch-size 256 --workers 2"
 
69
  ]
70
  },
71
  {
@@ -74,11 +91,43 @@
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
  ],
@@ -89,5 +138,5 @@
89
  "language_info": {"name": "python"}
90
  },
91
  "nbformat": 4,
92
- "nbformat_minor": 4
93
  }
 
4
  "cell_type": "markdown",
5
  "metadata": {},
6
  "source": [
7
+ "# Togyzkumalak OCR — Colab training (moves + diagram)\n",
8
  "\n",
9
  "Runtime → Change runtime type → **GPU**, then run the cells in order.\n",
10
  "\n",
11
+ "Trains both classifiers on on-the-fly synthetic data (nothing to upload):\n",
12
+ "1. **moves** — 163-class move-cell reader (`11`..`99x` + `empty`)\n",
13
+ "2. **diagram** — 85-class board-diagram reader (values `0`-`81`, `x`, `-`, `empty`) for the kazan boxes and 2×9 pit grids of the summary strips\n",
14
+ "\n",
15
+ "At the end you download one `artifacts.zip` with the ONNX models + class sidecars, ready to drop into `models/` of the HuggingFace Space."
16
  ]
17
  },
18
  {
 
21
  "metadata": {},
22
  "outputs": [],
23
  "source": [
24
+ "!nvidia-smi"
25
+ ]
26
+ },
27
+ {
28
+ "cell_type": "code",
29
+ "execution_count": null,
30
+ "metadata": {},
31
+ "outputs": [],
32
+ "source": [
33
+ "# Fetch the code at runtime.\n",
34
+ "# Option A (default): clone from GitHub\n",
35
+ "!git clone https://github.com/ansarzeinulla/9OCR.git project\n",
36
+ "%cd project\n",
37
  "\n",
38
+ "# Option B: upload a zip named project.zip instead, then:\n",
39
+ "# from google.colab import files; files.upload()\n",
40
+ "# !unzip -q project.zip -d project && cd project"
41
  ]
42
  },
43
  {
 
46
  "metadata": {},
47
  "outputs": [],
48
  "source": [
49
+ "!pip install -q -r requirements-train.txt"
50
  ]
51
  },
52
  {
 
55
  "metadata": {},
56
  "outputs": [],
57
  "source": [
58
+ "# Build glyph pools at runtime: ARDIS (European handwriting) + EMNIST via\n",
59
+ "# torchvision (--emnist is needed on Colab because ingredients/ is gitignored;\n",
60
+ "# digit 0 comes from these pools).\n",
61
  "!python scripts/get_glyphs.py --emnist"
62
  ]
63
  },
 
67
  "metadata": {},
68
  "outputs": [],
69
  "source": [
70
+ "# Sanity check both synthesizers before burning GPU time.\n",
71
+ "!python -m togyz.synth --preview preview_moves.png --task moves\n",
72
+ "!python -m togyz.synth --preview preview_diagram.png --task diagram\n",
73
+ "from IPython.display import Image as IPImage, display\n",
74
+ "display(IPImage('preview_moves.png'))\n",
75
+ "display(IPImage('preview_diagram.png'))"
76
  ]
77
  },
78
  {
 
81
  "metadata": {},
82
  "outputs": [],
83
  "source": [
84
+ "# Train the move classifier (~30 epochs x 50k synthetic cells).\n",
85
+ "!python train.py --task moves --epochs 30 --samples-per-epoch 50000 --batch-size 256 --workers 2"
86
  ]
87
  },
88
  {
 
91
  "metadata": {},
92
  "outputs": [],
93
  "source": [
94
+ "# Train the unified board-diagram classifier.\n",
95
+ "!python train.py --task diagram --epochs 30 --samples-per-epoch 50000 --batch-size 256 --workers 2"
96
+ ]
97
+ },
98
+ {
99
+ "cell_type": "code",
100
+ "execution_count": null,
101
+ "metadata": {},
102
+ "outputs": [],
103
+ "source": [
104
+ "# Export both checkpoints to single-file ONNX (+ .classes.json sidecars).\n",
105
+ "!python scripts/export_onnx.py"
106
+ ]
107
+ },
108
+ {
109
+ "cell_type": "code",
110
+ "execution_count": null,
111
+ "metadata": {},
112
+ "outputs": [],
113
+ "source": [
114
+ "# Bundle the serving artifacts and download.\n",
115
+ "# In the Space's models/ dir: best.onnx stays best.onnx, the diagram model\n",
116
+ "# is renamed to diagram.onnx (+ diagram.classes.json).\n",
117
+ "import shutil, zipfile\n",
118
+ "from pathlib import Path\n",
119
+ "\n",
120
+ "out = Path('artifacts'); out.mkdir(exist_ok=True)\n",
121
+ "shutil.copy('checkpoints/best.onnx', out / 'best.onnx')\n",
122
+ "shutil.copy('checkpoints/best.classes.json', out / 'best.classes.json')\n",
123
+ "shutil.copy('checkpoints/diagram/best.onnx', out / 'diagram.onnx')\n",
124
+ "shutil.copy('checkpoints/diagram/best.classes.json', out / 'diagram.classes.json')\n",
125
+ "with zipfile.ZipFile('artifacts.zip', 'w') as zf:\n",
126
+ " for f in out.iterdir():\n",
127
+ " zf.write(f, arcname=f.name)\n",
128
  "\n",
 
129
  "from google.colab import files\n",
130
+ "files.download('artifacts.zip')"
131
  ]
132
  }
133
  ],
 
138
  "language_info": {"name": "python"}
139
  },
140
  "nbformat": 4,
141
+ "nbformat_minor": 5
142
  }
read_game.py CHANGED
@@ -30,9 +30,10 @@ def main() -> None:
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,
@@ -52,15 +53,15 @@ def main() -> None:
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,
@@ -71,9 +72,13 @@ def main() -> None:
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))
 
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("--diagram-onnx", default="checkpoints/diagram/best.onnx",
34
+ help="board-diagram ONNX (kazan boxes + pit cells); "
35
+ "checkpoint matching is skipped if the file "
36
+ "does not exist")
37
  parser.add_argument("--out", default=None, help="output dir (default: out/<image stem>)")
38
  parser.add_argument("--topk", type=int, default=5)
39
  parser.add_argument("--beam-width", type=int, default=1024,
 
53
  out_dir.mkdir(parents=True, exist_ok=True)
54
 
55
  moves_clf = load_classifier(args.onnx)
56
+ diagram_clf = None
57
+ if Path(args.diagram_onnx).exists():
58
+ diagram_clf = load_classifier(args.diagram_onnx)
59
  else:
60
+ print(f"No diagram classifier at {args.diagram_onnx} - checkpoint matching off.")
61
 
62
  print(f"Reading {args.image} ...")
63
  out = run_pipeline(
64
+ args.image, moves_clf, diagram_clf,
65
  result=args.result, topk=args.topk,
66
  beam_width=args.beam_width, per_ply=args.beam_top,
67
  temperature=args.temperature, tta=not args.no_tta,
 
72
  print(f"WARNING: median cell height is only {out['median_cell_height']}px - "
73
  "accuracy suffers at this resolution; re-photograph at full camera "
74
  "resolution if possible.")
75
+ if out["checkpoint_report"]:
76
+ kaz = [(r["move"], r["side"], r["read"])
77
+ for r in out["checkpoint_report"] if r["kind"] == "kazan"]
78
+ pits = sum(1 for r in out["checkpoint_report"] if r["kind"] == "pit")
79
+ print(f"Diagram checkpoints read: kazans {kaz}, {pits} pit cells")
80
+ for warning in out["warnings"]:
81
+ print(f"WARNING: {warning}")
82
 
83
  result = {"image": args.image, "onnx": args.onnx, **out["game_json"]}
84
  (out_dir / "game.json").write_text(json.dumps(result, indent=1))
scripts/export_onnx.py CHANGED
@@ -4,7 +4,7 @@
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
@@ -66,8 +66,8 @@ def main() -> None:
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:
 
4
 
5
  Writes, next to each source checkpoint:
6
  checkpoints/best.onnx + checkpoints/best.classes.json
7
+ checkpoints/diagram/best.onnx + checkpoints/diagram/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
 
66
  parser.add_argument(
67
  "checkpoints",
68
  nargs="*",
69
+ default=["checkpoints/best.pt", "checkpoints/diagram/best.pt"],
70
+ help="checkpoint(s) to export (default: moves + diagram best.pt)",
71
  )
72
  args = parser.parse_args()
73
  for path in args.checkpoints:
togyz/classes.py CHANGED
@@ -18,6 +18,10 @@ 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]
 
 
 
 
 
18
  CLASS_TO_IDX = {name: i for i, name in enumerate(CLASSES)}
19
  NUM_CLASSES = len(CLASSES) # 163
20
 
21
+ # Unified board-diagram classes for the summary strips: pit/kazan values
22
+ # 0-81 (rendered as 1 or 2 digits), 'x' = tuzdyk mark, '-' = zero mark.
23
+ # At inference the pipeline filters per context: kazan boxes allow only
24
+ # 10-81, pit cells allow '-', 'x' and 0-81.
25
+ DASH = "-"
26
+ TUZDYK = "x"
27
+ DIAGRAM_CLASSES = [str(v) for v in range(82)] + [TUZDYK, DASH, EMPTY]
togyz/dataset.py CHANGED
@@ -10,7 +10,7 @@ from torch.utils.data import Dataset
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):
@@ -22,12 +22,13 @@ class SyntheticCellDataset(Dataset):
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
@@ -38,7 +39,7 @@ class SyntheticCellDataset(Dataset):
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
 
 
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, synthesize_diagram_cell
14
 
15
 
16
  class SyntheticCellDataset(Dataset):
 
22
  """
23
 
24
  def __init__(self, num_samples: int, seed: int | None = None, project_root=".",
25
+ classes: list[str] | None = None, task: str = "moves"):
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
+ self.synth_fn = synthesize_diagram_cell if task == "diagram" else synthesize_cell
32
 
33
  def __len__(self) -> int:
34
  return self.num_samples
 
39
  else:
40
  rng = random.Random(random.getrandbits(64))
41
  class_name = self.classes[rng.randrange(len(self.classes))]
42
+ img = self.synth_fn(class_name, self.sampler, rng)
43
  return torch.from_numpy(preprocess_pil(img)), self.class_to_idx[class_name]
44
 
45
 
togyz/pipeline.py CHANGED
@@ -2,7 +2,7 @@
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
 
@@ -13,6 +13,7 @@ git history (sheet 1 = 36/39 exact plies). Keep it numerically equivalent.
13
 
14
  from __future__ import annotations
15
 
 
16
  import json
17
  import math
18
  from dataclasses import dataclass, field
@@ -22,10 +23,10 @@ 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
@@ -36,8 +37,16 @@ TTA_SHIFTS = [(0, 0), (-2, 0), (2, 0), (0, -2), (0, 2), (-2, -2), (2, 2)]
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
@@ -152,27 +161,66 @@ def _result_consistency(game: Game, result: str | None) -> int:
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)
@@ -188,12 +236,13 @@ def beam_decode(ply_probs, classes, beam_width=1024, per_ply=9, result=None,
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]))
@@ -212,8 +261,8 @@ def format_pgn(plies: list[str]) -> str:
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
 
@@ -222,13 +271,20 @@ def run_pipeline(image, moves_clf: Classifier, kazan_clf: Classifier | None = No
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
 
@@ -294,31 +350,62 @@ def run_pipeline(image, moves_clf: Classifier, kazan_clf: Classifier | None = No
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 = [
@@ -329,14 +416,19 @@ def run_pipeline(image, moves_clf: Classifier, kazan_clf: Classifier | None = No
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 {
@@ -348,7 +440,8 @@ def run_pipeline(image, moves_clf: Classifier, kazan_clf: Classifier | None = No
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,
 
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
+ board-diagram 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
 
 
13
 
14
  from __future__ import annotations
15
 
16
+ import heapq
17
  import json
18
  import math
19
  from dataclasses import dataclass, field
 
23
  import onnxruntime as ort
24
  from PIL import Image
25
 
26
+ from .classes import DASH, EMPTY, TUZDYK
27
  from .preprocess import preprocess_pil
28
  from .rules import Game
29
+ from .sheet import clean_cell, extract_cells, extract_diagram_cells, render_overlay
30
 
31
  MIN_CELL_HEIGHT = 35 # px; below this, resolution is low (surfaced as a warning)
32
  MIN_INK_RATIO = 0.02 # cells with less handwriting ink than this are empty
 
37
  EXACT_WEIGHT = 0.5 # weight of the exact class-string probability
38
  FACTOR_WEIGHT = 0.5 # weight of the factorized digit-marginal probability
39
  KAZAN_FLOOR = 1e-4 # a misread checkpoint must not single-handedly kill truth
40
+ PIT_FLOOR = 1e-3 # pit reads are noisier than kazans: weaker veto power
41
+ PIT_WEIGHT = 0.5 # 18 pit terms per checkpoint must not swamp the rest
42
  RESULT_CODES = {"1-0": 0, "0-1": 1, "draw": -1}
43
 
44
+ # Memory guard for user-selectable beam widths (the UI offers up to 2^30):
45
+ # a late-game hypothesis (copied move lists etc.) costs roughly this many
46
+ # bytes, and the expanded pool peaks at beam_width * per_ply hypotheses.
47
+ HYPOTHESIS_BYTES = 25_000
48
+ MAX_POOL_BYTES = 2 * 2**30 # ~2 GB hypothesis budget
49
+
50
 
51
  # --------------------------------------------------------------------------- #
52
  # ONNX session loading
 
161
  return 1 if leader == want else 0
162
 
163
 
164
+ def _filter_probs(probs: np.ndarray, classes: list[str], allowed: set[str]):
165
+ """Zero out contextually impossible classes and renormalize.
166
+
167
+ Returns None when nothing observable remains (e.g. a blank crop whose
168
+ mass sat entirely on 'empty') - the cell is then treated as unobserved.
169
+ """
170
+ filtered = np.array(
171
+ [p if c in allowed else 0.0 for c, p in zip(classes, probs)], np.float64
172
+ )
173
+ total = float(filtered.sum())
174
+ if total < 1e-6:
175
+ return None
176
+ return filtered / total
177
+
178
+
179
+ def _diagram_logp(game: Game, cp) -> float:
180
+ """Log-likelihood of a hypothesis' board state under one checkpoint's
181
+ diagram reads (kazan boxes + observed pit cells)."""
182
  logp = 0.0
183
  for side, player in (("W", 0), ("B", 1)):
184
+ probs = cp["kazan"].get(side)
185
  if probs is None:
186
  continue
187
  value = game.kazans[player]
188
  p = float(probs[value]) if value < len(probs) else 0.0
189
  logp += math.log(max(p, KAZAN_FLOOR))
190
+
191
+ for (side, index), (values, p_x, p_dash) in cp["pits"].items():
192
+ player = 0 if side == "W" else 1
193
+ pit_pos = player * 9 + (index - 1)
194
+ # a pit owned as tuzdyk by the opponent is written as 'x'
195
+ if game.tuzdyks[1 - player] == index - 1:
196
+ p = p_x
197
+ else:
198
+ count = game.pits[pit_pos]
199
+ if count == 0:
200
+ p = p_dash + float(values[0])
201
+ else:
202
+ p = float(values[count]) if count < len(values) else 0.0
203
+ logp += PIT_WEIGHT * math.log(max(p, PIT_FLOOR))
204
  return logp
205
 
206
 
207
+ def safe_beam_width(beam_width: int, per_ply: int = 9) -> int:
208
+ """Largest beam width that keeps the expanded pool inside the RAM budget."""
209
+ return max(1, min(beam_width, MAX_POOL_BYTES // (HYPOTHESIS_BYTES * max(per_ply, 1))))
210
+
211
+
212
  def beam_decode(ply_probs, classes, beam_width=1024, per_ply=9, result=None,
213
  checkpoints=None):
214
  """Longest fully-legal move sequence with the highest joint probability,
215
  exploring all `per_ply` legal continuations per ply; the final pool is
216
+ re-ranked by diagram-checkpoint and known-result consistency.
217
+
218
+ `beam_width` is capped by `safe_beam_width` so user-requested widths up
219
+ to 2^30 degrade into a narrower beam instead of exhausting RAM.
220
 
221
  Returns (annotated_moves, log_prob, per_ply_share_list).
222
  """
223
+ beam_width = safe_beam_width(beam_width, per_ply)
224
  checkpoints = checkpoints or {}
225
  class_idx = {c: i for i, c in enumerate(classes)}
226
  active = [(Game(), [], [], 0.0)] # (game, annotated_moves, chosen_shares, logp)
 
236
  annotated = played.notation + ("+" if played.captured else "")
237
  new_logp = logp + math.log(max(raw, 1e-12))
238
  if len(moves) + 1 in checkpoints:
239
+ new_logp += _diagram_logp(nxt, checkpoints[len(moves) + 1])
240
  hyp = (nxt, moves + [annotated], chosen + [share], new_logp)
241
  (finished if nxt.is_over else expanded).append(hyp)
242
  if not expanded:
243
  break
244
+ # nlargest instead of a full sort: wide beams stay tractable
245
+ active = heapq.nlargest(beam_width, expanded, key=lambda h: h[3])
246
 
247
  pool = active + finished
248
  best = max(pool, key=lambda h: (len(h[1]), _result_consistency(h[0], result), h[3]))
 
261
  # --------------------------------------------------------------------------- #
262
  # The single entry point
263
  # --------------------------------------------------------------------------- #
264
+ def run_pipeline(image, moves_clf: Classifier, diagram_clf: Classifier | None = None,
265
+ result=None, topk=5, beam_width=8192, per_ply=9,
266
  temperature=1.0, tta=True, save_cells_dir=None):
267
  """Read one scoresheet image into game records.
268
 
 
271
  (the CLI uses it, the Gradio app does not).
272
 
273
  Keys: beam_pgn, raw_pgn, legal_pgn, game_json (dict), stopped,
274
+ plies_scanned, legal_plies, beam_plies, checkpoint_report, warnings,
275
+ annotated_image (PIL), median_cell_height, low_resolution (bool).
276
  """
277
  classes = moves_clf.classes
278
  empty_idx = classes.index(EMPTY)
279
+ warnings: list[str] = []
280
+ effective_width = safe_beam_width(beam_width, per_ply)
281
+ if effective_width < beam_width:
282
+ warnings.append(
283
+ f"beam width {beam_width} trimmed to {effective_width} to stay "
284
+ "inside the memory budget"
285
+ )
286
 
287
+ cells, sheet, gridlines = extract_cells(image, with_gridlines=True)
288
  cells.sort(key=lambda c: (c.move_no, c.side != "W")) # game order: 1W 1B 2W ...
289
  median_h = sorted(c.bbox[3] for c in cells)[len(cells) // 2]
290
 
 
350
  in_sync = False
351
  raw_pgn.append(raw)
352
 
353
+ # board-diagram evidence from the summary strips (every 10 moves):
354
+ # kazan boxes (values 10-81) plus the 2x9 pit grid ('x'/'-'/0-81)
355
  checkpoints, checkpoint_report = {}, []
356
+ diagram_cells, diagram_labels = [], {}
357
+ if diagram_clf is not None:
358
+ dg_classes = diagram_clf.classes
359
+ value_names = [str(v) for v in range(82)]
360
+ kazan_allowed = set(str(v) for v in range(10, 82))
361
+ pit_allowed = set(value_names) | {TUZDYK, DASH}
362
+ x_idx = dg_classes.index(TUZDYK)
363
+ dash_idx = dg_classes.index(DASH)
364
+
365
+ dg_cells = [(c, *clean_cell(c.image)) for c in extract_diagram_cells(sheet)]
366
+ dg_cells = [
367
+ (c, img) for c, img, ink in dg_cells
368
+ if ink >= (MIN_KAZAN_INK if c.kind == "kazan" else MIN_INK_RATIO)
369
+ ]
370
+ if dg_cells:
371
+ dg_probs = classify_cells(
372
+ [img for _, img in dg_cells], diagram_clf,
373
  temperature=temperature, tta=tta,
374
  )
375
+ for (cell, img), probs in zip(dg_cells, dg_probs):
376
+ if dg_classes[int(np.argmax(probs))] == EMPTY:
377
+ continue
378
+ allowed = kazan_allowed if cell.kind == "kazan" else pit_allowed
379
+ filtered = _filter_probs(probs, dg_classes, allowed)
380
+ if filtered is None:
381
  continue
382
+ cp = checkpoints.setdefault(
383
+ cell.move_no * 2, {"kazan": {}, "pits": {}}
 
 
 
 
384
  )
385
+ values = filtered[:82]
386
+ top_idx = int(np.argmax(filtered))
387
+ read = dg_classes[top_idx]
388
+ if cell.kind == "kazan":
389
+ # renormalized over 10-81 only, value-indexed for scoring
390
+ cp["kazan"][cell.side] = values / max(float(values.sum()), 1e-9)
391
+ else:
392
+ cp["pits"][(cell.side, cell.pit_index)] = (
393
+ values, float(filtered[x_idx]), float(filtered[dash_idx])
394
+ )
395
+ diagram_cells.append(cell)
396
+ diagram_labels[(cell.move_no, cell.kind, cell.side, cell.pit_index)] = read
397
+ entry = {"move": cell.move_no, "kind": cell.kind, "side": cell.side,
398
+ "read": read, "prob": round(float(filtered[top_idx]), 4)}
399
+ if cell.kind == "pit":
400
+ entry["index"] = cell.pit_index
401
+ checkpoint_report.append(entry)
402
  if cells_dir:
403
+ suffix = f"_{cell.pit_index}" if cell.kind == "pit" else ""
404
+ img.save(cells_dir /
405
+ f"{cell.kind}_{cell.move_no:02d}_{cell.side}{suffix}.png")
406
 
407
  beam_moves, beam_logp, beam_shares = beam_decode(
408
+ ply_prob_arrays, classes, effective_width, per_ply,
409
  result=result, checkpoints=checkpoints,
410
  )
411
  beam_detail = [
 
416
  for p, m in zip(plies, beam_moves):
417
  labels[(p["move"], p["side"])] = m
418
 
419
+ annotated = render_overlay(
420
+ sheet, cells[: len(plies)], labels,
421
+ diagram_cells=diagram_cells, diagram_labels=diagram_labels,
422
+ gridlines=gridlines,
423
+ )
424
 
425
  game_json = {
426
  "plies_scanned": len(plies),
427
  "legal_plies": len(legal_pgn),
428
  "beam": {"plies": len(beam_moves), "log_prob": round(beam_logp, 3), "moves": beam_detail},
429
+ "checkpoints": checkpoint_report,
430
  "stopped": stopped or {"reason": "sheet exhausted"},
431
+ "warnings": warnings,
432
  "plies": plies,
433
  }
434
  return {
 
440
  "plies_scanned": len(plies),
441
  "legal_plies": len(legal_pgn),
442
  "beam_plies": len(beam_moves),
443
+ "checkpoint_report": checkpoint_report,
444
+ "warnings": warnings,
445
  "annotated_image": annotated,
446
  "median_cell_height": int(median_h),
447
  "low_resolution": median_h < MIN_CELL_HEIGHT,
togyz/sheet.py CHANGED
@@ -26,11 +26,13 @@ CELL_MARGIN = 0.15 # expand crops; handwriting overflows the printed cells
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:
@@ -82,6 +84,12 @@ def _find_table_quads(gray: np.ndarray) -> list[np.ndarray]:
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())
@@ -139,12 +147,26 @@ def _row_lines(warped: np.ndarray) -> list[int]:
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")
@@ -154,11 +176,17 @@ def extract_cells(image) -> tuple[list[Cell], Image.Image]:
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")):
@@ -180,37 +208,132 @@ def extract_cells(image) -> tuple[list[Cell], Image.Image]:
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:
@@ -245,6 +368,7 @@ def extract_kazan_cells(sheet: Image.Image) -> list[Cell]:
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
 
@@ -281,9 +405,19 @@ def clean_cell(image: Image.Image) -> tuple[Image.Image, float]:
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:
@@ -295,4 +429,13 @@ def render_overlay(sheet: Image.Image, cells: list[Cell], labels: dict | None =
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))
 
26
 
27
  @dataclass
28
  class Cell:
29
+ move_no: int # move cells: 1-80; diagram cells: checkpoint move (10..80)
30
+ side: str # "W" (Bast.) or "B" (Kost.); for pits this is the row owner
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
+ kind: str = "move" # "move" | "kazan" | "pit"
35
+ pit_index: int = 0 # 1-9 for pit cells (scoresheet numbering), else 0
36
 
37
 
38
  def _grid_mask(gray: np.ndarray) -> np.ndarray:
 
84
  f"Found only {len(quads)} move tables (need {TABLES}). "
85
  "Check photo quality/framing."
86
  )
87
+ # the 8 move tables all have near-identical area; drop outliers like the
88
+ # footer/signature block, which can out-size a genuine table
89
+ median_area = float(np.median([a for a, _ in quads]))
90
+ consistent = [(a, q) for a, q in quads if 0.5 * median_area < a < 2.0 * median_area]
91
+ if len(consistent) >= TABLES:
92
+ quads = consistent
93
  quads = [q for _, q in sorted(quads, key=lambda t: -t[0])[:TABLES]]
94
  # split into top/bottom bands by y, order each band left to right
95
  quads.sort(key=lambda q: q[:, 1].mean())
 
147
  return rows
148
 
149
 
150
+ def _map_segment(inverse: np.ndarray, p0, p1) -> tuple[tuple[int, int], tuple[int, int]]:
151
+ """Map a rectified-table segment back onto the original sheet."""
152
+ pts = np.array([p0, p1], np.float32).reshape(-1, 1, 2)
153
+ mapped = cv2.perspectiveTransform(pts, inverse).reshape(2, 2)
154
+ return (
155
+ (int(mapped[0][0]), int(mapped[0][1])),
156
+ (int(mapped[1][0]), int(mapped[1][1])),
157
+ )
158
+
159
+
160
+ def extract_cells(image, with_gridlines: bool = False):
161
  """All 160 move cells (80 moves x W/B) plus the upright sheet image.
162
 
163
  `image` is a path (str/Path) or a PIL.Image. Each table is
164
  perspective-rectified before being split, so cell crops stay aligned even
165
  on tilted phone photos.
166
+
167
+ Returns (cells, sheet) or, with `with_gridlines`, (cells, sheet, gridlines)
168
+ where gridlines is a list of ((x0, y0), (x1, y1)) sheet-coordinate segments
169
+ for every row/column separator used during segmentation.
170
  """
171
  if isinstance(image, Image.Image):
172
  pil = ImageOps.exif_transpose(image).convert("L")
 
176
  gray = np.asarray(pil)
177
 
178
  cells: list[Cell] = []
179
+ gridlines: list[tuple[tuple[int, int], tuple[int, int]]] = []
180
  for t, quad in enumerate(_find_table_quads(gray)):
181
  warped, inverse = _rectify_table(gray, quad)
182
  th, tw = warped.shape
183
  rows = _row_lines(warped)
184
  cols = [round(f * tw) for f in COLUMN_FRACTIONS]
185
+ if with_gridlines:
186
+ for y in rows:
187
+ gridlines.append(_map_segment(inverse, (0, y), (tw, y)))
188
+ for x in cols:
189
+ gridlines.append(_map_segment(inverse, (x, rows[0]), (x, rows[-1])))
190
  for r in range(ROWS_PER_TABLE):
191
  y0, y1 = rows[r + 1], rows[r + 2] # rows[0..1] is the header
192
  for col_index, side in ((1, "W"), (2, "B")):
 
208
  cells.append(
209
  Cell(t * ROWS_PER_TABLE + r + 1, side, bbox, crop, sheet_quad)
210
  )
211
+ if with_gridlines:
212
+ return cells, pil, gridlines
213
  return cells, pil
214
 
215
 
216
+ PIT_MARGIN_Y = 0.25 # pit digits regularly overflow the printed row height
217
+ PIT_MARGIN_X = 0.15 # and bleed a little into neighboring columns
218
 
219
+
220
+ def _strip_pit_grid(mask: np.ndarray, gray: np.ndarray,
221
+ strip_bbox: tuple[int, int, int, int], checkpoint: int) -> list[Cell]:
222
+ """The 18 pit cells of the 2x9 board diagram inside one summary strip.
223
+
224
+ Upper row = Kost. (Black), pit indexes 9..1 left-to-right; lower row =
225
+ Bast. (White), indexes 1..9. Grid lines are detected inside the strip;
226
+ when detection is incomplete the known uniform 2x9 structure is used.
227
+ """
228
+ x, y, bw, bh = strip_bbox
229
+ h, w = gray.shape
230
+ sub = mask[y : y + bh, x : x + bw]
231
+
232
+ # horizontal separators: top / middle / bottom of the 2x9 grid
233
+ hlines = _detect_lines(sub, axis=0, min_frac=0.55)
234
+ if len(hlines) < 2:
235
+ return []
236
+ top, bottom = hlines[0], hlines[-1]
237
+ if bottom - top < 6:
238
+ return []
239
+ mid_target = (top + bottom) / 2
240
+ inner = [v for v in hlines if top + 2 < v < bottom - 2]
241
+ middle = min(inner, key=lambda v: abs(v - mid_target)) if inner else int(round(mid_target))
242
+
243
+ # vertical separators: 10 column lines across the grid band
244
+ band = sub[top : bottom + 1]
245
+ vlines = _detect_lines(band, axis=1, min_frac=0.5)
246
+ if len(vlines) != 10:
247
+ if len(vlines) >= 2:
248
+ x0, x1 = vlines[0], vlines[-1]
249
+ else:
250
+ xs = np.where(band.sum(axis=0) > 0)[0]
251
+ x0, x1 = (int(xs.min()), int(xs.max())) if len(xs) else (0, bw - 1)
252
+ vlines = [round(x0 + j * (x1 - x0) / 9) for j in range(10)]
253
+
254
+ cells: list[Cell] = []
255
+ rows = ((top, middle, "B"), (middle, bottom, "W"))
256
+ for ry0, ry1, owner in rows:
257
+ my = round((ry1 - ry0) * PIT_MARGIN_Y)
258
+ for j in range(9):
259
+ cx0, cx1 = vlines[j], vlines[j + 1]
260
+ mx = round((cx1 - cx0) * PIT_MARGIN_X)
261
+ gx0 = max(0, x + cx0 - mx)
262
+ gx1 = min(w, x + cx1 + mx)
263
+ gy0 = max(0, y + ry0 - my)
264
+ gy1 = min(h, y + ry1 + my)
265
+ if gx1 - gx0 < 4 or gy1 - gy0 < 4:
266
+ continue
267
+ pit_index = 9 - j if owner == "B" else j + 1
268
+ cells.append(Cell(
269
+ checkpoint, owner, (gx0, gy0, gx1 - gx0, gy1 - gy0),
270
+ Image.fromarray(gray[gy0:gy1, gx0:gx1]),
271
+ kind="pit", pit_index=pit_index,
272
+ ))
273
+ return cells
274
+
275
+
276
+ def extract_diagram_cells(sheet: Image.Image) -> list[Cell]:
277
+ """All board-diagram cells from the 8 summary strips between table bands.
278
+
279
+ After every 10 moves the scorer fills a small board diagram: the box
280
+ protruding ABOVE the strip is the Kost. (Black) kazan, the box BELOW is
281
+ the Bast. (White) kazan, and the 2x9 grid between them holds the pit
282
+ counts ('x' = tuzdyk, '-' = 0, or a 1-2 digit number). Returned as Cell
283
+ objects with move_no = the checkpoint move (10, 20, ... 80), kind
284
+ "kazan"/"pit", side "B"/"W" (row owner for pits) and pit_index 1-9.
285
  """
286
  gray = np.asarray(sheet)
287
  h, w = gray.shape
288
  mask = _grid_mask(gray)
289
+ quads = _find_table_quads(gray)
290
+ table_area = float(np.median([cv2.contourArea(q) for q in quads]))
291
+
292
+ # Diagram strips are located by table position, not by enumeration order:
293
+ # the strip for checkpoint 10*(4*band + col + 1) sits below table `col` of
294
+ # band `band`, in the gap under that band. Handwriting sometimes bridges
295
+ # neighboring strips into one connected component, so wide components are
296
+ # kept and carved by each table's x-range.
297
  n, _, stats, _ = cv2.connectedComponentsWithStats(mask)
298
+ components = [
299
  tuple(int(v) for v in stats[i][:4])
300
  for i in range(1, n)
301
+ if stats[i][2] > w * 0.05
302
  and h * 0.02 < stats[i][3] < h * 0.12
303
  and stats[i][2] > 1.5 * stats[i][3]
304
+ and stats[i][2] * stats[i][3] < 2.5 * table_area
305
  ]
306
+
307
+ top_quads, bottom_quads = quads[:4], quads[4:]
308
+ zones = (
309
+ (top_quads, (max(q[:, 1].max() for q in top_quads),
310
+ min(q[:, 1].min() for q in bottom_quads))),
311
+ (bottom_quads, (max(q[:, 1].max() for q in bottom_quads), h)),
312
+ )
313
+
314
+ strips: list[tuple[int, tuple[int, int, int, int]]] = []
315
+ for band, (band_quads, (zy0, zy1)) in enumerate(zones):
316
+ for col, quad in enumerate(band_quads):
317
+ tx0, tx1 = float(quad[:, 0].min()), float(quad[:, 0].max())
318
+ parts = []
319
+ for cx, cy, cw, ch in components:
320
+ yc = cy + ch / 2
321
+ overlap = min(cx + cw, tx1) - max(cx, tx0)
322
+ if zy0 < yc < zy1 and overlap > 0.3 * (tx1 - tx0):
323
+ parts.append((cx, cy, cw, ch))
324
+ if not parts:
325
+ continue
326
+ pad = int(0.02 * w)
327
+ sx0 = max(int(tx0) - pad, min(cx for cx, *_ in parts))
328
+ sx1 = min(int(tx1) + pad, max(cx + cw for cx, _, cw, _ in parts))
329
+ sy0 = min(cy for _, cy, *_ in parts)
330
+ sy1 = max(cy + ch for _, cy, _, ch in parts)
331
+ checkpoint = (band * 4 + col + 1) * 10
332
+ strips.append((checkpoint, (sx0, sy0, sx1 - sx0, sy1 - sy0)))
333
 
334
  cells: list[Cell] = []
335
+ for checkpoint, (x, y, bw, bh) in strips:
336
+ cells.extend(_strip_pit_grid(mask, gray, (x, y, bw, bh), checkpoint))
337
  sub = mask[y : y + bh, x : x + bw]
338
  long_rows = np.where((sub > 0).sum(axis=1) > bw * 0.55)[0]
339
  if len(long_rows) == 0:
 
368
  cells.append(Cell(
369
  checkpoint, side, (cx0, cy0, cx1 - cx0, cy1 - cy0),
370
  Image.fromarray(gray[cy0:cy1, cx0:cx1]),
371
+ kind="kazan",
372
  ))
373
  return cells
374
 
 
405
  return Image.fromarray(cleaned), ink_ratio
406
 
407
 
408
+ def render_overlay(sheet: Image.Image, cells: list[Cell], labels: dict | None = None,
409
+ diagram_cells: list[Cell] | None = None,
410
+ diagram_labels: dict | None = None,
411
+ gridlines: list | None = None) -> Image.Image:
412
+ """Debug image: move-cell boxes (green) with predicted labels (red),
413
+ kazan boxes (blue), pit cells (orange) with their reads, and the
414
+ segmentation gridlines (gray).
415
+
416
+ `diagram_labels` is keyed by (move_no, kind, side, pit_index).
417
+ """
418
  vis = cv2.cvtColor(np.asarray(sheet), cv2.COLOR_GRAY2BGR)
419
+ for p0, p1 in gridlines or []:
420
+ cv2.line(vis, p0, p1, (160, 160, 160), 1, cv2.LINE_AA)
421
  for cell in cells:
422
  x, y, w, h = cell.bbox
423
  if cell.quad is not None:
 
429
  if text:
430
  cv2.putText(vis, text, (x + 2, y + h - 3),
431
  cv2.FONT_HERSHEY_SIMPLEX, 0.45, (0, 0, 255), 1, cv2.LINE_AA)
432
+ for cell in diagram_cells or []:
433
+ x, y, w, h = cell.bbox
434
+ color = (200, 120, 0) if cell.kind == "kazan" else (0, 140, 255) # BGR
435
+ cv2.rectangle(vis, (x, y), (x + w, y + h), color, 1)
436
+ if diagram_labels:
437
+ text = diagram_labels.get((cell.move_no, cell.kind, cell.side, cell.pit_index))
438
+ if text:
439
+ cv2.putText(vis, text, (x + 1, y + h - 2),
440
+ cv2.FONT_HERSHEY_SIMPLEX, 0.35, (0, 0, 255), 1, cv2.LINE_AA)
441
  return Image.fromarray(cv2.cvtColor(vis, cv2.COLOR_BGR2RGB))
togyz/synth.py CHANGED
@@ -20,7 +20,7 @@ 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
@@ -206,12 +206,116 @@ def synthesize_cell(
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
@@ -221,7 +325,7 @@ def _preview(path: str, seed: int, count: int) -> None:
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)
@@ -234,5 +338,6 @@ if __name__ == "__main__":
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)
 
20
  import numpy as np
21
  from PIL import Image
22
 
23
+ from .classes import CLASSES, DASH, DIAGRAM_CLASSES, EMPTY, TUZDYK
24
  from .glyphs import GlyphSampler
25
 
26
  # working resolution: glyphs are rendered at this digit height, the finished
 
206
  return Image.fromarray(_camera_effects(img, rng))
207
 
208
 
209
+ def _paste_clipped(alpha: np.ndarray, mask: np.ndarray, x: int, y: int) -> None:
210
+ """max-blend `mask` onto `alpha` at (x, y), clipping at the canvas edges
211
+ (diagram digits regularly overflow the printed cell borders)."""
212
+ h, w = alpha.shape
213
+ mh, mw = mask.shape
214
+ x0, y0 = max(0, x), max(0, y)
215
+ x1, y1 = min(w, x + mw), min(h, y + mh)
216
+ if x1 <= x0 or y1 <= y0:
217
+ return
218
+ region = alpha[y0:y1, x0:x1]
219
+ np.maximum(region, mask[y0 - y : y1 - y, x0 - x : x1 - x], out=region)
220
+
221
+
222
+ def _dash_mask(rng: random.Random) -> np.ndarray:
223
+ """A handwritten '-' (zero mark): one short, slightly curved stroke."""
224
+ s = 64
225
+ canvas = np.zeros((s // 2, s), np.float32)
226
+ y = s // 4 + rng.randint(-4, 4)
227
+ x0, x1 = rng.randint(2, 12), rng.randint(52, 62)
228
+ mid = ((x0 + x1) // 2, y + rng.randint(-6, 6))
229
+ pts = np.array([(x0, y), mid, (x1, y + rng.randint(-4, 4))], np.int32)
230
+ cv2.polylines(canvas, [pts], False, 1.0, rng.randint(3, 6), cv2.LINE_AA)
231
+ ys, xs = np.where(canvas > 0.1)
232
+ return canvas[ys.min() : ys.max() + 1, xs.min() : xs.max() + 1]
233
+
234
+
235
+ def synthesize_diagram_cell(
236
+ class_name: str, sampler: GlyphSampler, rng: random.Random
237
+ ) -> Image.Image:
238
+ """Render one board-diagram cell photo: a pit/kazan value 0-81, a tuzdyk
239
+ 'x', a zero '-', or 'empty'.
240
+
241
+ Diagram crops differ from move cells: they are squarer, always framed by
242
+ printed grid lines (drawn on several edges), digits often overflow the
243
+ row height and neighboring digits bleed into the crop sides.
244
+ """
245
+ # --- ink masks for the written symbol ---
246
+ if class_name == EMPTY:
247
+ masks: list[np.ndarray] = []
248
+ elif class_name == DASH:
249
+ m = _dash_mask(rng)
250
+ target_w = max(10, int(RENDER_DIGIT_H * rng.uniform(0.35, 0.8)))
251
+ target_h = max(4, int(m.shape[0] * target_w / m.shape[1]))
252
+ masks = [cv2.resize(m, (target_w, target_h), interpolation=cv2.INTER_LINEAR)]
253
+ else:
254
+ chars = class_name if class_name == TUZDYK else str(int(class_name))
255
+ slant = rng.uniform(-0.10, 0.40)
256
+ masks = []
257
+ for char in chars:
258
+ mask = sampler.sample(char, rng)
259
+ rel_h = rng.uniform(0.55, 0.95) if char == "x" else rng.uniform(0.85, 1.15)
260
+ target_h = max(8, int(RENDER_DIGIT_H * rel_h))
261
+ target_w = max(4, int(mask.shape[1] * target_h / mask.shape[0]))
262
+ mask = cv2.resize(mask, (target_w, target_h), interpolation=cv2.INTER_LINEAR)
263
+ mask = _transform_mask(mask, rng.uniform(-7, 7), slant + rng.uniform(-0.05, 0.05))
264
+ masks.append(_vary_thickness(mask, rng))
265
+
266
+ gaps = [int(rng.uniform(-0.22, 0.10) * masks[i].shape[1]) for i in range(len(masks) - 1)]
267
+ total_w = sum(m.shape[1] for m in masks) + sum(gaps)
268
+ max_h = max((m.shape[0] for m in masks), default=RENDER_DIGIT_H)
269
+
270
+ # --- cell geometry: squarish crop, digits may overflow it vertically ---
271
+ fill = rng.uniform(0.55, 1.25) # >1: taller than the printed row
272
+ cell_h = max(24, int(max_h / fill))
273
+ aspect = rng.uniform(0.9, 1.9)
274
+ cell_w = max(int(cell_h * aspect), int(total_w * rng.uniform(0.85, 1.15)))
275
+
276
+ alpha = np.zeros((cell_h, cell_w), np.float32)
277
+ x = int((cell_w - total_w) * rng.uniform(0.1, 0.9)) if masks else 0
278
+ for i, mask in enumerate(masks):
279
+ mh = mask.shape[0]
280
+ y = int((cell_h - mh) * rng.uniform(0.25, 0.75)) if cell_h > mh else int(
281
+ (cell_h - mh) * rng.uniform(0.3, 0.7)
282
+ )
283
+ _paste_clipped(alpha, mask, x, y)
284
+ if i < len(gaps):
285
+ x += mask.shape[1] + gaps[i]
286
+
287
+ # --- neighbor bleed: fragments of adjacent cells' digits at the sides ---
288
+ if rng.random() < 0.3:
289
+ frag = sampler.sample(rng.choice("0123456789"), rng)
290
+ fh = max(8, int(RENDER_DIGIT_H * rng.uniform(0.7, 1.1)))
291
+ fw = max(4, int(frag.shape[1] * fh / frag.shape[0]))
292
+ frag = cv2.resize(frag, (fw, fh), interpolation=cv2.INTER_LINEAR)
293
+ outside = rng.uniform(0.6, 0.9)
294
+ fx = -int(fw * outside) if rng.random() < 0.5 else cell_w - int(fw * (1 - outside))
295
+ _paste_clipped(alpha, frag, fx, int((cell_h - fh) * rng.uniform(0.2, 0.8)))
296
+
297
+ img = _compose_ink(_paper(rng, cell_h, cell_w), alpha, rng)
298
+ # printed grid borders: diagram crops nearly always contain them
299
+ for _ in range(rng.randint(1, 4) if rng.random() < 0.85 else 0):
300
+ img = _add_border_lines(img, rng)
301
+
302
+ out_h = rng.randint(20, 90) # pit rows are small on the sheet
303
+ out_w = max(12, int(cell_w * out_h / cell_h))
304
+ img = cv2.resize(img, (out_w, out_h), interpolation=cv2.INTER_AREA)
305
+ return Image.fromarray(_camera_effects(img, rng))
306
+
307
+
308
+ def _preview(path: str, seed: int, count: int, task: str = "moves") -> None:
309
  """Render a comparison grid: the real-crop classes first, then random ones."""
310
  sampler = GlyphSampler()
311
  rng = random.Random(seed)
312
+ if task == "diagram":
313
+ fixed = ["x", "-", "0", "7", "12", "45", "81", "empty"]
314
+ pool, synth_fn = DIAGRAM_CLASSES, synthesize_diagram_cell
315
+ else:
316
+ fixed = ["96", "77", "13x", "37", "85", "42"] * 2
317
+ pool, synth_fn = CLASSES, synthesize_cell
318
+ names = fixed + [rng.choice(pool) for _ in range(max(0, count - len(fixed)))]
319
 
320
  tile_w, tile_h, caption = 200, 100, 14
321
  cols = 6
 
325
 
326
  draw = ImageDraw.Draw(sheet)
327
  for i, name in enumerate(names):
328
+ cell = synth_fn(name, sampler, rng).resize((tile_w, tile_h))
329
  cx, cy = (i % cols) * tile_w, (i // cols) * (tile_h + caption)
330
  sheet.paste(cell, (cx, cy + caption))
331
  draw.text((cx + 4, cy + 1), name, fill=0)
 
338
  parser.add_argument("--preview", default="preview.png", help="output image path")
339
  parser.add_argument("--seed", type=int, default=0)
340
  parser.add_argument("--n", type=int, default=48, help="number of cells")
341
+ parser.add_argument("--task", choices=["moves", "diagram"], default="moves")
342
  args = parser.parse_args()
343
+ _preview(args.preview, args.seed, args.n, args.task)
train.py CHANGED
@@ -19,7 +19,7 @@ import torch
19
  import torch.nn as nn
20
  from torch.utils.data import DataLoader
21
 
22
- from togyz.classes import CLASSES, KAZAN_CLASSES
23
  from togyz.dataset import RealCropDataset, SyntheticCellDataset
24
  from togyz.model import auto_device, build_model, save_checkpoint
25
 
@@ -52,9 +52,11 @@ def evaluate_real(model, real: RealCropDataset, device) -> tuple[float, float]:
52
 
53
  def main() -> None:
54
  parser = argparse.ArgumentParser(description=__doc__)
55
- parser.add_argument("--task", choices=["moves", "kazan"], default="moves",
56
- help="moves: 163-class cell classifier; kazan: 00-81 "
57
- "checkpoint numbers from the summary strips")
 
 
58
  parser.add_argument("--epochs", type=int, default=20)
59
  parser.add_argument("--samples-per-epoch", type=int, default=50_000)
60
  parser.add_argument("--val-size", type=int, default=4_000)
@@ -63,7 +65,7 @@ def main() -> None:
63
  parser.add_argument("--weight-decay", type=float, default=1e-4)
64
  parser.add_argument("--workers", type=int, default=4)
65
  parser.add_argument("--out", default=None,
66
- help="default: checkpoints (moves) / checkpoints/kazan")
67
  parser.add_argument("--resume", default=None, help="path to last.pt to continue")
68
  parser.add_argument("--device", default=None, help="cuda / mps / cpu (default: auto)")
69
  args = parser.parse_args()
@@ -71,13 +73,16 @@ def main() -> None:
71
  device = torch.device(args.device) if args.device else auto_device()
72
  print(f"Device: {device}, task: {args.task}")
73
 
74
- classes = KAZAN_CLASSES if args.task == "kazan" else CLASSES
75
- out_dir = Path(args.out or ("checkpoints/kazan" if args.task == "kazan" else "checkpoints"))
76
  out_dir.mkdir(parents=True, exist_ok=True)
77
  (out_dir / "classes.json").write_text(json.dumps(classes))
78
 
79
- train_set = SyntheticCellDataset(args.samples_per_epoch, seed=None, classes=classes)
80
- val_set = SyntheticCellDataset(args.val_size, seed=1234, classes=classes)
 
 
 
81
  real_set = RealCropDataset() if args.task == "moves" else RealCropDataset("/nonexistent")
82
  print(f"Glyph pools:\n{train_set.sampler.describe()}")
83
  print(f"Real eval crops: {len(real_set)}")
 
19
  import torch.nn as nn
20
  from torch.utils.data import DataLoader
21
 
22
+ from togyz.classes import CLASSES, DIAGRAM_CLASSES
23
  from togyz.dataset import RealCropDataset, SyntheticCellDataset
24
  from togyz.model import auto_device, build_model, save_checkpoint
25
 
 
52
 
53
  def main() -> None:
54
  parser = argparse.ArgumentParser(description=__doc__)
55
+ parser.add_argument("--task", choices=["moves", "diagram"], default="moves",
56
+ help="moves: 163-class cell classifier; diagram: unified "
57
+ "board-diagram reader (0-81, 'x', '-') for the kazan "
58
+ "boxes and pit cells of the summary strips "
59
+ "(replaces the old 'kazan' task)")
60
  parser.add_argument("--epochs", type=int, default=20)
61
  parser.add_argument("--samples-per-epoch", type=int, default=50_000)
62
  parser.add_argument("--val-size", type=int, default=4_000)
 
65
  parser.add_argument("--weight-decay", type=float, default=1e-4)
66
  parser.add_argument("--workers", type=int, default=4)
67
  parser.add_argument("--out", default=None,
68
+ help="default: checkpoints (moves) / checkpoints/diagram")
69
  parser.add_argument("--resume", default=None, help="path to last.pt to continue")
70
  parser.add_argument("--device", default=None, help="cuda / mps / cpu (default: auto)")
71
  args = parser.parse_args()
 
73
  device = torch.device(args.device) if args.device else auto_device()
74
  print(f"Device: {device}, task: {args.task}")
75
 
76
+ classes = DIAGRAM_CLASSES if args.task == "diagram" else CLASSES
77
+ out_dir = Path(args.out or ("checkpoints/diagram" if args.task == "diagram" else "checkpoints"))
78
  out_dir.mkdir(parents=True, exist_ok=True)
79
  (out_dir / "classes.json").write_text(json.dumps(classes))
80
 
81
+ train_set = SyntheticCellDataset(args.samples_per_epoch, seed=None,
82
+ classes=classes, task=args.task)
83
+ val_set = SyntheticCellDataset(args.val_size, seed=1234,
84
+ classes=classes, task=args.task)
85
+ # the labeled real crops are move cells; other tasks have no real eval set
86
  real_set = RealCropDataset() if args.task == "moves" else RealCropDataset("/nonexistent")
87
  print(f"Glyph pools:\n{train_set.sampler.describe()}")
88
  print(f"Real eval crops: {len(real_set)}")