Spaces:
Sleeping
Sleeping
Commit ·
cfc041e
1
Parent(s): e5800bc
Speed up beam, stream per-image results with progress, single upload
Browse files- pipeline: hoist per-ply digit marginals out of the per-hypothesis loop
(_ply_marginals), ~2.1x faster beam with identical output; thread a
progress_cb through run_pipeline and beam_decode (phase + per-ply)
- app: single multi-file uploader replaces the 5 fixed image slots; convert
is now a generator that streams each image's result as it finishes and
builds the unified zip only at the end; gr.Progress bar shows per-image
phase/ply; default beam lowered 8192 -> 2048 (a full game was ~2min at
8192, ~23s at 2048, and diagram evidence prunes well)
- app.py +67 -23
- togyz/pipeline.py +43 -12
app.py
CHANGED
|
@@ -56,7 +56,9 @@ from togyz.pipeline import load_classifier, run_pipeline
|
|
| 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 |
-
|
|
|
|
|
|
|
| 60 |
|
| 61 |
# accepted spellings of a game result -> pipeline result code
|
| 62 |
RESULT_MAP = {
|
|
@@ -170,10 +172,11 @@ def _pgn_tags(meta: dict, game: dict) -> str:
|
|
| 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", "")
|
|
@@ -194,11 +197,32 @@ def _process_one(image_path, game: dict, meta: dict, beam_width: int,
|
|
| 194 |
return row, (out["annotated_image"], caption), files, warnings
|
| 195 |
|
| 196 |
|
| 197 |
-
def
|
| 198 |
-
"""
|
| 199 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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)
|
|
@@ -210,39 +234,54 @@ def convert(meta_text, games_text, beam_choice, img1, img2, img3, img4, img5):
|
|
| 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 |
-
|
| 216 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
| 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 |
-
|
| 240 |
|
| 241 |
|
| 242 |
-
def _busy_wrapper(
|
| 243 |
"""Turn infrastructure overload into a friendly message instead of a 500."""
|
| 244 |
try:
|
| 245 |
-
|
| 246 |
except gr.Error:
|
| 247 |
raise
|
| 248 |
except Exception as exc: # noqa: BLE001 - surface anything else gracefully
|
|
@@ -255,9 +294,10 @@ 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
|
| 259 |
"describe the round and the games in the two text fields, then "
|
| 260 |
-
"**Convert**
|
|
|
|
| 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."
|
|
@@ -280,16 +320,20 @@ with gr.Blocks(title="Togyzkumalak Scoresheet Reader") as demo:
|
|
| 280 |
"thorough; very large values are trimmed to fit memory)",
|
| 281 |
)
|
| 282 |
|
| 283 |
-
|
| 284 |
-
|
| 285 |
-
|
| 286 |
-
|
|
|
|
|
|
|
|
|
|
| 287 |
|
| 288 |
convert_btn = gr.Button("Convert", variant="primary")
|
| 289 |
|
| 290 |
results_table = gr.Dataframe(
|
| 291 |
headers=["game", "legal plies", "stopped", "beam PGN"],
|
| 292 |
-
label="Results
|
|
|
|
| 293 |
)
|
| 294 |
warnings_md = gr.Markdown()
|
| 295 |
gallery = gr.Gallery(label="Annotated reconstruction", columns=2, height="auto")
|
|
@@ -297,7 +341,7 @@ with gr.Blocks(title="Togyzkumalak Scoresheet Reader") as demo:
|
|
| 297 |
|
| 298 |
convert_btn.click(
|
| 299 |
_busy_wrapper,
|
| 300 |
-
inputs=[meta_box, games_box, beam_dd,
|
| 301 |
outputs=[results_table, gallery, zip_out, warnings_md],
|
| 302 |
)
|
| 303 |
|
|
|
|
| 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 |
+
# A full ~160-ply game costs roughly (width/1024) x 12s of beam search, so the
|
| 60 |
+
# default stays modest; the board-diagram evidence prunes well at this width.
|
| 61 |
+
DEFAULT_BEAM = "2048"
|
| 62 |
|
| 63 |
# accepted spellings of a game result -> pipeline result code
|
| 64 |
RESULT_MAP = {
|
|
|
|
| 172 |
|
| 173 |
|
| 174 |
def _process_one(image_path, game: dict, meta: dict, beam_width: int,
|
| 175 |
+
base_name: str, out_dir: Path, progress_cb=None):
|
| 176 |
"""Run the pipeline on one image; return (row, gallery_item, files, warnings)."""
|
| 177 |
out = run_pipeline(image_path, _MOVES, _DIAGRAM,
|
| 178 |
+
result=game["result"], beam_width=beam_width,
|
| 179 |
+
progress_cb=progress_cb)
|
| 180 |
|
| 181 |
stop = out["stopped"]
|
| 182 |
stop_txt = stop.get("reason", "")
|
|
|
|
| 197 |
return row, (out["annotated_image"], caption), files, warnings
|
| 198 |
|
| 199 |
|
| 200 |
+
def _as_paths(files) -> list[str]:
|
| 201 |
+
"""Normalize the multi-file uploader value into a list of file paths."""
|
| 202 |
+
if not files:
|
| 203 |
+
return []
|
| 204 |
+
if isinstance(files, (str, os.PathLike)):
|
| 205 |
+
files = [files]
|
| 206 |
+
paths = []
|
| 207 |
+
for f in files:
|
| 208 |
+
# gr.File yields str paths (type="filepath") or objects with .name
|
| 209 |
+
paths.append(f if isinstance(f, str) else getattr(f, "name", str(f)))
|
| 210 |
+
return paths
|
| 211 |
+
|
| 212 |
+
|
| 213 |
+
def convert(meta_text, games_text, beam_choice, files, progress=gr.Progress()):
|
| 214 |
+
"""One batch: up to 5 images sharing tournament metadata.
|
| 215 |
+
|
| 216 |
+
A generator: it yields (table, gallery, zip, warnings) after each image so
|
| 217 |
+
results stream in one by one; the zip download is assembled only at the
|
| 218 |
+
end and contains every game's PGNs appended together.
|
| 219 |
+
"""
|
| 220 |
+
images = _as_paths(files)
|
| 221 |
if not images:
|
| 222 |
raise gr.Error("Please upload at least one scoresheet image.")
|
| 223 |
+
if len(images) > MAX_IMAGES:
|
| 224 |
+
raise gr.Error(f"This demo handles at most {MAX_IMAGES} images per run "
|
| 225 |
+
f"(got {len(images)}).")
|
| 226 |
|
| 227 |
# parse everything up front - all input errors surface before any heavy work
|
| 228 |
meta = _parse_meta(meta_text)
|
|
|
|
| 234 |
|
| 235 |
out_dir = Path(tempfile.mkdtemp(prefix="togyz_"))
|
| 236 |
round_slug = _safe_slug(meta["round"])
|
| 237 |
+
n = len(images)
|
| 238 |
rows, gallery, all_files, all_warnings = [], [], [], []
|
| 239 |
|
| 240 |
+
def warn_md():
|
| 241 |
+
return "\n".join(f"⚠ {w}" for w in dict.fromkeys(all_warnings))
|
| 242 |
+
|
| 243 |
+
for i, (img, game) in enumerate(zip(images, games)):
|
| 244 |
+
name = Path(img).name
|
| 245 |
+
prefix = f"round{round_slug}_game{i + 1}" if round_slug else f"game{i + 1}"
|
| 246 |
+
|
| 247 |
+
def cb(frac, desc, i=i, name=name):
|
| 248 |
+
# blend per-image progress into an overall 0..1 bar
|
| 249 |
+
progress((i + frac) / n, desc=f"Image {i + 1}/{n} ({name}): {desc}")
|
| 250 |
+
|
| 251 |
+
cb(0.0, "starting")
|
| 252 |
try:
|
| 253 |
row, gal, files, warns = _process_one(
|
| 254 |
+
img, game, meta, beam_width, prefix, out_dir, progress_cb=cb
|
| 255 |
)
|
| 256 |
except Exception as exc: # one bad image must not kill the batch
|
| 257 |
+
rows.append([f"{prefix} ({name})", 0, f"error: {exc}", ""])
|
| 258 |
+
yield rows[:], gallery[:], None, warn_md()
|
| 259 |
continue
|
| 260 |
+
row[0] = f"{row[0]} ({name})"
|
| 261 |
rows.append(row)
|
| 262 |
gallery.append(gal)
|
| 263 |
all_files.extend(files)
|
| 264 |
all_warnings.extend(warns)
|
| 265 |
+
# stream this image's result immediately; zip is built only at the end
|
| 266 |
+
yield rows[:], gallery[:], None, warn_md()
|
| 267 |
|
|
|
|
| 268 |
if not all_files:
|
| 269 |
# every image errored - still return the table so the user sees why
|
| 270 |
+
yield rows[:], gallery[:], None, warn_md()
|
| 271 |
+
return
|
| 272 |
|
| 273 |
zip_path = out_dir / (f"round{round_slug}_pgns.zip" if round_slug else "pgns.zip")
|
| 274 |
with zipfile.ZipFile(zip_path, "w") as zf:
|
| 275 |
for f in all_files:
|
| 276 |
zf.write(f, arcname=Path(f).name)
|
| 277 |
|
| 278 |
+
yield rows[:], gallery[:], str(zip_path), warn_md()
|
| 279 |
|
| 280 |
|
| 281 |
+
def _busy_wrapper(meta_text, games_text, beam_choice, files, progress=gr.Progress()):
|
| 282 |
"""Turn infrastructure overload into a friendly message instead of a 500."""
|
| 283 |
try:
|
| 284 |
+
yield from convert(meta_text, games_text, beam_choice, files, progress)
|
| 285 |
except gr.Error:
|
| 286 |
raise
|
| 287 |
except Exception as exc: # noqa: BLE001 - surface anything else gracefully
|
|
|
|
| 294 |
with gr.Blocks(title="Togyzkumalak Scoresheet Reader") as demo:
|
| 295 |
gr.Markdown(
|
| 296 |
"# Togyzkumalak Scoresheet Reader\n"
|
| 297 |
+
"Upload up to **5** scoresheet photos of one round in a single batch, "
|
| 298 |
"describe the round and the games in the two text fields, then "
|
| 299 |
+
"**Convert**. Results stream in image by image with a live progress "
|
| 300 |
+
"bar; the combined PGN download appears once every image is done.\n\n"
|
| 301 |
"Outputs per game: `beam` (best legal reconstruction), `raw` (pure OCR), "
|
| 302 |
"`legal` (strict replay). Free demo — a first run may wake the Space, and "
|
| 303 |
"images are processed one at a time."
|
|
|
|
| 320 |
"thorough; very large values are trimmed to fit memory)",
|
| 321 |
)
|
| 322 |
|
| 323 |
+
image_files = gr.File(
|
| 324 |
+
label=f"Scoresheet photos (up to {MAX_IMAGES}, in the same order as the "
|
| 325 |
+
"game lines above)",
|
| 326 |
+
file_count="multiple",
|
| 327 |
+
file_types=["image"],
|
| 328 |
+
type="filepath",
|
| 329 |
+
)
|
| 330 |
|
| 331 |
convert_btn = gr.Button("Convert", variant="primary")
|
| 332 |
|
| 333 |
results_table = gr.Dataframe(
|
| 334 |
headers=["game", "legal plies", "stopped", "beam PGN"],
|
| 335 |
+
label="Results (stream in as each image finishes)",
|
| 336 |
+
wrap=True, interactive=False,
|
| 337 |
)
|
| 338 |
warnings_md = gr.Markdown()
|
| 339 |
gallery = gr.Gallery(label="Annotated reconstruction", columns=2, height="auto")
|
|
|
|
| 341 |
|
| 342 |
convert_btn.click(
|
| 343 |
_busy_wrapper,
|
| 344 |
+
inputs=[meta_box, games_box, beam_dd, image_files],
|
| 345 |
outputs=[results_table, gallery, zip_out, warnings_md],
|
| 346 |
)
|
| 347 |
|
togyz/pipeline.py
CHANGED
|
@@ -117,14 +117,12 @@ def classify_cells(images, clf: Classifier, temperature=1.0, tta=True, batch_siz
|
|
| 117 |
# --------------------------------------------------------------------------- #
|
| 118 |
# Legal-move scoring, kazan evidence, beam search (numpy port)
|
| 119 |
# --------------------------------------------------------------------------- #
|
| 120 |
-
def
|
| 121 |
-
"""
|
| 122 |
|
| 123 |
-
|
| 124 |
-
|
| 125 |
-
|
| 126 |
-
hypothesis whose legal set explains the observation poorly accumulates a
|
| 127 |
-
genuinely low joint likelihood (measured: 24/39 vs 29/39 when normalized).
|
| 128 |
"""
|
| 129 |
source_marginal = [0.0] * 9
|
| 130 |
landing_marginal = [0.0] * 9
|
|
@@ -136,7 +134,19 @@ def _legal_scores(probs, legal_moves, classes, class_idx):
|
|
| 136 |
landing_marginal[int(c[1]) - 1] += p
|
| 137 |
if c.endswith("x"):
|
| 138 |
x_marginal += p
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 139 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 140 |
scored = []
|
| 141 |
for move in legal_moves:
|
| 142 |
exact = float(probs[class_idx[move.notation]])
|
|
@@ -210,13 +220,14 @@ def safe_beam_width(beam_width: int, per_ply: int = 9) -> int:
|
|
| 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 |
"""
|
|
@@ -225,11 +236,16 @@ def beam_decode(ply_probs, classes, beam_width=1024, per_ply=9, result=None,
|
|
| 225 |
class_idx = {c: i for i, c in enumerate(classes)}
|
| 226 |
active = [(Game(), [], [], 0.0)] # (game, annotated_moves, chosen_shares, logp)
|
| 227 |
finished = []
|
|
|
|
| 228 |
|
| 229 |
-
for probs in ply_probs:
|
|
|
|
|
|
|
|
|
|
|
|
|
| 230 |
expanded = []
|
| 231 |
for game, moves, chosen, logp in active:
|
| 232 |
-
candidates = _legal_scores(probs, game.legal_moves(),
|
| 233 |
for raw, share, move in candidates[:per_ply]:
|
| 234 |
nxt = game.copy()
|
| 235 |
played = nxt.play(move.action)
|
|
@@ -263,17 +279,24 @@ def format_pgn(plies: list[str]) -> str:
|
|
| 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 |
|
| 269 |
`image` is a path or a PIL.Image. Returns a dict with the reconstructed
|
| 270 |
PGNs and diagnostics; no files are written unless `save_cells_dir` is set
|
| 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] = []
|
|
@@ -284,10 +307,12 @@ def run_pipeline(image, moves_clf: Classifier, diagram_clf: Classifier | None =
|
|
| 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 |
|
|
|
|
| 291 |
cleaned = [clean_cell(c.image) for c in cells] # (image, ink_ratio) pairs
|
| 292 |
all_probs = classify_cells(
|
| 293 |
[img for img, _ in cleaned], moves_clf, temperature=temperature, tta=tta
|
|
@@ -355,6 +380,7 @@ def run_pipeline(image, moves_clf: Classifier, diagram_clf: Classifier | None =
|
|
| 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))
|
|
@@ -404,10 +430,15 @@ def run_pipeline(image, moves_clf: Classifier, diagram_clf: Classifier | None =
|
|
| 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 = [
|
| 412 |
{"move": p["move"], "side": p["side"], "chosen": m,
|
| 413 |
"prob": round(q, 4), "agrees_with_raw": m.rstrip("+") == p["raw"]}
|
|
|
|
| 117 |
# --------------------------------------------------------------------------- #
|
| 118 |
# Legal-move scoring, kazan evidence, beam search (numpy port)
|
| 119 |
# --------------------------------------------------------------------------- #
|
| 120 |
+
def _ply_marginals(probs, classes):
|
| 121 |
+
"""Source/landing digit marginals and the x-mark mass for one ply.
|
| 122 |
|
| 123 |
+
Depends only on the ply's class distribution, so the beam computes it
|
| 124 |
+
once per ply and shares it across every hypothesis (this used to be
|
| 125 |
+
recomputed per hypothesis - the dominant cost at large beam widths).
|
|
|
|
|
|
|
| 126 |
"""
|
| 127 |
source_marginal = [0.0] * 9
|
| 128 |
landing_marginal = [0.0] * 9
|
|
|
|
| 134 |
landing_marginal[int(c[1]) - 1] += p
|
| 135 |
if c.endswith("x"):
|
| 136 |
x_marginal += p
|
| 137 |
+
return source_marginal, landing_marginal, x_marginal
|
| 138 |
+
|
| 139 |
+
|
| 140 |
+
def _legal_scores(probs, legal_moves, class_idx, marginals):
|
| 141 |
+
"""Redistribute the 162-class distribution over the current legal moves.
|
| 142 |
|
| 143 |
+
A legal move earns credit from (a) its exact class string and (b) a
|
| 144 |
+
factorized term P(source digit) * P(landing digit), times agreement with
|
| 145 |
+
the written x-mark probability. Scores stay UNNORMALIZED for the beam so a
|
| 146 |
+
hypothesis whose legal set explains the observation poorly accumulates a
|
| 147 |
+
genuinely low joint likelihood (measured: 24/39 vs 29/39 when normalized).
|
| 148 |
+
"""
|
| 149 |
+
source_marginal, landing_marginal, x_marginal = marginals
|
| 150 |
scored = []
|
| 151 |
for move in legal_moves:
|
| 152 |
exact = float(probs[class_idx[move.notation]])
|
|
|
|
| 220 |
|
| 221 |
|
| 222 |
def beam_decode(ply_probs, classes, beam_width=1024, per_ply=9, result=None,
|
| 223 |
+
checkpoints=None, progress_cb=None):
|
| 224 |
"""Longest fully-legal move sequence with the highest joint probability,
|
| 225 |
exploring all `per_ply` legal continuations per ply; the final pool is
|
| 226 |
re-ranked by diagram-checkpoint and known-result consistency.
|
| 227 |
|
| 228 |
`beam_width` is capped by `safe_beam_width` so user-requested widths up
|
| 229 |
to 2^30 degrade into a narrower beam instead of exhausting RAM.
|
| 230 |
+
`progress_cb(ply_index, total_plies)` is called once per ply if given.
|
| 231 |
|
| 232 |
Returns (annotated_moves, log_prob, per_ply_share_list).
|
| 233 |
"""
|
|
|
|
| 236 |
class_idx = {c: i for i, c in enumerate(classes)}
|
| 237 |
active = [(Game(), [], [], 0.0)] # (game, annotated_moves, chosen_shares, logp)
|
| 238 |
finished = []
|
| 239 |
+
total_plies = len(ply_probs)
|
| 240 |
|
| 241 |
+
for ply_index, probs in enumerate(ply_probs):
|
| 242 |
+
if progress_cb:
|
| 243 |
+
progress_cb(ply_index, total_plies)
|
| 244 |
+
# marginals are identical across all hypotheses this ply - compute once
|
| 245 |
+
marginals = _ply_marginals(probs, classes)
|
| 246 |
expanded = []
|
| 247 |
for game, moves, chosen, logp in active:
|
| 248 |
+
candidates = _legal_scores(probs, game.legal_moves(), class_idx, marginals)
|
| 249 |
for raw, share, move in candidates[:per_ply]:
|
| 250 |
nxt = game.copy()
|
| 251 |
played = nxt.play(move.action)
|
|
|
|
| 279 |
# --------------------------------------------------------------------------- #
|
| 280 |
def run_pipeline(image, moves_clf: Classifier, diagram_clf: Classifier | None = None,
|
| 281 |
result=None, topk=5, beam_width=8192, per_ply=9,
|
| 282 |
+
temperature=1.0, tta=True, save_cells_dir=None, progress_cb=None):
|
| 283 |
"""Read one scoresheet image into game records.
|
| 284 |
|
| 285 |
`image` is a path or a PIL.Image. Returns a dict with the reconstructed
|
| 286 |
PGNs and diagnostics; no files are written unless `save_cells_dir` is set
|
| 287 |
(the CLI uses it, the Gradio app does not).
|
| 288 |
|
| 289 |
+
`progress_cb(frac, desc)` - if given - is called with a 0..1 fraction and
|
| 290 |
+
a human-readable phase description as the read progresses.
|
| 291 |
+
|
| 292 |
Keys: beam_pgn, raw_pgn, legal_pgn, game_json (dict), stopped,
|
| 293 |
plies_scanned, legal_plies, beam_plies, checkpoint_report, warnings,
|
| 294 |
annotated_image (PIL), median_cell_height, low_resolution (bool).
|
| 295 |
"""
|
| 296 |
+
def report(frac, desc):
|
| 297 |
+
if progress_cb:
|
| 298 |
+
progress_cb(frac, desc)
|
| 299 |
+
|
| 300 |
classes = moves_clf.classes
|
| 301 |
empty_idx = classes.index(EMPTY)
|
| 302 |
warnings: list[str] = []
|
|
|
|
| 307 |
"inside the memory budget"
|
| 308 |
)
|
| 309 |
|
| 310 |
+
report(0.0, "detecting tables and cells")
|
| 311 |
cells, sheet, gridlines = extract_cells(image, with_gridlines=True)
|
| 312 |
cells.sort(key=lambda c: (c.move_no, c.side != "W")) # game order: 1W 1B 2W ...
|
| 313 |
median_h = sorted(c.bbox[3] for c in cells)[len(cells) // 2]
|
| 314 |
|
| 315 |
+
report(0.05, f"classifying {len(cells)} move cells")
|
| 316 |
cleaned = [clean_cell(c.image) for c in cells] # (image, ink_ratio) pairs
|
| 317 |
all_probs = classify_cells(
|
| 318 |
[img for img, _ in cleaned], moves_clf, temperature=temperature, tta=tta
|
|
|
|
| 380 |
checkpoints, checkpoint_report = {}, []
|
| 381 |
diagram_cells, diagram_labels = [], {}
|
| 382 |
if diagram_clf is not None:
|
| 383 |
+
report(0.30, "reading board diagrams")
|
| 384 |
dg_classes = diagram_clf.classes
|
| 385 |
value_names = [str(v) for v in range(82)]
|
| 386 |
kazan_allowed = set(str(v) for v in range(10, 82))
|
|
|
|
| 430 |
img.save(cells_dir /
|
| 431 |
f"{cell.kind}_{cell.move_no:02d}_{cell.side}{suffix}.png")
|
| 432 |
|
| 433 |
+
def beam_progress(ply_index, total):
|
| 434 |
+
report(0.45 + 0.53 * (ply_index / max(total, 1)),
|
| 435 |
+
f"reconstructing game (ply {ply_index + 1}/{total})")
|
| 436 |
+
|
| 437 |
beam_moves, beam_logp, beam_shares = beam_decode(
|
| 438 |
ply_prob_arrays, classes, effective_width, per_ply,
|
| 439 |
+
result=result, checkpoints=checkpoints, progress_cb=beam_progress,
|
| 440 |
)
|
| 441 |
+
report(1.0, "done")
|
| 442 |
beam_detail = [
|
| 443 |
{"move": p["move"], "side": p["side"], "chosen": m,
|
| 444 |
"prob": round(q, 4), "agrees_with_raw": m.rstrip("+") == p["raw"]}
|