Codex_Extractor / smoke_signal /scripts /09_finetune_surya.py
Pointf5ive
Training phase: scripts 09+10, --model flag in 03, gold schema patch
5ffc645
Raw
History Blame Contribute Delete
21.2 kB
#!/usr/bin/env python3
"""
Smoke Signal — Stage 11: Surya Fine-Tuning Orchestrator
========================================================
Builds a governance-safe OCR fine-tuning dataset from gold corrections,
optionally uploads it to Hugging Face, and can launch Surya OCR finetuning.
Governance controls enforced:
- unknown/excluded rights are always blocked
- mixed rights classes are blocked by default
- checkpoint/config changes are versioned and logged
- every run logs model version, dataset size, and gold set hash
Primary sources used for integration choices:
- Surya README finetune entrypoint and args
- Surya example dataset shape (`image` + `text`)
"""
import argparse
import csv
import hashlib
import json
import subprocess
import sys
from datetime import datetime, timezone
from pathlib import Path
from typing import Dict, List, Optional, Tuple
ROOT = Path(__file__).resolve().parents[1]
MANIFEST_CSV = ROOT / "manifest" / "source_manifest.csv"
RUN_LOG_CSV = ROOT / "manifest" / "run_log.csv"
GOLD_FILE = ROOT / "gold" / "gold_corrections.jsonl"
RENDERS_DIR = ROOT / "renders"
TRAINING_DIR = ROOT / "training"
TRAINING_DATASETS_DIR = TRAINING_DIR / "datasets"
TRAINING_RUNS_DIR = TRAINING_DIR / "runs"
ELIGIBLE_RIGHTS = {"public-domain", "licensed-owned", "controlled-internal"}
BLOCKED_RIGHTS = {"unknown", "excluded"}
RUN_LOG_FIELDS = [
"run_id",
"date",
"operator",
"config_version",
"schema_version",
"source_batch",
"pages_processed",
"errors",
"cost_usd",
"output_path",
"notes",
]
def utc_now() -> datetime:
return datetime.now(timezone.utc)
def utc_iso() -> str:
return utc_now().isoformat().replace("+00:00", "Z")
def ensure_dirs() -> None:
TRAINING_DIR.mkdir(parents=True, exist_ok=True)
TRAINING_DATASETS_DIR.mkdir(parents=True, exist_ok=True)
TRAINING_RUNS_DIR.mkdir(parents=True, exist_ok=True)
def sha256_file(path: Path) -> str:
h = hashlib.sha256()
with open(path, "rb") as f:
for block in iter(lambda: f.read(1 << 20), b""):
h.update(block)
return h.hexdigest()
def sha256_text(text: str) -> str:
return hashlib.sha256(text.encode("utf-8")).hexdigest()
def ensure_run_log() -> None:
RUN_LOG_CSV.parent.mkdir(parents=True, exist_ok=True)
if RUN_LOG_CSV.exists():
return
with open(RUN_LOG_CSV, "w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=RUN_LOG_FIELDS)
writer.writeheader()
def append_run_log(row: Dict[str, str]) -> None:
ensure_run_log()
with open(RUN_LOG_CSV, "a", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=RUN_LOG_FIELDS)
writer.writerow({k: row.get(k, "") for k in RUN_LOG_FIELDS})
def load_manifest() -> Dict[str, Dict[str, str]]:
records: Dict[str, Dict[str, str]] = {}
if not MANIFEST_CSV.exists():
return records
with open(MANIFEST_CSV, newline="", encoding="utf-8") as f:
for row in csv.DictReader(f):
book_id = str(row.get("book_id", "")).strip()
if book_id:
records[book_id] = row
return records
def load_gold_records(path: Path) -> List[Dict]:
records: List[Dict] = []
if not path.exists():
return records
with open(path, encoding="utf-8") as f:
for idx, line in enumerate(f, start=1):
line = line.strip()
if not line:
continue
try:
rec = json.loads(line)
rec["_gold_line"] = idx
records.append(rec)
except json.JSONDecodeError:
# Keep pipeline resilient: skip malformed line.
continue
return records
def parse_page_number(value) -> Optional[int]:
if value is None:
return None
try:
return int(value)
except (TypeError, ValueError):
return None
def resolve_image_path(record: Dict, book_id: str, page_num: Optional[int]) -> Optional[Path]:
path_fields = [
"page_image_path",
"page_image",
"image_path",
"crop_path",
"render_path",
]
for field in path_fields:
raw = record.get(field)
if not raw:
continue
candidate = Path(str(raw))
if not candidate.is_absolute():
candidate = ROOT / candidate
if candidate.exists() and candidate.is_file():
return candidate.resolve()
if page_num is not None:
render_dir = RENDERS_DIR / book_id
if render_dir.exists():
candidates = sorted(render_dir.glob(f"{book_id}_page_{page_num:04d}_*.*"))
if candidates:
return candidates[0].resolve()
return None
def select_transcript(record: Dict) -> str:
for key in ("final_text", "corrected_text", "text", "raw_text", "raw_ocr"):
value = record.get(key)
if value is not None:
text = str(value).strip()
if text:
return text
return ""
def build_examples(
gold_records: List[Dict],
manifest: Dict[str, Dict[str, str]],
rights_class_filter: Optional[str],
allow_mixed_rights: bool,
) -> Tuple[List[Dict], Dict]:
stats = {
"gold_rows": len(gold_records),
"kept": 0,
"skipped_missing_book": 0,
"skipped_missing_manifest": 0,
"skipped_status": 0,
"skipped_blocked_rights": 0,
"skipped_rights_filter": 0,
"skipped_missing_text": 0,
"skipped_missing_image": 0,
}
rights_seen = set()
examples: List[Dict] = []
for rec in gold_records:
book_id = str(rec.get("book_id", "")).strip()
if not book_id:
stats["skipped_missing_book"] += 1
continue
manifest_row = manifest.get(book_id)
if not manifest_row:
stats["skipped_missing_manifest"] += 1
continue
status = str(rec.get("status", "")).strip().lower()
if status and status not in {"accepted", "edited"}:
stats["skipped_status"] += 1
continue
rights_class = str(manifest_row.get("rights_class", "unknown")).strip().lower()
if rights_class in BLOCKED_RIGHTS or rights_class not in ELIGIBLE_RIGHTS:
stats["skipped_blocked_rights"] += 1
continue
if rights_class_filter and rights_class != rights_class_filter:
stats["skipped_rights_filter"] += 1
continue
text = select_transcript(rec)
if not text:
stats["skipped_missing_text"] += 1
continue
page_num = parse_page_number(rec.get("page") or rec.get("page_number"))
image_path = resolve_image_path(rec, book_id, page_num)
if not image_path:
stats["skipped_missing_image"] += 1
continue
rights_seen.add(rights_class)
example = {
"image": str(image_path),
"text": text,
"book_id": book_id,
"page": page_num,
"region_class": str(rec.get("region_class", "narration")),
"rights_class": rights_class,
"confidence": float(rec.get("confidence", 0) or 0),
"gold_line": int(rec.get("_gold_line", 0)),
}
examples.append(example)
if not allow_mixed_rights and len(rights_seen) > 1:
raise ValueError(
f"Mixed rights classes found in training set: {sorted(rights_seen)}. "
"Run separate jobs per rights class or pass --allow-mixed-rights explicitly."
)
stats["kept"] = len(examples)
stats["rights_seen"] = sorted(rights_seen)
return examples, stats
def save_training_artifacts(run_id: str, examples: List[Dict], stats: Dict, metadata: Dict) -> Path:
run_dataset_dir = TRAINING_DATASETS_DIR / run_id
run_dataset_dir.mkdir(parents=True, exist_ok=True)
examples_jsonl = run_dataset_dir / "training_examples.jsonl"
with open(examples_jsonl, "w", encoding="utf-8") as f:
for row in examples:
f.write(json.dumps(row, ensure_ascii=False) + "\n")
manifest_path = run_dataset_dir / "dataset_manifest.json"
with open(manifest_path, "w", encoding="utf-8") as f:
json.dump({"stats": stats, "metadata": metadata}, f, indent=2, ensure_ascii=False)
return run_dataset_dir
def build_hf_dataset(examples: List[Dict]):
try:
from datasets import Dataset, Image
except Exception as exc: # pragma: no cover - environment-dependent
raise RuntimeError(
"datasets[vision] is required. Install with: pip install datasets[vision]"
) from exc
ds = Dataset.from_dict({
"image": [e["image"] for e in examples],
"text": [e["text"] for e in examples],
}).cast_column("image", Image())
return ds
def push_dataset_to_hub(ds, repo_id: str, private: bool, token: Optional[str], run_id: str) -> None:
ds.push_to_hub(
repo_id,
private=private,
token=token,
commit_message=f"Smoke Signal finetune dataset {run_id}",
)
def resolve_finetune_entrypoint(explicit_script: Optional[str], surya_repo: Optional[str]) -> List[str]:
if explicit_script:
script_path = Path(explicit_script).expanduser().resolve()
if not script_path.exists():
raise FileNotFoundError(f"Surya finetune script not found: {script_path}")
return [sys.executable, str(script_path)]
try:
import importlib.util
spec = importlib.util.find_spec("surya.scripts.finetune_ocr")
if spec is not None:
return [sys.executable, "-m", "surya.scripts.finetune_ocr"]
except Exception:
pass
if surya_repo:
candidate = Path(surya_repo).expanduser().resolve() / "surya" / "scripts" / "finetune_ocr.py"
if candidate.exists():
return [sys.executable, str(candidate)]
raise RuntimeError(
"Could not resolve Surya finetune entrypoint. "
"Install surya-ocr or pass --surya-finetune-script /path/to/finetune_ocr.py"
)
def _detect_hub_model_arg_name(entrypoint_cmd: List[str]) -> str:
"""
TrainingArguments changed over time. Detect supported hub model id arg
from finetune --help output.
"""
try:
probe = subprocess.run(
entrypoint_cmd + ["--help"],
capture_output=True,
text=True,
check=False,
)
help_text = (probe.stdout or "") + "\\n" + (probe.stderr or "")
if "--hub_model_id" in help_text:
return "--hub_model_id"
if "--push_to_hub_model_id" in help_text:
return "--push_to_hub_model_id"
except Exception:
pass
# Default to current TrainingArguments key.
return "--hub_model_id"
def build_train_command(args, run_output_dir: Path) -> List[str]:
cmd = resolve_finetune_entrypoint(args.surya_finetune_script, args.surya_repo)
hub_model_arg = _detect_hub_model_arg_name(cmd)
cmd += [
"--output_dir",
str(run_output_dir),
"--dataset_name",
args.dataset_repo_id,
"--per_device_train_batch_size",
str(args.per_device_train_batch_size),
"--gradient_checkpointing",
"true" if args.gradient_checkpointing else "false",
"--max_sequence_length",
str(args.max_sequence_length),
"--num_train_epochs",
str(args.num_train_epochs),
"--learning_rate",
str(args.learning_rate),
"--logging_steps",
str(args.logging_steps),
"--save_steps",
str(args.save_steps),
"--save_total_limit",
str(args.save_total_limit),
"--remove_unused_columns",
"false",
"--push_to_hub",
"true",
hub_model_arg,
args.model_repo_id,
]
if args.pretrained_checkpoint_path:
cmd += ["--pretrained_checkpoint_path", args.pretrained_checkpoint_path]
if args.hf_token:
cmd += ["--hub_token", args.hf_token]
return cmd
def write_run_summary(path: Path, summary: Dict) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with open(path, "w", encoding="utf-8") as f:
json.dump(summary, f, indent=2, ensure_ascii=False)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Smoke Signal — Stage 11: Surya Fine-Tuning")
parser.add_argument("--gold-file", default=str(GOLD_FILE), help="Path to gold corrections JSONL")
parser.add_argument(
"--rights-class",
default=None,
choices=sorted(ELIGIBLE_RIGHTS),
help="Restrict to one rights class (recommended governance mode)",
)
parser.add_argument(
"--allow-mixed-rights",
action="store_true",
help="Allow mixed rights classes in one run (off by default)",
)
parser.add_argument(
"--dataset-repo-id",
default="Pointf5ive/smoke-signal-ocr-finetune",
help="HF dataset repo id (<namespace>/<name>)",
)
parser.add_argument(
"--model-repo-id",
default="Pointf5ive/smoke-signal-surya-ft",
help="HF model repo id (<namespace>/<name>)",
)
parser.add_argument("--hf-token", default=None, help="HF token (or set HF_TOKEN env var)")
parser.add_argument("--private-dataset", action="store_true", help="Create/push dataset repo as private")
parser.add_argument("--private-model", action="store_true", help="Create model repo as private")
parser.add_argument("--operator", default="codex", help="Operator name for governance logs")
parser.add_argument("--pretrained-checkpoint-path", default=None, help="Optional Surya init checkpoint")
parser.add_argument("--surya-finetune-script", default=None, help="Path to surya/scripts/finetune_ocr.py")
parser.add_argument("--surya-repo", default=None, help="Path to local Surya repo (fallback resolver)")
parser.add_argument("--per-device-train-batch-size", type=int, default=16)
parser.add_argument("--max-sequence-length", type=int, default=1024)
parser.add_argument("--num-train-epochs", type=float, default=2.0)
parser.add_argument("--learning-rate", type=float, default=5e-5)
parser.add_argument("--gradient-checkpointing", action="store_true")
parser.add_argument("--logging-steps", type=int, default=25)
parser.add_argument("--save-steps", type=int, default=200)
parser.add_argument("--save-total-limit", type=int, default=2)
parser.add_argument(
"--prepare-only",
action="store_true",
help="Stop after dataset prep/upload; do not start finetuning",
)
parser.add_argument(
"--skip-upload",
action="store_true",
help="Prepare local dataset artifacts but skip HF push",
)
parser.add_argument("--run-id", default=None, help="Optional explicit run id")
return parser.parse_args()
def main() -> None:
args = parse_args()
ensure_dirs()
token = args.hf_token
if not token:
import os
token = os.getenv("HF_TOKEN") or os.getenv("HUGGINGFACE_HUB_TOKEN")
args.hf_token = token
run_id = args.run_id or f"SS-FT-{utc_now().strftime('%Y%m%d-%H%M%S')}"
run_output_dir = TRAINING_RUNS_DIR / run_id
run_output_dir.mkdir(parents=True, exist_ok=True)
gold_path = Path(args.gold_file).expanduser().resolve()
if not gold_path.exists():
raise FileNotFoundError(f"Gold file not found: {gold_path}")
manifest = load_manifest()
if not manifest:
raise RuntimeError("Manifest is empty. Run 01_register_sources.py and set rights_class first.")
gold_records = load_gold_records(gold_path)
if not gold_records:
raise RuntimeError("Gold set is empty or unreadable; cannot fine-tune.")
examples, stats = build_examples(
gold_records=gold_records,
manifest=manifest,
rights_class_filter=args.rights_class,
allow_mixed_rights=args.allow_mixed_rights,
)
if not examples:
raise RuntimeError(f"No valid training examples after governance filters. Stats: {stats}")
rights_for_run = stats.get("rights_seen", [])
if not args.allow_mixed_rights and len(rights_for_run) > 1:
raise RuntimeError(
f"Mixed rights classes in run: {rights_for_run}. This violates governance by default."
)
# Governance metadata
gold_hash = sha256_file(gold_path)
text_concat = "\n".join(f"{e['book_id']}|{e['page']}|{e['text']}" for e in examples)
dataset_hash = sha256_text(text_concat)
config_version = "ss_surya_finetune_config_v0.1"
if args.pretrained_checkpoint_path:
ck_hash = sha256_text(args.pretrained_checkpoint_path)[:8]
config_version = f"{config_version}_ckpt_{ck_hash}"
metadata = {
"run_id": run_id,
"created_at": utc_iso(),
"config_version": config_version,
"schema_version": "ss_surya_ocr_finetune_dataset_v1",
"gold_file": str(gold_path),
"gold_hash": gold_hash,
"dataset_hash": dataset_hash,
"dataset_size": len(examples),
"rights_seen": rights_for_run,
"rights_filter": args.rights_class,
"dataset_repo_id": args.dataset_repo_id,
"model_repo_id": args.model_repo_id,
"pretrained_checkpoint_path": args.pretrained_checkpoint_path or "base",
"operator": args.operator,
}
dataset_artifact_dir = save_training_artifacts(run_id, examples, stats, metadata)
print(f"\nPrepared dataset artifacts: {dataset_artifact_dir}")
print(f"Examples kept: {len(examples)} | Rights: {rights_for_run} | Gold hash: {gold_hash[:12]}...")
summary = {
"metadata": metadata,
"stats": stats,
"train_command": None,
"train_returncode": None,
"train_stdout_path": None,
"train_stderr_path": None,
}
if not args.skip_upload:
if not args.hf_token:
raise RuntimeError("HF token required for upload. Pass --hf-token or set HF_TOKEN.")
ds = build_hf_dataset(examples)
push_dataset_to_hub(ds, args.dataset_repo_id, args.private_dataset, args.hf_token, run_id)
print(f"Pushed dataset to HF: {args.dataset_repo_id}")
else:
print("Skipped HF upload (--skip-upload).")
if args.prepare_only:
print("Prepare-only mode complete. Finetuning not started.")
else:
if args.skip_upload:
raise RuntimeError(
"Cannot start finetuning with --skip-upload because Surya expects --dataset_name. "
"Upload dataset first or run with --prepare-only."
)
if not args.hf_token:
raise RuntimeError("HF token required for model push during finetuning.")
train_cmd = build_train_command(args, run_output_dir)
stdout_path = run_output_dir / "finetune_stdout.log"
stderr_path = run_output_dir / "finetune_stderr.log"
summary["train_command"] = train_cmd
summary["train_stdout_path"] = str(stdout_path)
summary["train_stderr_path"] = str(stderr_path)
print("Launching Surya finetune...")
print(" ".join(train_cmd))
with open(stdout_path, "w", encoding="utf-8") as out, open(stderr_path, "w", encoding="utf-8") as err:
proc = subprocess.run(train_cmd, stdout=out, stderr=err, text=True)
summary["train_returncode"] = proc.returncode
if proc.returncode != 0:
raise RuntimeError(
f"Surya finetune failed with return code {proc.returncode}. "
f"See {stdout_path} and {stderr_path}."
)
print(f"Finetune complete. Model pushed to: {args.model_repo_id}")
summary_path = TRAINING_RUNS_DIR / f"{run_id}_summary.json"
write_run_summary(summary_path, summary)
notes = {
"model_version": args.pretrained_checkpoint_path or "base",
"dataset_size": len(examples),
"gold_hash": gold_hash,
"dataset_repo": args.dataset_repo_id,
"model_repo": args.model_repo_id,
}
append_run_log(
{
"run_id": run_id,
"date": utc_now().date().isoformat(),
"operator": args.operator,
"config_version": config_version,
"schema_version": "ss_surya_ocr_finetune_dataset_v1",
"source_batch": args.rights_class or ",".join(rights_for_run),
"pages_processed": str(len(examples)),
"errors": str(
stats["skipped_missing_book"]
+ stats["skipped_missing_manifest"]
+ stats["skipped_status"]
+ stats["skipped_blocked_rights"]
+ stats["skipped_rights_filter"]
+ stats["skipped_missing_text"]
+ stats["skipped_missing_image"]
),
"cost_usd": "",
"output_path": args.model_repo_id,
"notes": json.dumps(notes, ensure_ascii=False),
}
)
print(f"Run summary: {summary_path}")
print(f"Governance log updated: {RUN_LOG_CSV}")
if __name__ == "__main__":
main()