ansarzeinulla commited on
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)

Files changed (2) hide show
  1. app.py +67 -23
  2. 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
- DEFAULT_BEAM = "8192"
 
 
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 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)
@@ -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
- 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):
243
  """Turn infrastructure overload into a friendly message instead of a 500."""
244
  try:
245
- return convert(*args)
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 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."
@@ -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
- 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
 
290
  results_table = gr.Dataframe(
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")
@@ -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, *image_slots],
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 _legal_scores(probs, legal_moves, classes, class_idx):
121
- """Redistribute the 162-class distribution over the current legal moves.
122
 
123
- A legal move earns credit from (a) its exact class string and (b) a
124
- factorized term P(source digit) * P(landing digit), times agreement with
125
- the written x-mark probability. Scores stay UNNORMALIZED for the beam so a
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(), classes, class_idx)
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"]}