Spaces:
Sleeping
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
- README.md +26 -15
- app.py +168 -45
- notebooks/colab_train.ipynb +70 -21
- read_game.py +16 -11
- scripts/export_onnx.py +3 -3
- togyz/classes.py +7 -3
- togyz/dataset.py +4 -3
- togyz/pipeline.py +127 -34
- togyz/sheet.py +163 -20
- togyz/synth.py +111 -6
- train.py +14 -9
|
@@ -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
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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/
|
| 31 |
-
cp checkpoints/
|
| 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
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
| 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
|
| 121 |
-
moves (Black's box above the strip, White's below,
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
|
|
@@ -1,8 +1,9 @@
|
|
| 1 |
"""Gradio demo for the Togyzkumalak scoresheet reader (HuggingFace Space).
|
| 2 |
|
| 3 |
-
Upload up to 5 scoresheet photos
|
| 4 |
-
and
|
| 5 |
-
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 61 |
-
|
| 62 |
-
if
|
| 63 |
-
_log("loading
|
| 64 |
-
|
|
|
|
|
|
|
|
|
|
| 65 |
_log("models loaded")
|
| 66 |
|
| 67 |
|
|
@@ -70,10 +86,94 @@ def _safe_slug(text: str) -> str:
|
|
| 70 |
return keep.strip("_")
|
| 71 |
|
| 72 |
|
| 73 |
-
def
|
| 74 |
-
"""
|
| 75 |
-
|
| 76 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
| 102 |
raise gr.Error("Please upload at least one scoresheet image.")
|
| 103 |
-
|
| 104 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 105 |
|
| 106 |
out_dir = Path(tempfile.mkdtemp(prefix="togyz_"))
|
| 107 |
-
round_slug = _safe_slug(
|
| 108 |
-
rows, gallery, all_files = [], [], []
|
| 109 |
|
| 110 |
-
for i, (img,
|
| 111 |
-
prefix = f"{round_slug}
|
| 112 |
try:
|
| 113 |
-
row, gal, files = _process_one(
|
|
|
|
|
|
|
| 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
|
| 150 |
-
"
|
|
|
|
| 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 |
-
|
| 156 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 157 |
|
| 158 |
-
|
| 159 |
-
|
| 160 |
-
|
| 161 |
-
|
| 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=[
|
| 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
|
|
@@ -4,12 +4,15 @@
|
|
| 4 |
"cell_type": "markdown",
|
| 5 |
"metadata": {},
|
| 6 |
"source": [
|
| 7 |
-
"# Togyzkumalak
|
| 8 |
"\n",
|
| 9 |
"Runtime → Change runtime type → **GPU**, then run the cells in order.\n",
|
| 10 |
"\n",
|
| 11 |
-
"
|
| 12 |
-
"
|
|
|
|
|
|
|
|
|
|
| 13 |
]
|
| 14 |
},
|
| 15 |
{
|
|
@@ -18,13 +21,23 @@
|
|
| 18 |
"metadata": {},
|
| 19 |
"outputs": [],
|
| 20 |
"source": [
|
| 21 |
-
"
|
| 22 |
-
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 26 |
"\n",
|
| 27 |
-
"
|
|
|
|
|
|
|
| 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
|
| 46 |
-
"# (--emnist is needed because ingredients/ is gitignored
|
|
|
|
| 47 |
"!python scripts/get_glyphs.py --emnist"
|
| 48 |
]
|
| 49 |
},
|
|
@@ -53,10 +67,12 @@
|
|
| 53 |
"metadata": {},
|
| 54 |
"outputs": [],
|
| 55 |
"source": [
|
| 56 |
-
"# Sanity check
|
| 57 |
-
"!python -m togyz.synth --preview
|
| 58 |
-
"
|
| 59 |
-
"
|
|
|
|
|
|
|
| 60 |
]
|
| 61 |
},
|
| 62 |
{
|
|
@@ -65,7 +81,8 @@
|
|
| 65 |
"metadata": {},
|
| 66 |
"outputs": [],
|
| 67 |
"source": [
|
| 68 |
-
"
|
|
|
|
| 69 |
]
|
| 70 |
},
|
| 71 |
{
|
|
@@ -74,11 +91,43 @@
|
|
| 74 |
"metadata": {},
|
| 75 |
"outputs": [],
|
| 76 |
"source": [
|
| 77 |
-
"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 78 |
"\n",
|
| 79 |
-
"# download the trained model\n",
|
| 80 |
"from google.colab import files\n",
|
| 81 |
-
"files.download('
|
| 82 |
]
|
| 83 |
}
|
| 84 |
],
|
|
@@ -89,5 +138,5 @@
|
|
| 89 |
"language_info": {"name": "python"}
|
| 90 |
},
|
| 91 |
"nbformat": 4,
|
| 92 |
-
"nbformat_minor":
|
| 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 |
}
|
|
@@ -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("--
|
| 34 |
-
help="
|
| 35 |
-
"if the file
|
|
|
|
| 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 |
-
|
| 56 |
-
if Path(args.
|
| 57 |
-
|
| 58 |
else:
|
| 59 |
-
print(f"No
|
| 60 |
|
| 61 |
print(f"Reading {args.image} ...")
|
| 62 |
out = run_pipeline(
|
| 63 |
-
args.image, moves_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["
|
| 75 |
-
|
| 76 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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))
|
|
@@ -4,7 +4,7 @@
|
|
| 4 |
|
| 5 |
Writes, next to each source checkpoint:
|
| 6 |
checkpoints/best.onnx + checkpoints/best.classes.json
|
| 7 |
-
checkpoints/
|
| 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/
|
| 70 |
-
help="checkpoint(s) to export (default: moves +
|
| 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:
|
|
@@ -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 |
-
#
|
| 22 |
-
#
|
| 23 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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]
|
|
@@ -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 =
|
| 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 |
|
|
@@ -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 |
-
|
| 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,
|
| 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
|
| 156 |
-
"""
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 157 |
logp = 0.0
|
| 158 |
for side, player in (("W", 0), ("B", 1)):
|
| 159 |
-
probs =
|
| 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
|
|
|
|
|
|
|
|
|
|
| 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 +=
|
| 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 |
-
|
|
|
|
| 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,
|
| 216 |
-
result=None, topk=5, beam_width=
|
| 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,
|
| 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 |
-
#
|
|
|
|
| 298 |
checkpoints, checkpoint_report = {}, []
|
| 299 |
-
|
| 300 |
-
|
| 301 |
-
|
| 302 |
-
|
| 303 |
-
|
| 304 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 305 |
temperature=temperature, tta=tta,
|
| 306 |
)
|
| 307 |
-
for (cell, img), probs in zip(
|
| 308 |
-
if
|
|
|
|
|
|
|
|
|
|
|
|
|
| 309 |
continue
|
| 310 |
-
|
| 311 |
-
|
| 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 |
-
|
|
|
|
|
|
|
| 319 |
|
| 320 |
beam_moves, beam_logp, beam_shares = beam_decode(
|
| 321 |
-
ply_prob_arrays, classes,
|
| 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(
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
"
|
| 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 |
-
"
|
|
|
|
| 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,
|
|
@@ -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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 187 |
-
|
| 188 |
|
| 189 |
-
|
| 190 |
-
|
| 191 |
-
|
| 192 |
-
|
| 193 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 194 |
"""
|
| 195 |
gray = np.asarray(sheet)
|
| 196 |
h, w = gray.shape
|
| 197 |
mask = _grid_mask(gray)
|
| 198 |
-
|
| 199 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 200 |
n, _, stats, _ = cv2.connectedComponentsWithStats(mask)
|
| 201 |
-
|
| 202 |
tuple(int(v) for v in stats[i][:4])
|
| 203 |
for i in range(1, n)
|
| 204 |
-
if stats[i][2] > w * 0.
|
| 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] <
|
| 208 |
]
|
| 209 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 210 |
|
| 211 |
cells: list[Cell] = []
|
| 212 |
-
for
|
| 213 |
-
|
| 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
|
| 285 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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))
|
|
@@ -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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 210 |
"""Render a comparison grid: the real-crop classes first, then random ones."""
|
| 211 |
sampler = GlyphSampler()
|
| 212 |
rng = random.Random(seed)
|
| 213 |
-
|
| 214 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 =
|
| 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)
|
|
@@ -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,
|
| 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", "
|
| 56 |
-
help="moves: 163-class cell classifier;
|
| 57 |
-
"
|
|
|
|
|
|
|
| 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/
|
| 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 =
|
| 75 |
-
out_dir = Path(args.out or ("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,
|
| 80 |
-
|
|
|
|
|
|
|
|
|
|
| 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)}")
|