From b1e6a98380a3d8892c8052f94a5c178e0e940bab Mon Sep 17 00:00:00 2001 From: tkgstrator <29420801+tkgstrator@users.noreply.github.com> Date: Mon, 13 Jul 2026 03:38:27 +0000 Subject: [PATCH] feat(train_board_ocr): replace --ckpt-dir with --model-name and add hand MAE - Replace --ckpt-dir with --model-name (+ --ckpt-root, default ./runs). Checkpoints now land at //. model_name defaults to board-ocr-, matching the sweep naming. - --resume latest no longer errors when latest.pt is missing; it falls through to a fresh run so sweep scripts can pass it unconditionally. - Add hand/mae to metrics and surface board_acc, hand_acc, and hand_mae in train + val log lines. MAE treats the count head as ordinal, so "off by 1" reads very differently from "off by 9". - Update train.sh, train_multi_gpu.sh, train_backbones.sh to pass --model-name. - Bump version to 0.3.0 (CLI flag rename is breaking; pre-1.0 => minor). Co-Authored-By: Claude Opus 4.8 (1M context) --- mito_train/training/train_board_ocr.py | 60 ++++++++++++++++++-------- pyproject.toml | 2 +- scripts/train.sh | 6 +-- scripts/train_backbones.sh | 10 +++-- scripts/train_multi_gpu.sh | 6 +-- uv.lock | 2 +- 6 files changed, 55 insertions(+), 31 deletions(-) diff --git a/mito_train/training/train_board_ocr.py b/mito_train/training/train_board_ocr.py index e2dc504..fab556a 100644 --- a/mito_train/training/train_board_ocr.py +++ b/mito_train/training/train_board_ocr.py @@ -102,7 +102,7 @@ def compute_loss( METRIC_KEYS = ( "board/cell_acc", "board/full_acc", - "hand/slot_acc", "hand/full_acc", "sfen/full_acc", + "hand/slot_acc", "hand/full_acc", "hand/mae", "sfen/full_acc", ) @@ -131,6 +131,9 @@ def compute_metrics( hand_slot_acc = hand_ok.float().mean() hand_all = hand_ok.all(dim=1) hand_full = hand_all.float().mean() + # MAE over all 14 slots — treats the count head as ordinal so + # "off by 1" is much better than "off by 9". + hand_mae = (hand_pred - hand_target).abs().float().mean() sfen_full = (board_all & hand_all).float().mean() return { @@ -138,6 +141,7 @@ def compute_metrics( "board/full_acc": board_full, "hand/slot_acc": hand_slot_acc, "hand/full_acc": hand_full, + "hand/mae": hand_mae, "sfen/full_acc": sfen_full, } @@ -276,25 +280,31 @@ def log_main(msg: str) -> None: model.parameters(), lr=args.lr, weight_decay=1e-4, fused=use_cuda, ) + ckpt_dir = args.ckpt_root / args.model_name if is_main(): - args.ckpt_dir.mkdir(parents=True, exist_ok=True) + ckpt_dir.mkdir(parents=True, exist_ok=True) start_epoch = 1 resumed_from: str | None = None resumed_wandb_id: str | None = None if args.resume is not None: - ckpt_path = args.ckpt_dir / "latest.pt" if str(args.resume) == "latest" else args.resume - log_main(f"[board_ocr] resume from {ckpt_path}") - ckpt = torch.load(ckpt_path, map_location=device, weights_only=False) - _unwrap(model).load_state_dict(ckpt["model"]) - optimizer.load_state_dict(ckpt["optimizer"]) - if ckpt.get("scaler") is not None and scaler.is_enabled(): - scaler.load_state_dict(ckpt["scaler"]) - start_epoch = ckpt["epoch"] + 1 - resumed_from = str(ckpt_path) - resumed_wandb_id = ckpt.get("wandb_run_id") - log_main(f"[board_ocr] resumed at epoch={start_epoch} (ckpt was epoch {ckpt['epoch']})") - if resumed_wandb_id: - log_main(f"[board_ocr] will resume wandb run id={resumed_wandb_id}") + ckpt_path = ckpt_dir / "latest.pt" if str(args.resume) == "latest" else args.resume + # `--resume latest` on a fresh model_name has no checkpoint to load — treat + # as fresh start so sweep scripts can pass --resume latest unconditionally. + if str(args.resume) == "latest" and not ckpt_path.exists(): + log_main(f"[board_ocr] --resume latest requested but {ckpt_path} not found — starting fresh") + else: + log_main(f"[board_ocr] resume from {ckpt_path}") + ckpt = torch.load(ckpt_path, map_location=device, weights_only=False) + _unwrap(model).load_state_dict(ckpt["model"]) + optimizer.load_state_dict(ckpt["optimizer"]) + if ckpt.get("scaler") is not None and scaler.is_enabled(): + scaler.load_state_dict(ckpt["scaler"]) + start_epoch = ckpt["epoch"] + 1 + resumed_from = str(ckpt_path) + resumed_wandb_id = ckpt.get("wandb_run_id") + log_main(f"[board_ocr] resumed at epoch={start_epoch} (ckpt was epoch {ckpt['epoch']})") + if resumed_wandb_id: + log_main(f"[board_ocr] will resume wandb run id={resumed_wandb_id}") if args.wandb_run_id: resumed_wandb_id = args.wandb_run_id @@ -408,6 +418,9 @@ def log_main(msg: str) -> None: f"[board_ocr] epoch={epoch:3d} " f"loss={avg_loss:.4f} (board={avg_bl:.4f} hand={avg_hl:.4f}) " f"cell_acc={avg['board/cell_acc']:.3f} " + f"board_acc={avg['board/full_acc']:.3f} " + f"hand_acc={avg['hand/full_acc']:.3f} " + f"hand_mae={avg['hand/mae']:.3f} " f"sfen_acc={avg['sfen/full_acc']:.3f}" ) @@ -448,6 +461,9 @@ def log_main(msg: str) -> None: } log_main( f"[board_ocr] val cell_acc={v_avg['board/cell_acc']:.3f} " + f"board_acc={v_avg['board/full_acc']:.3f} " + f"hand_acc={v_avg['hand/full_acc']:.3f} " + f"hand_mae={v_avg['hand/mae']:.3f} " f"sfen_acc={v_avg['sfen/full_acc']:.3f}" ) wandb_log.update({f"val/{k}": v for k, v in v_avg.items()}) @@ -471,9 +487,9 @@ def log_main(msg: str) -> None: "hand_weight": args.hand_weight, "wandb_run_id": wandb_run.id if wandb_run is not None else None, } - torch.save(ckpt_payload, args.ckpt_dir / "latest.pt") + torch.save(ckpt_payload, ckpt_dir / "latest.pt") if epoch % args.save_every == 0 or epoch == args.epochs: - torch.save(ckpt_payload, args.ckpt_dir / f"epoch-{epoch:03d}.pt") + torch.save(ckpt_payload, ckpt_dir / f"epoch-{epoch:03d}.pt") if wandb_run is not None: wandb_run.finish() @@ -539,14 +555,20 @@ def main() -> None: p.add_argument("--val-every", type=int, default=2) p.add_argument("--log-every", type=int, default=20, help="Log per-step train metrics to W&B every N batches.") - p.add_argument("--ckpt-dir", type=Path, default=Path("./runs/board-ocr")) + p.add_argument("--model-name", type=str, default=None, + help="Run identifier. Checkpoints go to //. " + "Defaults to 'board-ocr-'.") + p.add_argument("--ckpt-root", type=Path, default=Path("./runs"), + help="Directory under which each run gets its own / subdir.") p.add_argument("--save-every", type=int, default=5, help="Interval for saving epoch-{N}.pt snapshots. latest.pt is saved every epoch.") p.add_argument("--resume", type=Path, default=None, - help="Checkpoint path. Pass 'latest' to load --ckpt-dir/latest.pt.") + help="Checkpoint path. Pass 'latest' to load ./runs//latest.pt.") p.add_argument("--wandb-run-id", type=str, default=None, help="Force resume this W&B run id (overrides ckpt's stored id).") args = p.parse_args() + if args.model_name is None: + args.model_name = f"board-ocr-{args.backbone}" run(args) diff --git a/pyproject.toml b/pyproject.toml index ef2cf36..65b9dfb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "mito-train" -version = "0.2.0" +version = "0.3.0" description = "ぴよ将棋OCR 3モデル (board-detector / piece-classifier / hand-classifier) の学習・ONNX出力" requires-python = ">=3.11" dependencies = [ diff --git a/scripts/train.sh b/scripts/train.sh index a422457..64116e8 100755 --- a/scripts/train.sh +++ b/scripts/train.sh @@ -1,7 +1,7 @@ #!/usr/bin/env bash # board OCR full 学習: mobilenet_v3_small + Apple Silicon MPS 前提。 # 環境変数で上書き可: EPOCHS, BATCH_SIZE, NUM_WORKERS, LR, BACKBONE, IMAGE_SIZE, -# CKPT_DIR, SAVE_EVERY, RESUME, WANDB_RUN_ID, HF_REPO_ID +# MODEL_NAME, SAVE_EVERY, RESUME, WANDB_RUN_ID, HF_REPO_ID # 途中再開したいときは: RESUME=latest EPOCHS=30 ./scripts/train.sh # 過去の W&B run に強制的に紐づけたいときは: # WANDB_RUN_ID=xxxxxxxx RESUME=latest EPOCHS=30 ./scripts/train.sh @@ -15,7 +15,7 @@ NUM_WORKERS=${NUM_WORKERS:-8} LR=${LR:-6e-4} BACKBONE=${BACKBONE:-mobilenet_v3_small} IMAGE_SIZE=${IMAGE_SIZE:-288} -CKPT_DIR=${CKPT_DIR:-./runs/board-ocr-v2} +MODEL_NAME=${MODEL_NAME:-board-ocr-v2} SAVE_EVERY=${SAVE_EVERY:-5} extra_args=() @@ -37,6 +37,6 @@ uv run python -m mito_train.training.train_board_ocr \ --batch-size "${BATCH_SIZE}" \ --num-workers "${NUM_WORKERS}" \ --lr "${LR}" \ - --ckpt-dir "${CKPT_DIR}" \ + --model-name "${MODEL_NAME}" \ --save-every "${SAVE_EVERY}" \ "${extra_args[@]}" diff --git a/scripts/train_backbones.sh b/scripts/train_backbones.sh index 7c753ad..1aa74f7 100755 --- a/scripts/train_backbones.sh +++ b/scripts/train_backbones.sh @@ -139,7 +139,8 @@ resolve_run_state() { # and DDP modes. Returns the training exit code. run_single() { local BB="$1" - local ckpt_dir="${CKPT_ROOT}/board-ocr-${BB}" + local model_name="board-ocr-${BB}" + local ckpt_dir="${CKPT_ROOT}/${model_name}" local latest_ckpt="${ckpt_dir}/latest.pt" local final_ckpt="${ckpt_dir}/${final_ckpt_name}" @@ -187,7 +188,7 @@ run_single() { --num-workers "${NUM_WORKERS}" \ --prefetch-factor "${PREFETCH_FACTOR}" \ --lr "${LR}" \ - --ckpt-dir "${ckpt_dir}" \ + --model-name "${model_name}" \ "${extra_args[@]}" \ "${per_run_extra[@]}" then @@ -205,7 +206,8 @@ run_single() { # Its stdout+stderr is redirected to per-backbone log file. Sets $!. launch_parallel() { local BB="$1" gpu="$2" - local ckpt_dir="${CKPT_ROOT}/board-ocr-${BB}" + local model_name="board-ocr-${BB}" + local ckpt_dir="${CKPT_ROOT}/${model_name}" local latest_ckpt="${ckpt_dir}/latest.pt" local final_ckpt="${ckpt_dir}/${final_ckpt_name}" local log_file="${PARALLEL_LOG_DIR}/${BB}.log" @@ -242,7 +244,7 @@ launch_parallel() { --num-workers "${NUM_WORKERS}" \ --prefetch-factor "${PREFETCH_FACTOR}" \ --lr "${LR}" \ - --ckpt-dir "${ckpt_dir}" \ + --model-name "${model_name}" \ "${extra_args[@]}" \ "${per_run_extra[@]}" \ >"${log_file}" 2>&1 & diff --git a/scripts/train_multi_gpu.sh b/scripts/train_multi_gpu.sh index 41ace65..04360b7 100755 --- a/scripts/train_multi_gpu.sh +++ b/scripts/train_multi_gpu.sh @@ -31,7 +31,7 @@ IMAGE_SIZE=${IMAGE_SIZE:-224} LR=${LR:-3e-4} BACKBONE=${BACKBONE:-mobilenet_v3_small} HF_REPO_ID=${HF_REPO_ID:-ultemica/piyoshogi} -CKPT_DIR=${CKPT_DIR:-./runs/board-ocr-${BACKBONE}} +MODEL_NAME=${MODEL_NAME:-board-ocr-${BACKBONE}} SAVE_EVERY=${SAVE_EVERY:-5} extra_args=(--preload) @@ -48,7 +48,7 @@ fi echo "[train_multi_gpu] nproc=${NPROC_PER_NODE} backbone=${BACKBONE} batch=${BATCH_SIZE} (per rank)" echo "[train_multi_gpu] effective batch = ${BATCH_SIZE} x ${NPROC_PER_NODE} = $(( BATCH_SIZE * NPROC_PER_NODE ))" echo "[train_multi_gpu] workers=${PER_GPU_WORKERS} per rank ($(( PER_GPU_WORKERS * NPROC_PER_NODE )) total)" -echo "[train_multi_gpu] ckpt_dir=${CKPT_DIR}" +echo "[train_multi_gpu] model_name=${MODEL_NAME} (ckpt_dir=./runs/${MODEL_NAME})" # --standalone: single-node rendezvous, avoids setting MASTER_ADDR/MASTER_PORT # manually for the common single-host case. @@ -66,6 +66,6 @@ uv run torchrun \ --num-workers "${PER_GPU_WORKERS}" \ --prefetch-factor "${PREFETCH_FACTOR}" \ --lr "${LR}" \ - --ckpt-dir "${CKPT_DIR}" \ + --model-name "${MODEL_NAME}" \ --save-every "${SAVE_EVERY}" \ "${extra_args[@]}" diff --git a/uv.lock b/uv.lock index f017df5..348ec64 100644 --- a/uv.lock +++ b/uv.lock @@ -1119,7 +1119,7 @@ wheels = [ [[package]] name = "mito-train" -version = "0.2.0" +version = "0.3.0" source = { editable = "." } dependencies = [ { name = "albumentations" },