Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
60 changes: 41 additions & 19 deletions mito_train/training/train_board_ocr.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
)


Expand Down Expand Up @@ -131,13 +131,17 @@ 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 {
"board/cell_acc": board_cell_acc,
"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,
}

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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}"
)

Expand Down Expand Up @@ -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()})
Expand All @@ -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()
Expand Down Expand Up @@ -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 <ckpt-root>/<model-name>/. "
"Defaults to 'board-ocr-<backbone>'.")
p.add_argument("--ckpt-root", type=Path, default=Path("./runs"),
help="Directory under which each run gets its own <model-name>/ 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/<model-name>/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)


Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
@@ -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 = [
Expand Down
6 changes: 3 additions & 3 deletions scripts/train.sh
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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=()
Expand All @@ -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[@]}"
10 changes: 6 additions & 4 deletions scripts/train_backbones.sh
Original file line number Diff line number Diff line change
Expand Up @@ -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}"

Expand Down Expand Up @@ -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
Expand All @@ -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"
Expand Down Expand Up @@ -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 &
Expand Down
6 changes: 3 additions & 3 deletions scripts/train_multi_gpu.sh
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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.
Expand All @@ -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[@]}"
2 changes: 1 addition & 1 deletion uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading