AnemiaLens / backend /scripts /train_archive_model_v8.py
asnannp's picture
sync: sync backend code, models, schemas, and API routers to Hugging Face Space cleanly
f559cc0
Raw
History Blame Contribute Delete
2.38 kB
"""
Train the AnemiaLens v8 clinical-robust archive model.
Usage::
python scripts/train_archive_model_v8.py [--dataset PATH] [--output PATH] [--report PATH]
"""
from __future__ import annotations
import argparse
import json
import sys
import time
from pathlib import Path
import joblib
ROOT = Path(__file__).resolve().parents[2]
BACKEND_ROOT = ROOT / "backend"
sys.path.insert(0, str(BACKEND_ROOT))
from app.config import DEFAULT_TRAINING_REPORT_PATH # noqa: E402
from app.ml.archive_model_v8 import V8_VERSION, train_archive_model_v8 # noqa: E402
DEFAULT_DATASET = ROOT / "archive" / "dataset anemia"
DEFAULT_OUTPUT = BACKEND_ROOT / "models" / f"{V8_VERSION}.joblib"
def _build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(
description="Train the v8 AnemiaLens archive screening model.",
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
parser.add_argument("--dataset", type=Path, default=DEFAULT_DATASET)
parser.add_argument("--output", type=Path, default=DEFAULT_OUTPUT)
parser.add_argument("--report", type=Path, default=DEFAULT_TRAINING_REPORT_PATH)
parser.add_argument("--splits", type=int, default=10)
parser.add_argument("--test-size", type=float, default=0.2)
return parser
def main(argv: list[str] | None = None) -> int:
args = _build_parser().parse_args(argv)
if not args.dataset.exists():
print(f"ERROR: Dataset directory not found: {args.dataset}", file=sys.stderr)
return 1
t0 = time.perf_counter()
try:
artifact, report = train_archive_model_v8(
args.dataset,
n_splits=args.splits,
test_size=args.test_size,
)
except Exception as exc:
print(f"ERROR: Training failed — {exc}", file=sys.stderr)
return 1
args.output.parent.mkdir(parents=True, exist_ok=True)
args.report.parent.mkdir(parents=True, exist_ok=True)
joblib.dump(artifact, args.output)
args.report.write_text(json.dumps(report, indent=2), encoding="utf-8")
elapsed = time.perf_counter() - t0
print(f"Saved artifact to {args.output}")
print(f"Saved report to {args.report}")
print(f"Validation metrics: {json.dumps(report['metrics'], indent=2)}")
print(f"Completed in {elapsed:.1f}s")
return 0
if __name__ == "__main__":
raise SystemExit(main())