diff --git a/app/storage/r2_client.py b/app/storage/r2_client.py index 17afe56..cc03368 100644 --- a/app/storage/r2_client.py +++ b/app/storage/r2_client.py @@ -1,6 +1,7 @@ from __future__ import annotations from dataclasses import dataclass +from pathlib import Path from socket import timeout as SocketTimeout import boto3 @@ -135,3 +136,33 @@ def get_object_bytes( if code in {"AccessDenied", "403"}: raise ObjectAccessDeniedError("Object access denied") from exc raise ObjectDownloadError("Object download failed") from exc + + def download_object(self, bucket: str, object_key: str, destination: Path) -> None: + try: + self._client.download_file(bucket, object_key, str(destination)) + except (ConnectTimeoutError, ReadTimeoutError, SocketTimeout) as exc: + raise ObjectTimeoutError("Object download timed out") from exc + except ClientError as exc: + code = exc.response.get("Error", {}).get("Code", "") + if code in {"NoSuchKey", "NoSuchBucket", "404"}: + raise ObjectNotFoundError("Object not found") from exc + if code in {"AccessDenied", "403"}: + raise ObjectAccessDeniedError("Object access denied") from exc + raise ObjectDownloadError("Object download failed") from exc + + def delete_object(self, bucket: str, object_key: str) -> None: + try: + self._client.delete_object(Bucket=bucket, Key=object_key) + except ClientError as exc: + code = exc.response.get("Error", {}).get("Code", "") + if code in {"AccessDenied", "403"}: + raise ObjectAccessDeniedError("Object delete access denied") from exc + raise ObjectDownloadError("Object delete failed") from exc + + def generate_presigned_put_url(self, bucket: str, object_key: str, expires_in_seconds: int) -> str: + """A short-lived, single-object write URL a remote runtime can use without holding credentials.""" + return self._client.generate_presigned_url( + "put_object", + Params={"Bucket": bucket, "Key": object_key}, + ExpiresIn=expires_in_seconds, + ) diff --git a/tests/test_r2_client_checkpoint_transfer.py b/tests/test_r2_client_checkpoint_transfer.py new file mode 100644 index 0000000..bab720e --- /dev/null +++ b/tests/test_r2_client_checkpoint_transfer.py @@ -0,0 +1,75 @@ +from __future__ import annotations + +from pathlib import Path + +import pytest +from botocore.exceptions import ClientError + +from app.storage.r2_client import ObjectAccessDeniedError, ObjectNotFoundError, R2ObjectStore + + +def _client_error(code: str) -> ClientError: + return ClientError({"Error": {"Code": code, "Message": code}}, "operation") + + +@pytest.fixture +def store() -> R2ObjectStore: + return R2ObjectStore.from_settings( + endpoint_url="https://example.invalid", + access_key_id="key", + secret_access_key="secret", + region_name="auto", + default_bucket="bucket", + allowed_buckets_raw="", + read_timeout_seconds=5, + ) + + +def test_generate_presigned_put_url_targets_the_given_bucket_and_key(store: R2ObjectStore) -> None: + url = store.generate_presigned_put_url("bucket", "colab-runs/run-1/checkpoint.pdparams", expires_in_seconds=900) + + assert "colab-runs/run-1/checkpoint.pdparams" in url + assert "X-Amz-Signature" in url + + +def test_download_object_writes_the_destination_file(store: R2ObjectStore, tmp_path: Path) -> None: + destination = tmp_path / "checkpoint.pdparams" + + def fake_download_file(bucket: str, key: str, filename: str) -> None: + assert (bucket, key) == ("bucket", "checkpoint.pdparams") + Path(filename).write_bytes(b"weights") + + store._client.download_file = fake_download_file # type: ignore[attr-defined] + + store.download_object("bucket", "checkpoint.pdparams", destination) + + assert destination.read_bytes() == b"weights" + + +def test_download_object_missing_key_raises_not_found(store: R2ObjectStore, tmp_path: Path) -> None: + def fake_download_file(*_args: object) -> None: + raise _client_error("NoSuchKey") + + store._client.download_file = fake_download_file # type: ignore[attr-defined] + + with pytest.raises(ObjectNotFoundError): + store.download_object("bucket", "missing.pdparams", tmp_path / "out") + + +def test_delete_object_access_denied_raises_access_denied_error(store: R2ObjectStore) -> None: + def fake_delete_object(**_kwargs: object) -> None: + raise _client_error("AccessDenied") + + store._client.delete_object = fake_delete_object # type: ignore[attr-defined] + + with pytest.raises(ObjectAccessDeniedError): + store.delete_object("bucket", "checkpoint.pdparams") + + +def test_delete_object_succeeds_without_error(store: R2ObjectStore) -> None: + calls: list[dict[str, object]] = [] + store._client.delete_object = lambda **kwargs: calls.append(kwargs) # type: ignore[attr-defined] + + store.delete_object("bucket", "checkpoint.pdparams") + + assert calls == [{"Bucket": "bucket", "Key": "checkpoint.pdparams"}] diff --git a/tests/test_run_rec_colab.py b/tests/test_run_rec_colab.py new file mode 100644 index 0000000..5d924fd --- /dev/null +++ b/tests/test_run_rec_colab.py @@ -0,0 +1,324 @@ +from __future__ import annotations + +import hashlib +import json +import sys +import tarfile +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from app.storage.r2_client import ObjectNotFoundError +from training import run_rec_colab + +SESSION_STOP = "stop" + + +class FakeR2Store: + """Stands in for R2ObjectStore: an in-memory object map keyed by (bucket, key).""" + + def __init__(self) -> None: + self.default_bucket = "test-bucket" + self.objects: dict[tuple[str, str], bytes] = {} + self.deleted: list[tuple[str, str]] = [] + + def generate_presigned_put_url(self, bucket: str, key: str, expires_in_seconds: int) -> str: + return f"https://example.invalid/put/{bucket}/{key}?expires={expires_in_seconds}" + + def download_object(self, bucket: str, key: str, destination: Path) -> None: + data = self.objects.get((bucket, key)) + if data is None: + raise ObjectNotFoundError(f"no such object: {bucket}/{key}") + destination.write_bytes(data) + + def delete_object(self, bucket: str, key: str) -> None: + self.objects.pop((bucket, key), None) + self.deleted.append((bucket, key)) + + +@pytest.fixture +def colab_run(tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + labels = tmp_path / "labels-dir" + (labels / "labels").mkdir(parents=True) + for name in ("train.txt", "holdout.txt"): + (labels / "labels" / name).write_text("a.png\ta\n", encoding="utf-8") + checkpoint = tmp_path / "base.pdparams" + checkpoint.write_bytes(b"base") + runs = tmp_path / "runs" + calls: list[list[str]] = [] + checkpoint_bytes = b"weights" + behavior = { + "exec_status": 0, + "remote_status": "success", + "stop_status": 0, + "eval_status": 0, + "upload_checkpoint": True, + "run_request": None, + "remote_log_text": "remote log contents\n", + "checkpoint_config": "Global: {}\n", + "train_log": "epoch 1 done\n", + "metadata_download_status": 0, + "session_active_after_stop_failure": False, + } + fake_r2 = FakeR2Store() + + def fake_capture(command: list[str]) -> SimpleNamespace: + assert command[1] == "sessions" + if behavior["session_active_after_stop_failure"]: + session_name = next(c[3] for c in calls if len(c) > 1 and c[1] == "new") + return SimpleNamespace(returncode=0, stdout=f"[{session_name}] fake-endpoint | Hardware: T4\n", stderr="") + return SimpleNamespace(returncode=0, stdout="[colab] No active sessions found on server.\n", stderr="") + + def fake_stream(command: list[str], _log: Path) -> int: + calls.append(command) + if command[0].endswith("evaluate_rec_checkpoint.sh"): + if behavior["eval_status"]: + return behavior["eval_status"] + output = Path(command[-1]) + output.mkdir(parents=True) + (output / "fixture_report.json").write_text('{"field_accuracy": 0.99, "run_code": {"field_accuracy": 1.0}}') + return 0 + verb = command[1] + if verb == "exec": + # Stands in for colab_remote.py actually running: populate the fake R2 object and + # the small metadata payload that fetch_remote_metadata/retrieve_checkpoint consume. + run_request = behavior["run_request"] + metadata: dict[str, object] = {"status": behavior["remote_status"]} + if behavior["remote_status"] != "success": + metadata["error"] = "RuntimeError: training crashed" + if behavior["upload_checkpoint"]: + upload = run_request["checkpoint_upload"] + fake_r2.objects[(upload["bucket"], upload["key"])] = checkpoint_bytes + metadata["checkpoint"] = { + "bucket": upload["bucket"], + "key": upload["key"], + "size_bytes": len(checkpoint_bytes), + "sha256": hashlib.sha256(checkpoint_bytes).hexdigest(), + } + metadata["checkpoint_config"] = behavior["checkpoint_config"] + metadata["train_log"] = behavior["train_log"] + behavior["remote_metadata"] = metadata + return behavior["exec_status"] + if verb == "download": + remote_path, local_path = command[4], Path(command[-1]) + if remote_path.endswith("run.json"): + if behavior["metadata_download_status"]: + return behavior["metadata_download_status"] + local_path.write_text(json.dumps(behavior["remote_metadata"]), encoding="utf-8") + return 0 + if remote_path.endswith("remote.log"): + local_path.write_text(behavior["remote_log_text"], encoding="utf-8") + return 0 + raise AssertionError(f"unexpected download path: {remote_path}") + if verb == SESSION_STOP: + return behavior["stop_status"] + return 0 + + monkeypatch.setattr(run_rec_colab, "RUNS", runs) + monkeypatch.setattr(run_rec_colab.shutil, "which", lambda _name: "colab") + monkeypatch.setattr(run_rec_colab, "validate_labels", lambda _path: (1, ["a.png"])) + monkeypatch.setattr(run_rec_colab, "PART_BYTES", 4) + monkeypatch.setattr(run_rec_colab, "require_r2_store", lambda _parser: fake_r2) + + def fake_stage(archive: Path, run_request: dict, *_rest) -> None: + behavior["run_request"] = run_request + archive.write_bytes(b"input-archive") + + monkeypatch.setattr(run_rec_colab, "stage_inputs", fake_stage) + monkeypatch.setattr(run_rec_colab, "stream_command", fake_stream) + monkeypatch.setattr(run_rec_colab, "capture_command", fake_capture) + monkeypatch.setattr( + sys, "argv", ["run_rec_colab.py", "--labels-dir", str(labels), "--pretrained-checkpoint", str(checkpoint)] + ) + return runs, calls, behavior, fake_r2 + + +def only_run(runs: Path) -> Path: + (run_dir,) = runs.iterdir() + return run_dir + + +def test_success_retrieves_checkpoint_from_r2_and_deletes_it(colab_run) -> None: + runs, calls, behavior, fake_r2 = colab_run + + assert run_rec_colab.main() == 0 + + request = behavior["run_request"] + assert request["training"]["epochs"] == 10 + assert set(request["paddle_wheel_mirror"]) == {"url", "sha256"} + assert request["checkpoint_upload"]["bucket"] == fake_r2.default_bucket + uploads = [command[-1] for command in calls if command[1] == "upload"] + assert uploads == [f"/content/ocrkit-input.part{i:04d}" for i in range(4)] + exec_command = next(command for command in calls if command[1] == "exec") + assert float(exec_command[exec_command.index("--timeout") + 1]) == 6 * 3600 + + run_dir = only_run(runs) + assert (run_dir / "checkpoint/best_accuracy.pdparams").read_bytes() == b"weights" + assert (run_dir / "checkpoint/config.yml").read_text() == "Global: {}\n" + assert (run_dir / "remote.log").read_text() == "remote log contents\n" + assert (run_dir / "evaluation/fixture_report.json").is_file() + assert json.loads((run_dir / "run.json").read_text())["evaluation"]["field_accuracy"] == 0.99 + assert json.loads((run_dir / "status.json").read_text())["runtime_stopped"] is True + stop_index = next(i for i, command in enumerate(calls) if command[1] == SESSION_STOP) + evaluation_index = next(i for i, command in enumerate(calls) if command[0].endswith("evaluate_rec_checkpoint.sh")) + assert stop_index < evaluation_index + assert not (run_dir / "accepted").exists() + assert fake_r2.deleted == [(request["checkpoint_upload"]["bucket"], request["checkpoint_upload"]["key"])] + assert not fake_r2.objects + + +def test_training_failure_keeps_partial_checkpoint_and_stops_runtime(colab_run) -> None: + runs, calls, behavior, fake_r2 = colab_run + behavior.update(exec_status=1, remote_status="failed") + + assert run_rec_colab.main() == 1 + + run_dir = only_run(runs) + assert not (run_dir / "checkpoint").exists() + assert not (run_dir / "run.json").exists() + assert (run_dir / "partial/checkpoint/best_accuracy.pdparams").is_file() + assert calls[-1][1] == SESSION_STOP + assert json.loads((run_dir / "status.json").read_text())["status"] == "failed" + assert not fake_r2.objects # still deleted even though the run overall failed + + +def test_local_evaluation_failure_demotes_checkpoint_to_partial(colab_run) -> None: + runs, _calls, behavior, _fake_r2 = colab_run + behavior["eval_status"] = 1 + + assert run_rec_colab.main() == 1 + + run_dir = only_run(runs) + assert not (run_dir / "checkpoint").exists() + assert not (run_dir / "run.json").exists() + assert (run_dir / "partial/checkpoint/best_accuracy.pdparams").is_file() + assert "local checkpoint evaluation failed" in json.loads((run_dir / "status.json").read_text())["error"] + assert json.loads((run_dir / "status.json").read_text())["runtime_stopped"] is True + + +def test_teardown_failure_with_a_still_active_session_demotes_checkpoint_to_partial(colab_run) -> None: + runs, _calls, behavior, _fake_r2 = colab_run + behavior["stop_status"] = 1 + behavior["session_active_after_stop_failure"] = True + + assert run_rec_colab.main() == 1 + + run_dir = only_run(runs) + assert not (run_dir / "checkpoint").exists() + assert (run_dir / "partial/checkpoint/best_accuracy.pdparams").is_file() + assert "colab stop" in json.loads((run_dir / "status.json").read_text())["error"] + + +def test_teardown_failure_when_colab_already_released_the_session_still_succeeds(colab_run) -> None: + """`colab stop` can 404 simply because Colab already reclaimed a finished runtime.""" + runs, _calls, behavior, _fake_r2 = colab_run + behavior["stop_status"] = 1 + behavior["session_active_after_stop_failure"] = False + + assert run_rec_colab.main() == 0 + + run_dir = only_run(runs) + assert (run_dir / "checkpoint/best_accuracy.pdparams").is_file() + status = json.loads((run_dir / "status.json").read_text()) + assert status["status"] == "success" + assert status["runtime_stopped"] is True + assert status["error"] is None + + +def test_provisioning_failure_stops_runtime_and_uploads_nothing(colab_run) -> None: + runs, calls, _behavior, fake_r2 = colab_run + original = run_rec_colab.stream_command + run_rec_colab.stream_command = lambda command, log: 1 if command[1] == "new" else original(command, log) + try: + assert run_rec_colab.main() == 1 + finally: + run_rec_colab.stream_command = original + + assert [command[1] for command in calls] == [SESSION_STOP] + assert not (only_run(runs) / "checkpoint").exists() + assert not fake_r2.objects + + +def test_remote_metadata_download_failure_is_reported_clearly(colab_run) -> None: + runs, _calls, behavior, _fake_r2 = colab_run + behavior["metadata_download_status"] = 1 + + assert run_rec_colab.main() == 1 + + error = json.loads((only_run(runs) / "status.json").read_text())["error"] + assert "run metadata" in error + + +def test_tampered_checkpoint_is_rejected_and_still_deleted(tmp_path: Path) -> None: + fake_r2 = FakeR2Store() + fake_r2.objects[("bucket", "key")] = b"tampered-bytes" + remote_metadata = { + "status": "success", + "checkpoint": {"bucket": "bucket", "key": "key", "size_bytes": 999, "sha256": "0" * 64}, + } + run_dir = tmp_path / "run" + run_dir.mkdir() + + with pytest.raises(ValueError, match="checksum"): + run_rec_colab.retrieve_checkpoint(fake_r2, remote_metadata, None, run_dir) + + assert not fake_r2.objects # cleaned up even though verification failed + assert json.loads((run_dir / "accepted" / "run.json").read_text())["status"] == "success" + + +def test_stage_inputs_skips_uploading_the_official_checkpoint(tmp_path: Path) -> None: + dataset_root = tmp_path / "dataset" + (dataset_root / "labels" / "images").mkdir(parents=True) + (dataset_root / "labels" / "train.txt").write_text("a.png\ta\n", encoding="utf-8") + (dataset_root / "labels" / "holdout.txt").write_text("", encoding="utf-8") + (dataset_root / "labels" / "images" / "a.png").write_bytes(b"png") + archive_path = tmp_path / "input.tar.gz" + + run_rec_colab.stage_inputs(archive_path, {"run_id": "x"}, dataset_root, None, ["a.png"], []) + + with tarfile.open(archive_path) as archive: + names = archive.getnames() + assert not any("pretrained" in name for name in names) + assert "dataset/a.png" in names + + +def test_transfer_retries_a_transient_failure(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + statuses = iter([1, 1, 0]) + attempts: list[list[str]] = [] + + def flaky(command: list[str], _log: Path) -> int: + attempts.append(command) + return next(statuses) + + monkeypatch.setattr(run_rec_colab, "stream_command", flaky) + + assert run_rec_colab.transfer_with_retry(["colab", "download"], tmp_path / "log") == 0 + assert len(attempts) == 3 + + +def test_remote_join_input_parts_round_trip(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + from training import colab_remote + + monkeypatch.setattr(colab_remote, "CONTENT", tmp_path) + monkeypatch.setattr(colab_remote, "INPUT_ARCHIVE", tmp_path / "in.tar.gz") + payload = b"0123456789" + for index in range(2): + (tmp_path / f"ocrkit-input.part{index:04d}").write_bytes(payload[index * 5 : index * 5 + 5]) + + colab_remote.join_input_parts() + + assert (tmp_path / "in.tar.gz").read_bytes() == payload + assert not list(tmp_path.glob("ocrkit-input.part*")) + + +def test_remote_upload_checkpoint_raises_on_transport_failure(tmp_path: Path) -> None: + from training import colab_remote + + checkpoint = tmp_path / "best_accuracy.pdparams" + checkpoint.write_bytes(b"weights") + + with pytest.raises(RuntimeError, match="uploading the checkpoint"): + colab_remote.upload_checkpoint( + {"bucket": "b", "key": "k", "url": "http://127.0.0.1:1/unreachable"}, checkpoint + ) diff --git a/training/README.md b/training/README.md index d6521c7..c0da217 100644 --- a/training/README.md +++ b/training/README.md @@ -398,10 +398,12 @@ Prepare the offline environment once, then run the CPU recognition Smoke: `run_rec_smoke.sh` accepts `--labels-dir`, `--output-dir`, `--epochs` (the target total epoch), and `--resume-checkpoint` (a checkpoint base path without -`.pdparams`, `.pdopt`, or `.states`). It validates both label files, fine-tunes -recognition only, and leaves `latest` plus `best_accuracy` under -`training/.work/`. Per-epoch `iter_epoch_*` dumps and PaddleOCR's duplicate -`best_model/` copy are pruned after training. To reclaim space from older runs: +`.pdparams`, `.pdopt`, or `.states`). `--device cpu|cuda` selects PaddleOCR's +device and defaults to `cpu`; `--train-only` stops after training and pruning +without running the evaluation. It validates both label files, fine-tunes recognition only, and +leaves `latest` plus `best_accuracy` under `training/.work/`. Per-epoch +`iter_epoch_*` dumps and PaddleOCR's duplicate `best_model/` copy are pruned +after training. To reclaim space from older runs: ```bash uv run python training/scripts/prune_rec_checkpoints.py --root training/.work @@ -419,6 +421,83 @@ The current release gate is field accuracy at least `364/379` (`0.9604221635883905`). A failed Smoke keeps the checkpoint and evaluation report for inspection but returns a non-zero status. +### Run recognition training on Colab GPU + +Install the [Google Colab CLI](https://github.com/googlecolab/google-colab-cli) +and prepare the reviewed/materialized dataset. `setup_rec_environment.sh` is +only needed locally if you also want CPU Smoke or a non-default (custom) base +checkpoint; the default official base checkpoint is fetched by Colab itself. + +```bash +uv tool install google-colab-cli +``` + +The first CLI request can prompt for Google OAuth authentication in the +terminal. The CLI keeps those credentials locally; OCRKit does not send +platform or release credentials to Colab. The runner also requires +`OCRKIT_R2_ENDPOINT_URL`, `OCRKIT_R2_ACCESS_KEY_ID`, `OCRKIT_R2_SECRET_ACCESS_KEY`, +and `OCRKIT_R2_DEFAULT_BUCKET` (see `.env.model.example`): the trained +checkpoint is far too large to transfer efficiently through the Colab CLI, so +Colab uploads it straight to that private R2 bucket using a short-lived, +single-object presigned URL that the runner generates locally and never +writes to disk or a log; Colab never receives R2 credentials. The runner +downloads the checkpoint from R2 and deletes the object once it has done so. + +Train on Colab and evaluate the retrieved checkpoint locally with one command: + +```bash +uv run python training/run_rec_colab.py +``` + +The default dataset is `datasets/labeled/rec`. To train from a materialized +platform snapshot, select that snapshot's output directory explicitly: + +```bash +uv run python training/run_rec_colab.py \ + --labels-dir datasets/labeled/rec/platform/@ \ + --gpu T4 \ + --epochs 10 +``` + +`--gpu` is a Colab allocation preference (default `T4`), not a model or +training requirement. PaddlePaddle's CUDA runtime and device are checked before +training; an unavailable or unsupported GPU request fails without falling back +to CPU. `--timeout-seconds` (default 6 hours) bounds the remote run; the Colab +CLI's own `exec` default of 30 seconds is always overridden. The local CPU command remains `./training/run_rec_smoke.sh`. + +The 2.9 GB `paddlepaddle-gpu` wheel is slow to fetch from the official +CDN outside China, so the runner installs a checksummed mirror of the official +cu129 build (`PADDLE_WHEEL` in `run_rec_colab.py`) when the runtime selects the +cu129 index, and otherwise falls back to the official index. + +Only training runs on Colab. CUDA builds of PaddlePaddle export `nn.Linear` as +`linear_v2`, which `paddle2onnx` cannot convert, so the runner retrieves the +checkpoint, stops the runtime, and then runs the unchanged +`training/evaluate_rec_checkpoint.sh` on your machine (the local training +environment from `setup_rec_environment.sh` is required). The run only +succeeds if that evaluation passes the same gate as a local run. + +Through the Colab CLI, the runner transfers only the selected train/holdout +labels and referenced crops, available review/snapshot provenance files, the +training scripts, and (when `--pretrained-checkpoint` names a checkpoint other +than the official default) that custom checkpoint. The official default base +checkpoint is instead fetched by Colab directly from its public URL and +checksum-verified there, and the trained checkpoint returns through R2 rather +than the CLI. It records source revisions, input checksums, the effective +training configuration, PaddleOCR revision, allocated GPU details, checkpoint +checksums, and the local evaluation summary in the returned `run.json`. + +The checkpoint, `fixture_report.json`, provenance, the remote training log, +the Colab CLI log, and `status.json` are stored below the ignored +`training/.work/colab-runs//` directory. A successful run does not +publish a candidate or change the stable model channel. Provisioning, +staging, training, metadata retrieval, and handled failures all stop the +Colab runtime after it has been allocated; it is stopped before the local +evaluation starts. Failed runs keep diagnostics and any partial output under +`partial/`; if teardown itself fails, `status.json` includes the named +`colab stop` command to release that session. The R2 checkpoint object is +deleted once retrieved, on both success and failure. + To evaluate a checkpoint explicitly, use a new output directory: ```bash diff --git a/training/bootstrap.sh b/training/bootstrap.sh index 02b1de5..09a029a 100644 --- a/training/bootstrap.sh +++ b/training/bootstrap.sh @@ -15,7 +15,15 @@ printf 'PaddleOCR checkout: %s\n' "${paddleocr_dir}" venv_dir="${work_dir}/venv" if [[ ! -x "${venv_dir}/bin/python" ]]; then - python3.12 -m venv "${venv_dir}" + if command -v python3.12 >/dev/null && python3.12 -m venv "${venv_dir}"; then + : + elif command -v uv >/dev/null; then + # Managed runtimes such as Colab ship python3.12 without python3-venv/ensurepip. + uv venv --clear --seed --python 3.12 "${venv_dir}" + else + printf 'Python 3.12 with venv support, or uv, is required to create the training environment.\n' >&2 + exit 1 + fi fi printf 'Training environment: %s\n' "${venv_dir}" diff --git a/training/colab_remote.py b/training/colab_remote.py new file mode 100644 index 0000000..43cbddf --- /dev/null +++ b/training/colab_remote.py @@ -0,0 +1,320 @@ +#!/usr/bin/env python3 +from __future__ import annotations + +import hashlib +import json +import os +import shlex +import shutil +import subprocess +import sys +import tarfile +import time +import traceback +import urllib.request +from datetime import UTC, datetime +from pathlib import Path +from typing import Any + +COLAB_ROOT = Path("/content/ocrkit-colab") +CONTENT = Path("/content") +INPUT_ARCHIVE = CONTENT / "ocrkit-input.tar.gz" +PART_BYTES = 32 * 1024 * 1024 +REPO = COLAB_ROOT / "repo" +DATASET = COLAB_ROOT / "dataset" +RESULTS = COLAB_ROOT / "results" +CHECKPOINTS = RESULTS / "checkpoint" +RUN_METADATA = RESULTS / "run.json" +REMOTE_LOG = RESULTS / "remote.log" +BASE_CHECKPOINT_PATH = REPO / "training/.work/pretrained/PP-OCRv6_small_rec_pretrained.pdparams" + + +def sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as stream: + for chunk in iter(lambda: stream.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +def safe_extract(archive_path: Path, destination: Path) -> None: + root = destination.resolve() + with tarfile.open(archive_path, "r:gz") as archive: + members = archive.getmembers() + for member in members: + target = (destination / member.name).resolve() + if not target.is_relative_to(root) or not (member.isfile() or member.isdir()): + raise RuntimeError("Colab input archive contains an unsupported path or file type") + archive.extractall(destination) + + +def run_logged(command: list[str], *, cwd: Path, log: Any, env: dict[str, str] | None = None) -> None: + rendered = shlex.join(command) + print(f"$ {rendered}", flush=True) + log.write(f"$ {rendered}\n") + log.flush() + started = time.monotonic() + with subprocess.Popen( + command, + cwd=cwd, + env=env, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + bufsize=1, + ) as process: + assert process.stdout is not None + for line in process.stdout: + print(line, end="", flush=True) + log.write(line) + status = process.wait() + elapsed = round(time.monotonic() - started, 1) + log.write(f"$ {rendered} -> exit {status} in {elapsed}s\n") + log.flush() + if status: + raise RuntimeError(f"remote command exited with status {status}: {rendered}") + + +def timed(stages: dict[str, float], name: str): + class _Timer: + def __enter__(self): + self._start = time.monotonic() + return self + + def __exit__(self, *_exc): + stages[name] = round(time.monotonic() - self._start, 1) + + return _Timer() + + +def checked_gpu() -> dict[str, Any]: + if not shutil.which("nvidia-smi"): + raise RuntimeError("Colab did not provide an NVIDIA GPU runtime") + result = subprocess.run( + [ + "nvidia-smi", + "--query-gpu=name,compute_cap,driver_version,memory.total", + "--format=csv,noheader,nounits", + ], + check=False, + capture_output=True, + text=True, + ) + if result.returncode or not result.stdout.strip(): + raise RuntimeError("Colab could not report an allocated NVIDIA GPU") + devices = [] + for line in result.stdout.splitlines(): + name, capability, driver, memory = (part.strip() for part in line.split(",", 3)) + devices.append( + { + "name": name, + "compute_capability": capability, + "driver_version": driver, + "memory_mib": int(memory), + } + ) + return {"devices": devices, "nvidia_smi": subprocess.run( + ["nvidia-smi"], check=False, capture_output=True, text=True + ).stdout} + + +def verify_inputs(request: dict[str, Any], root: Path) -> None: + for record in request["input_files"]: + path = root / record["path"] + if not path.is_file() or path.stat().st_size != record["size_bytes"] or sha256(path) != record["sha256"]: + raise RuntimeError(f"staged input failed checksum verification: {record['path']}") + + +def join_input_parts() -> None: + parts = sorted(CONTENT.glob("ocrkit-input.part*")) + if not parts: + raise RuntimeError("OCRKit input parts were not uploaded to the Colab runtime") + with INPUT_ARCHIVE.open("wb") as archive: + for part in parts: + archive.write(part.read_bytes()) + part.unlink() + + +def fetch_base_checkpoint(base_checkpoint: dict[str, Any]) -> None: + """A checkpoint the runner did not upload; Colab downloads it directly instead.""" + BASE_CHECKPOINT_PATH.parent.mkdir(parents=True, exist_ok=True) + urllib.request.urlretrieve(base_checkpoint["url"], BASE_CHECKPOINT_PATH) + if sha256(BASE_CHECKPOINT_PATH) != base_checkpoint["sha256"]: + raise RuntimeError("downloaded base checkpoint failed checksum verification") + + +def upload_checkpoint(upload: dict[str, Any], checkpoint_path: Path) -> dict[str, Any]: + """PUTs the trained checkpoint straight to R2; the Colab CLI's own transfer is far slower.""" + result = subprocess.run( + [ + "curl", + "-fsS", + "-X", + "PUT", + "--data-binary", + f"@{checkpoint_path}", + "-H", + "Content-Type: application/octet-stream", + "--max-time", + "1800", + "-o", + "/dev/null", + "-w", + "%{http_code}", + upload["url"], + ], + capture_output=True, + text=True, + ) + if result.returncode or result.stdout.strip() not in {"200", "201"}: + raise RuntimeError(f"uploading the checkpoint to R2 failed (curl exit {result.returncode}: {result.stderr.strip()[-300:]})") + return { + "bucket": upload["bucket"], + "key": upload["key"], + "size_bytes": checkpoint_path.stat().st_size, + "sha256": sha256(checkpoint_path), + } + + +def main() -> int: + RESULTS.mkdir(parents=True, exist_ok=True) + stages: dict[str, float] = {} + result: dict[str, Any] = { + "schema_version": 2, + "status": "failed", + "completed_at": datetime.now(UTC).isoformat(), + } + try: + with REMOTE_LOG.open("w", encoding="utf-8") as log: + try: + join_input_parts() + COLAB_ROOT.mkdir(parents=True, exist_ok=True) + safe_extract(INPUT_ARCHIVE, COLAB_ROOT) + request_path = COLAB_ROOT / "request.json" + request = json.loads(request_path.read_text(encoding="utf-8")) + run_request = request["run"] + # The presigned checkpoint-upload URL is a write credential; never persist or log it. + result["request"] = { + key: (value if key != "checkpoint_upload" else {k: v for k, v in value.items() if k != "url"}) + for key, value in run_request.items() + } + verify_inputs(request, COLAB_ROOT) + gpu = checked_gpu() + REPO.mkdir(parents=True, exist_ok=True) + DATASET.mkdir(parents=True, exist_ok=True) + CHECKPOINTS.mkdir(parents=True, exist_ok=True) + + if run_request["base_checkpoint"]["source"] == "official-download": + with timed(stages, "fetch_base_checkpoint"): + fetch_base_checkpoint(run_request["base_checkpoint"]) + elif not BASE_CHECKPOINT_PATH.is_file(): + raise RuntimeError("uploaded base checkpoint was not staged at the expected path") + + if not shutil.which("uv"): + run_logged([sys.executable, "-m", "pip", "install", "uv"], cwd=COLAB_ROOT, log=log) + mirror = run_request["paddle_wheel_mirror"] + with timed(stages, "setup_environment"): + run_logged( + ["bash", "training/setup_rec_environment.sh", "--device", "cuda"], + cwd=REPO, + env={**os.environ, "OCRKIT_PADDLE_WHEEL_URL": mirror["url"], "OCRKIT_PADDLE_WHEEL_SHA256": mirror["sha256"]}, + log=log, + ) + + paddle = subprocess.run( + [ + str(REPO / "training/.work/venv/bin/python"), + "-c", + "import json, paddle; paddle.device.set_device('gpu:0'); paddle.to_tensor([1.0]).numpy(); print(json.dumps({'version': paddle.__version__, 'cuda': paddle.is_compiled_with_cuda(), 'device': paddle.device.get_device()}))", + ], + check=True, + capture_output=True, + text=True, + ) + paddle_info = json.loads(paddle.stdout.strip().splitlines()[-1]) + if not paddle_info["cuda"] or paddle_info["device"] != "gpu:0": + raise RuntimeError("PaddlePaddle did not select the allocated CUDA device") + + epochs = str(run_request["training"]["epochs"]) + with timed(stages, "train"): + run_logged( + [ + "bash", + "training/run_rec_smoke.sh", + "--labels-dir", + str(DATASET), + "--output-dir", + str(CHECKPOINTS), + "--epochs", + epochs, + "--device", + "cuda", + "--train-only", + ], + cwd=REPO, + log=log, + ) + best_checkpoint = CHECKPOINTS / "best_accuracy.pdparams" + if not best_checkpoint.is_file() or best_checkpoint.stat().st_size == 0: + raise RuntimeError("training did not produce the best-accuracy recognition checkpoint") + + with timed(stages, "upload_checkpoint"): + checkpoint_record = upload_checkpoint(run_request["checkpoint_upload"], best_checkpoint) + # Recorded immediately: if a later step fails, the caller still knows this object + # exists in R2 and can retrieve or delete it instead of leaking it silently. + result["checkpoint"] = checkpoint_record + + paddleocr_revision = subprocess.run( + ["git", "-C", str(REPO / "training/.work/PaddleOCR"), "rev-parse", "HEAD"], + check=True, + capture_output=True, + text=True, + ).stdout.strip() + config_path = CHECKPOINTS / "config.yml" + train_log_path = CHECKPOINTS / "train.log" + result.update( + { + "status": "success", + "runtime": { + **gpu, + "python": sys.version, + "paddle": paddle_info, + "paddleocr_revision": paddleocr_revision, + }, + "stage_seconds": stages, + "checkpoint": checkpoint_record, + "checkpoint_config": config_path.read_text(encoding="utf-8") if config_path.is_file() else None, + "train_log": train_log_path.read_text(encoding="utf-8", errors="replace") if train_log_path.is_file() else None, + } + ) + except BaseException as exc: + result["error"] = f"{type(exc).__name__}: {exc}" + result["stage_seconds"] = stages + traceback.print_exc(file=log) + log.flush() + finally: + result["completed_at"] = datetime.now(UTC).isoformat() + RUN_METADATA.write_text(json.dumps(result, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") + except BaseException as exc: + print(f"OCRKit Colab run failed: {type(exc).__name__}: {exc}", file=sys.stderr, flush=True) + try: + RUN_METADATA.write_text( + json.dumps( + {"schema_version": 2, "status": "failed", "error": f"{type(exc).__name__}: {exc}", "stage_seconds": stages}, + ensure_ascii=False, + indent=2, + ) + + "\n", + encoding="utf-8", + ) + except OSError: + pass + return 1 + if result.get("status") != "success": + print(f"OCRKit Colab run failed: {result.get('error', 'remote step failed')}", file=sys.stderr, flush=True) + return 1 + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/training/run_rec_colab.py b/training/run_rec_colab.py new file mode 100644 index 0000000..18e7f48 --- /dev/null +++ b/training/run_rec_colab.py @@ -0,0 +1,559 @@ +#!/usr/bin/env python3 +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import shlex +import shutil +import subprocess +import sys +import tarfile +from datetime import UTC, datetime +from pathlib import Path, PurePosixPath, PureWindowsPath +from typing import Any + +ROOT = Path(__file__).resolve().parents[1] +if str(ROOT) not in sys.path: + # Running this file directly (python training/run_rec_colab.py) puts training/, not the + # repo root, on sys.path; the repo's own `app` package needs the root added explicitly. + sys.path.insert(0, str(ROOT)) + +from app.core.config import settings # noqa: E402 +from app.storage.r2_client import ObjectNotFoundError, R2ObjectStore # noqa: E402 +RUNS = ROOT / "training/.work/colab-runs" +ACCEPTED_STAGING = "accepted" +PADDLE_WHEEL = { + "url": "https://cdn.owbastion.codes/ocrkit/wheels/cu129/paddlepaddle_gpu-3.3.1-cp312-cp312-linux_x86_64.whl", + "sha256": "03fc5211183ba20ef71e63a35e589fd386a8805c1011c3a409a0e5b118be4668", +} +OFFICIAL_BASE_CHECKPOINT = { + "model": "PP-OCRv6_small_rec_pretrained", + "url": "https://paddle-model-ecology.bj.bcebos.com/paddlex/official_pretrained_model/PP-OCRv6_small_rec_pretrained.pdparams", + "sha256": "25c9bd54b0e5900916e8bb6ada938abeffb1eac1baedac0ca54a45b1c9310825", +} +R2_KEY_PREFIX = "colab-runs" +R2_UPLOAD_URL_BUFFER_SECONDS = 900 +PART_BYTES = 32 * 1024 * 1024 +REMOTE_RUNNER = ROOT / "training/colab_remote.py" +PRETRAINED_CHECKPOINT = ROOT / "training/.work/pretrained/PP-OCRv6_small_rec_pretrained.pdparams" +SOURCE_FILES = ( + "training/bootstrap.sh", + "training/setup_rec_environment.sh", + "training/run_rec_smoke.sh", + "training/configs/rec_pp_ocrv6_small.yaml", + "training/scripts/prune_rec_checkpoints.py", + "training/scripts/validate_annotations.py", +) + + +def require_r2_store(parser: argparse.ArgumentParser) -> R2ObjectStore: + """The Colab backend retrieves checkpoints through R2 rather than the slow Colab CLI transfer.""" + if not ( + settings.r2_endpoint_url + and settings.r2_access_key_id + and settings.r2_secret_access_key + and settings.r2_default_bucket + ): + parser.error( + "OCRKIT_R2_ENDPOINT_URL, OCRKIT_R2_ACCESS_KEY_ID, OCRKIT_R2_SECRET_ACCESS_KEY, and " + "OCRKIT_R2_DEFAULT_BUCKET are required to run training on Colab" + ) + return R2ObjectStore.from_settings( + endpoint_url=settings.r2_endpoint_url, + access_key_id=settings.r2_access_key_id, + secret_access_key=settings.r2_secret_access_key, + region_name=settings.r2_region_name, + default_bucket=settings.r2_default_bucket, + allowed_buckets_raw=settings.r2_allowed_buckets, + read_timeout_seconds=settings.r2_read_timeout_seconds, + ) + + +def sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as stream: + for chunk in iter(lambda: stream.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +def add_file(files: dict[str, tuple[Path, str]], archive_path: str, source: Path, kind: str) -> None: + if source.is_symlink() or not source.is_file(): + raise ValueError(f"required Colab input is missing or is a symbolic link: {source}") + if archive_path in files: + if files[archive_path][0] != source: + raise ValueError(f"two Colab inputs map to the same path: {archive_path}") + return + files[archive_path] = (source, kind) + + +def safe_relative(value: str) -> Path: + path = Path(value) + if ( + not value + or "\\" in value + or path.is_absolute() + or PureWindowsPath(value).is_absolute() + or any(part in ("", ".", "..") for part in PurePosixPath(value).parts) + ): + raise ValueError(f"training input contains an unsafe image path: {value!r}") + return path + + +def validate_labels(label_path: Path) -> tuple[int, list[str]]: + result = subprocess.run( + [sys.executable, str(ROOT / "training/scripts/validate_annotations.py"), "rec", str(label_path)], + cwd=ROOT, + check=True, + capture_output=True, + text=True, + ) + sample_count = int(json.loads(result.stdout)["valid_samples"]) + if sample_count < 1: + raise ValueError(f"recognition split is empty: {label_path.name}") + images = [] + for line in label_path.read_text(encoding="utf-8").splitlines(): + if not line.strip(): + continue + image_name = line.split("\t", 1)[0] + relative = safe_relative(image_name) + candidates = (label_path.parent / "images" / relative, label_path.parent.parent / relative) + image_path = next((candidate for candidate in candidates if candidate.is_file()), None) + if image_path is None: + raise ValueError(f"validated recognition image disappeared: {image_name}") + image_path = image_path.resolve() + dataset_root = label_path.parent.parent.resolve() + if not image_path.is_relative_to(dataset_root): + raise ValueError(f"recognition image resolves outside the selected dataset: {image_name}") + images.append(image_name) + return sample_count, images + + +def git_metadata(path: Path) -> tuple[str | None, bool]: + revision = subprocess.run( + ["git", "-C", str(path), "rev-parse", "HEAD"], + capture_output=True, + text=True, + check=False, + ) + if revision.returncode: + return None, False + status = subprocess.run( + ["git", "-C", str(path), "status", "--porcelain", "--untracked-files=normal"], + capture_output=True, + text=True, + check=False, + ) + return revision.stdout.strip(), bool(status.stdout.strip()) + + +def read_dataset_provenance(dataset_root: Path) -> dict[str, Any]: + for name in ("provenance.json", "snapshot.json"): + path = dataset_root / name + if not path.is_file(): + continue + data = json.loads(path.read_text(encoding="utf-8")) + snapshot = data.get("snapshot", {}) + snapshot_id = snapshot.get("snapshot_id") or data.get("snapshot_id") + if snapshot_id: + return { + "snapshot_id": snapshot_id, + "snapshot_version": snapshot.get("version") or data.get("snapshot_version") or data.get("version"), + "code_revision": data.get("code_revision"), + } + return {} + + +def source_files(root: Path, files: dict[str, tuple[Path, str]]) -> None: + for relative in SOURCE_FILES: + add_file(files, f"repo/{relative}", root / relative, "ocrkit-source") + + +def stage_inputs( + archive_path: Path, + run_request: dict[str, Any], + dataset_root: Path, + checkpoint_path: Path | None, + train_images: list[str], + holdout_images: list[str], +) -> None: + files: dict[str, tuple[Path, str]] = {} + source_files(ROOT, files) + + for relative in ("labels/train.txt", "labels/holdout.txt"): + add_file(files, f"dataset/{relative}", dataset_root / relative, "reviewed-labels") + image_names = set(train_images + holdout_images) + for image_name in sorted(image_names): + relative = safe_relative(image_name) + label_parent = dataset_root / "labels" + candidates = (label_parent / "images" / relative, dataset_root / relative) + image_path = next((candidate for candidate in candidates if candidate.is_file()), None) + if image_path is None: + raise ValueError(f"training crop is missing: {image_name}") + add_file(files, f"dataset/{relative.as_posix()}", image_path, "reviewed-crop") + + for relative in ( + "provenance.json", + "snapshot.json", + "crop_manifest.json", + "review/train.jsonl", + "review/holdout.jsonl", + ): + path = dataset_root / relative + if path.is_file(): + add_file(files, f"dataset/{relative}", path, "dataset-provenance") + if checkpoint_path is not None: + add_file( + files, + "repo/training/.work/pretrained/PP-OCRv6_small_rec_pretrained.pdparams", + checkpoint_path, + "base-recognition-checkpoint", + ) + + records = [] + with tarfile.open(archive_path, "w:gz") as archive: + for path, (source, kind) in sorted(files.items()): + archive.add(source, arcname=path, recursive=False) + records.append( + { + "path": path, + "kind": kind, + "size_bytes": source.stat().st_size, + "sha256": sha256(source), + } + ) + request = {"run": run_request, "input_files": records} + request_path = archive_path.parent / "request.json" + request_path.write_text(json.dumps(request, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") + archive.add(request_path, arcname="request.json", recursive=False) + + +def stream_command(command: list[str], log_path: Path) -> int: + with log_path.open("a", encoding="utf-8") as log: + rendered = shlex.join(command) + print(f"$ {rendered}", flush=True) + log.write(f"$ {rendered}\n") + with subprocess.Popen( + command, + cwd=ROOT, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + bufsize=1, + ) as process: + assert process.stdout is not None + try: + for line in process.stdout: + print(line, end="", flush=True) + log.write(line) + except KeyboardInterrupt: + process.terminate() + process.wait() + raise + return process.wait() + + +def capture_command(command: list[str]) -> subprocess.CompletedProcess[str]: + return subprocess.run(command, cwd=ROOT, capture_output=True, text=True) + + +def session_still_active(colab: str, session: str, log_path: Path) -> bool: + """`colab stop` can fail simply because Colab already reclaimed an idle/finished runtime.""" + result = capture_command([colab, "sessions"]) + with log_path.open("a", encoding="utf-8") as log: + log.write(f"$ {colab} sessions\n{result.stdout}{result.stderr}\n") + if result.returncode: + return True # could not verify; assume active so a real leak is not silently dropped + return session in result.stdout + + +def transfer_with_retry(command: list[str], log_path: Path, attempts: int = 3) -> int: + status = 1 + for _ in range(attempts): + status = stream_command(command, log_path) + if not status: + break + return status + + +def upload_archive(colab: str, session: str, archive: Path, run_dir: Path, log_path: Path) -> int: + """Colab uploads are single JSON requests, so send the archive in bounded parts.""" + parts_dir = run_dir / "input-parts" + parts_dir.mkdir() + try: + with archive.open("rb") as stream: + for index, chunk in enumerate(iter(lambda: stream.read(PART_BYTES), b"")): + part = parts_dir / f"part{index:04d}" + part.write_bytes(chunk) + status = transfer_with_retry( + [colab, "upload", "-s", session, str(part), f"/content/ocrkit-input.part{index:04d}"], + log_path, + ) + part.unlink() + if status: + return status + return 0 + finally: + shutil.rmtree(parts_dir, ignore_errors=True) + + +def fetch_remote_metadata(colab: str, session: str, run_dir: Path, log_path: Path) -> tuple[dict[str, Any], Path | None]: + """Only run.json and remote.log travel through the Colab CLI; both are small, unlike the checkpoint.""" + metadata_path = run_dir / "remote-run.json" + status = transfer_with_retry( + [colab, "download", "-s", session, "/content/ocrkit-colab/results/run.json", str(metadata_path)], log_path + ) + if status: + raise RuntimeError("Colab did not return run metadata (results/run.json)") + remote_log_path = run_dir / "remote-log.txt" + if transfer_with_retry( + [colab, "download", "-s", session, "/content/ocrkit-colab/results/remote.log", str(remote_log_path)], log_path + ): + remote_log_path = None # best-effort: diagnostics, not required for the run's outcome + return json.loads(metadata_path.read_text(encoding="utf-8")), remote_log_path + + +def retrieve_checkpoint( + r2: R2ObjectStore, remote_metadata: dict[str, Any], remote_log_path: Path | None, run_dir: Path +) -> None: + """The checkpoint travels through R2 (uploaded by colab_remote.py), not the Colab CLI.""" + accepted = run_dir / ACCEPTED_STAGING + accepted.mkdir(parents=True, exist_ok=True) + (accepted / "run.json").write_text( + json.dumps(remote_metadata, ensure_ascii=False, indent=2) + "\n", encoding="utf-8" + ) + if remote_log_path is not None: + shutil.move(str(remote_log_path), accepted / "remote.log") + + checkpoint = remote_metadata.get("checkpoint") + if checkpoint is None: + return + checkpoint_dir = accepted / "checkpoint" + checkpoint_dir.mkdir(parents=True, exist_ok=True) + destination = checkpoint_dir / "best_accuracy.pdparams" + try: + try: + r2.download_object(checkpoint["bucket"], checkpoint["key"], destination) + except ObjectNotFoundError as exc: + raise ValueError( + f"Colab did not upload the checkpoint to R2: {checkpoint['bucket']}/{checkpoint['key']}" + ) from exc + if destination.stat().st_size != checkpoint["size_bytes"] or sha256(destination) != checkpoint["sha256"]: + raise ValueError("Colab checkpoint failed checksum verification after download from R2") + config_text = remote_metadata.get("checkpoint_config") + if config_text: + (checkpoint_dir / "config.yml").write_text(config_text, encoding="utf-8") + train_log_text = remote_metadata.get("train_log") + if train_log_text: + (checkpoint_dir / "train.log").write_text(train_log_text, encoding="utf-8") + finally: + try: + r2.delete_object(checkpoint["bucket"], checkpoint["key"]) + except Exception as exc: + print( + f"warning: failed to delete the Colab checkpoint object from R2 " + f"({checkpoint['bucket']}/{checkpoint['key']}): {exc}", + file=sys.stderr, + ) + + +def evaluate_locally(accepted: Path, log_path: Path) -> None: + """Run the existing local evaluation contract on the retrieved checkpoint.""" + evaluation = accepted / "evaluation" + status = stream_command( + [ + str(ROOT / "training/evaluate_rec_checkpoint.sh"), + str((accepted / "checkpoint/best_accuracy").resolve()), + str(evaluation.resolve()), + ], + log_path, + ) + if status: + raise RuntimeError(f"local checkpoint evaluation failed with exit status {status}") + report = json.loads((evaluation / "fixture_report.json").read_text(encoding="utf-8")) + metadata = json.loads((accepted / "run.json").read_text(encoding="utf-8")) + metadata["evaluation"] = { + "location": "local", + "report": "evaluation/fixture_report.json", + "field_accuracy": report.get("field_accuracy"), + "run_code_accuracy": report.get("run_code", {}).get("field_accuracy"), + } + (accepted / "run.json").write_text(json.dumps(metadata, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") + + +def main() -> int: + parser = argparse.ArgumentParser(description="Train the OCRKit recognition model on a Colab GPU, then evaluate it locally.") + parser.add_argument("--labels-dir", type=Path, default=ROOT / "datasets/labeled/rec") + parser.add_argument("--pretrained-checkpoint", type=Path, default=PRETRAINED_CHECKPOINT) + parser.add_argument("--epochs", type=int, default=10) + parser.add_argument("--gpu", default="T4", help="Colab GPU preference; no accelerator fallback is attempted") + parser.add_argument("--timeout-seconds", type=float, default=6 * 3600, help="Upper bound for the remote training run") + args = parser.parse_args() + if args.epochs < 1: + parser.error("--epochs must be a positive integer") + colab = shutil.which("colab") + if not colab: + parser.error("Google Colab CLI is required; install it with uv tool install google-colab-cli") + + r2 = require_r2_store(parser) + + dataset_root = args.labels_dir.resolve() + train_label = dataset_root / "labels/train.txt" + holdout_label = dataset_root / "labels/holdout.txt" + for path in (train_label, holdout_label): + if not path.is_file(): + parser.error(f"required training input is missing: {path}") + using_official_checkpoint = args.pretrained_checkpoint == PRETRAINED_CHECKPOINT + if not using_official_checkpoint and not args.pretrained_checkpoint.is_file(): + parser.error(f"required training input is missing: {args.pretrained_checkpoint}") + train_count, train_images = validate_labels(train_label) + holdout_count, holdout_images = validate_labels(holdout_label) + + run_id = datetime.now(UTC).strftime("%Y%m%dT%H%M%SZ") + f"-{os.getpid():x}" + run_dir = RUNS / run_id + run_dir.mkdir(parents=True, exist_ok=False) + input_archive = run_dir / "ocrkit-input.tar.gz" + colab_log = run_dir / "colab.log" + + source_revision, source_dirty = git_metadata(ROOT) + dataset_revision, dataset_dirty = git_metadata(dataset_root) + try: + dataset_source = dataset_root.relative_to(ROOT).as_posix() + except ValueError: + dataset_source = "external" + if using_official_checkpoint: + checkpoint_path = None + base_checkpoint = {**OFFICIAL_BASE_CHECKPOINT, "source": "official-download"} + else: + checkpoint_path = args.pretrained_checkpoint.resolve() + base_checkpoint = { + "model": "custom", + "source": "uploaded", + "source_name": checkpoint_path.name, + "sha256": sha256(checkpoint_path), + } + checkpoint_upload_key = f"{R2_KEY_PREFIX}/{run_id}/checkpoint.pdparams" + checkpoint_upload_url = r2.generate_presigned_put_url( + r2.default_bucket, checkpoint_upload_key, expires_in_seconds=int(args.timeout_seconds) + R2_UPLOAD_URL_BUFFER_SECONDS + ) + run_request = { + "run_id": run_id, + "ocrkit_revision": source_revision, + "ocrkit_worktree_dirty": source_dirty, + "dataset": { + **read_dataset_provenance(dataset_root), + "source": dataset_source, + "revision": dataset_revision, + "worktree_dirty": dataset_dirty, + "train_samples": train_count, + "holdout_samples": holdout_count, + }, + "base_checkpoint": base_checkpoint, + "paddle_wheel_mirror": PADDLE_WHEEL, + "checkpoint_upload": {"bucket": r2.default_bucket, "key": checkpoint_upload_key, "url": checkpoint_upload_url}, + "training": { + "epochs": args.epochs, + "device": "cuda", + "gpu_preference": args.gpu, + "recipe": "training/configs/rec_pp_ocrv6_small.yaml", + "paddleocr_recipe": "configs/rec/PP-OCRv6/PP-OCRv6_small_rec.yml", + }, + } + try: + stage_inputs(input_archive, run_request, dataset_root, checkpoint_path, train_images, holdout_images) + except Exception as exc: + input_archive.unlink(missing_ok=True) + (run_dir / "status.json").write_text( + json.dumps({"status": "failed", "error": f"{type(exc).__name__}: {exc}"}, indent=2) + "\n", + encoding="utf-8", + ) + raise + + session = f"ocrkit-rec-{run_id.lower()}" + session_attempted = False + remote_metadata: dict[str, Any] = {} + error: str | None = None + stop_status: int | None = None + exec_status: int | None = None + + try: + session_attempted = True + provision_status = stream_command([colab, "new", "-s", session, "--gpu", args.gpu], colab_log) + if provision_status: + raise RuntimeError( + f"Colab could not provision the requested GPU {args.gpu}; no fallback accelerator was selected." + ) + upload_status = upload_archive(colab, session, input_archive, run_dir, colab_log) + input_archive.unlink(missing_ok=True) + if upload_status: + raise RuntimeError("Colab failed to stage OCRKit training inputs") + exec_status = stream_command( + [colab, "exec", "-s", session, "--timeout", str(args.timeout_seconds), "-f", str(REMOTE_RUNNER)], + colab_log, + ) + remote_metadata, remote_log_path = fetch_remote_metadata(colab, session, run_dir, colab_log) + retrieve_checkpoint(r2, remote_metadata, remote_log_path, run_dir) + if exec_status: + raise RuntimeError(f"Colab training failed with exit status {exec_status}.") + if remote_metadata.get("status") != "success": + raise RuntimeError( + f"Colab training did not return a successful status: {remote_metadata.get('error', 'unknown error')}" + ) + except KeyboardInterrupt: + error = "Colab run interrupted by the operator." + except Exception as exc: + error = f"{type(exc).__name__}: {exc}" + finally: + input_archive.unlink(missing_ok=True) + if session_attempted: + try: + stop_status = stream_command([colab, "stop", "-s", session], colab_log) + except Exception as exc: + stop_status = -1 + stop_error = f"Colab runtime teardown command failed: {type(exc).__name__}: {exc}" + error = f"{error}; {stop_error}" if error else stop_error + if stop_status and not session_still_active(colab, session, colab_log): + # Colab had already released the runtime on its own; nothing was left running. + stop_status = 0 + if stop_status and error is None: + error = f"training completed, but Colab runtime teardown failed; run colab stop -s {session}" + + accepted = run_dir / ACCEPTED_STAGING + if error is None and stop_status in (None, 0) and accepted.is_dir(): + try: + evaluate_locally(accepted, colab_log) + except KeyboardInterrupt: + error = "Local checkpoint evaluation interrupted by the operator." + except Exception as exc: + error = f"{type(exc).__name__}: {exc}" + succeeded = error is None and stop_status in (None, 0) + if accepted.is_dir(): + if succeeded: + for child in accepted.iterdir(): + child.rename(run_dir / child.name) + accepted.rmdir() + else: + partial = run_dir / "partial" + partial.mkdir(exist_ok=True) + for child in accepted.iterdir(): + child.rename(partial / child.name) + accepted.rmdir() + status = { + "status": "success" if succeeded else "failed", + "runtime_stopped": stop_status == 0, + "colab_session": session, + "error": error, + } + (run_dir / "status.json").write_text(json.dumps(status, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") + if status["status"] != "success": + print(f"Colab OCRKit run failed: {error or 'runtime did not stop successfully'}", file=sys.stderr) + print(f"Local run logs and any partial outputs: {run_dir}", file=sys.stderr) + return 1 + print(f"Colab OCRKit run completed. Checkpoint, local evaluation, provenance, and logs: {run_dir}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/training/run_rec_smoke.sh b/training/run_rec_smoke.sh index 887f35e..015ae6f 100755 --- a/training/run_rec_smoke.sh +++ b/training/run_rec_smoke.sh @@ -10,6 +10,8 @@ config_path="${paddleocr_dir}/configs/rec/PP-OCRv6/PP-OCRv6_small_rec.yml" output_dir="${work_dir}/checkpoints/rec_pp_ocrv6_small" epoch_num=10 resume_checkpoint="" +device=cpu +train_only=false while [[ $# -gt 0 ]]; do case "$1" in @@ -29,13 +31,30 @@ while [[ $# -gt 0 ]]; do resume_checkpoint="$2" shift 2 ;; + --device) + device="$2" + shift 2 + ;; + --train-only) + train_only=true + shift + ;; *) - printf 'usage: %s [--labels-dir ] [--output-dir ] [--epochs ] [--resume-checkpoint ]\n' "$0" >&2 + printf 'usage: %s [--labels-dir ] [--output-dir ] [--epochs ] [--resume-checkpoint ] [--device cpu|cuda] [--train-only]\n' "$0" >&2 exit 2 ;; esac done +case "$device" in + cpu) use_gpu=False ;; + cuda) use_gpu=True ;; + *) + printf 'device must be cpu or cuda\n' >&2 + exit 2 + ;; +esac + if ! [[ "${epoch_num}" =~ ^[1-9][0-9]*$ ]]; then printf 'epochs must be a positive integer\n' >&2 exit 2 @@ -64,7 +83,7 @@ fi cd "${paddleocr_dir}" "${python_bin}" tools/train.py -c "${config_path}" -o \ - Global.use_gpu=False \ + "Global.use_gpu=${use_gpu}" \ "Global.epoch_num=${epoch_num}" \ Global.save_model_dir="${output_dir}" \ Global.save_epoch_step="$((epoch_num + 1))" \ @@ -85,6 +104,10 @@ cd "${paddleocr_dir}" cd "${root_dir}" "${python_bin}" "${root_dir}/training/scripts/prune_rec_checkpoints.py" "${output_dir}" +if [[ "${train_only}" == true ]]; then + exit 0 +fi + evaluation_dir="${work_dir}/evaluations/rec_pp_ocrv6_small/$(date -u +%Y.%m.%d-%H%M%S)-$$" "${root_dir}/training/evaluate_rec_checkpoint.sh" \ "${output_dir}/best_accuracy" \ diff --git a/training/setup_rec_environment.sh b/training/setup_rec_environment.sh index 5bbb089..bd89bd5 100755 --- a/training/setup_rec_environment.sh +++ b/training/setup_rec_environment.sh @@ -8,9 +8,22 @@ python_bin="${work_dir}/venv/bin/python" pretrained_dir="${work_dir}/pretrained" pretrained_model="${pretrained_dir}/PP-OCRv6_small_rec_pretrained.pdparams" +device=cpu +if [[ $# -gt 0 ]]; then + if [[ $# -ne 2 || "$1" != "--device" ]]; then + printf 'usage: %s [--device cpu|cuda]\n' "$0" >&2 + exit 2 + fi + device="$2" +fi +if [[ "${device}" != cpu && "${device}" != cuda ]]; then + printf 'device must be cpu or cuda\n' >&2 + exit 2 +fi + bash "${root_dir}/training/bootstrap.sh" -if [[ "$(uname -s)" == "Darwin" ]]; then +if [[ "${device}" == cpu && "$(uname -s)" == "Darwin" ]]; then if ! command -v brew >/dev/null; then printf 'Homebrew is required to install ccache on macOS. Install Homebrew, then rerun this script.\n' >&2 exit 1 @@ -24,15 +37,61 @@ if [[ "$(uname -s)" == "Darwin" ]]; then fi fi +paddle_package="paddlepaddle==3.3.1" +paddle_index="https://www.paddlepaddle.org.cn/packages/stable/cpu/" +if [[ "${device}" == cuda ]]; then + if [[ "$(uname -s)" != Linux ]] || ! command -v nvidia-smi >/dev/null; then + printf 'CUDA training requires a Linux runtime with an NVIDIA GPU and nvidia-smi.\n' >&2 + exit 1 + fi + + cuda_version="$(nvidia-smi 2>/dev/null | awk -F 'CUDA Version: ' 'NF > 1 {split($2, version, " "); print version[1]; exit}')" + if [[ -z "${cuda_version}" ]]; then + printf 'could not determine the NVIDIA driver CUDA version from nvidia-smi.\n' >&2 + exit 1 + fi + cuda_major="$(cut -d. -f1 <<< "${cuda_version}")" + cuda_minor="$(cut -d. -f2 <<< "${cuda_version}")" + if (( cuda_major > 12 || (cuda_major == 12 && cuda_minor >= 9) )); then + paddle_index="https://www.paddlepaddle.org.cn/packages/stable/cu129/" + elif (( cuda_major == 12 && cuda_minor >= 6 )); then + paddle_index="https://www.paddlepaddle.org.cn/packages/stable/cu126/" + elif (( cuda_major == 11 && cuda_minor >= 8 )) || (( cuda_major == 12 )); then + paddle_index="https://www.paddlepaddle.org.cn/packages/stable/cu118/" + else + printf 'NVIDIA CUDA %s is unsupported; PaddlePaddle 3.3.1 needs CUDA 11.8 or newer.\n' "${cuda_version}" >&2 + exit 1 + fi + + paddle_package="paddlepaddle-gpu==3.3.1" +fi + "${python_bin}" -m pip install --upgrade pip -"${python_bin}" -m pip install "paddlepaddle==3.3.1" -i https://www.paddlepaddle.org.cn/packages/stable/cpu/ +if [[ "${device}" == cuda && "${paddle_index}" == */cu129/ && -n "${OCRKIT_PADDLE_WHEEL_URL:-}" && -n "${OCRKIT_PADDLE_WHEEL_SHA256:-}" ]]; then + # Optional mirror of the official cu129 wheel for runtimes far from the official CDN. + mkdir -p "${work_dir}/wheels" + wheel_path="${work_dir}/wheels/${OCRKIT_PADDLE_WHEEL_URL##*/}" + curl -L --fail --retry 3 -o "${wheel_path}" "${OCRKIT_PADDLE_WHEEL_URL}" + wheel_sha256="$("${python_bin}" -c 'import hashlib, sys; print(hashlib.sha256(open(sys.argv[1], "rb").read()).hexdigest())' "${wheel_path}")" + if [[ "${wheel_sha256}" != "${OCRKIT_PADDLE_WHEEL_SHA256}" ]]; then + printf 'mirrored PaddlePaddle wheel failed checksum verification: %s\n' "${wheel_path}" >&2 + exit 1 + fi + "${python_bin}" -m pip install "${wheel_path}" +else + "${python_bin}" -m pip install "${paddle_package}" -i "${paddle_index}" +fi "${python_bin}" -m pip install -r "${paddleocr_dir}/requirements.txt" "${python_bin}" -m pip install "paddle2onnx==2.1.0" -if [[ "$(uname -s)" == "Darwin" ]]; then +if [[ "${device}" == cpu && "$(uname -s)" == "Darwin" ]]; then "${python_bin}" -c 'import platform; import paddle; assert platform.machine() == "arm64"; assert paddle.device.get_device() == "cpu"; assert not paddle.is_compiled_with_cuda(); print(f"PaddlePaddle {paddle.__version__}: {platform.machine()} {paddle.device.get_device()}")' fi +if [[ "${device}" == cuda ]]; then + "${python_bin}" -c 'import paddle; assert paddle.is_compiled_with_cuda(), "installed PaddlePaddle is not CUDA-enabled"; assert paddle.device.cuda.device_count() > 0, "no CUDA device is visible to PaddlePaddle"; paddle.device.set_device("gpu:0"); assert paddle.to_tensor([1.0]).numpy()[0] == 1.0; print(f"PaddlePaddle {paddle.__version__}: CUDA {paddle.device.get_device()}")' +fi + mkdir -p "${pretrained_dir}" if [[ ! -f "${pretrained_model}" ]]; then "${python_bin}" -c "from urllib.request import urlretrieve; urlretrieve('https://paddle-model-ecology.bj.bcebos.com/paddlex/official_pretrained_model/PP-OCRv6_small_rec_pretrained.pdparams', '${pretrained_model}')"