36 lines
985 B
Python
36 lines
985 B
Python
"""CLI entrypoint: python scripts/train.py [--config configs/default.yaml] [--epochs N]"""
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import os
|
|
import sys
|
|
|
|
import yaml
|
|
|
|
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
|
|
|
from engine.train import train # noqa: E402
|
|
|
|
|
|
def main():
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument("--config", default="configs/default.yaml")
|
|
parser.add_argument("--epochs", type=int, default=None, help="override train.epochs")
|
|
parser.add_argument("--num-threads", type=int, default=12)
|
|
parser.add_argument("--use-cuda", action="store_true", help="try CUDA if available")
|
|
args = parser.parse_args()
|
|
|
|
with open(args.config) as f:
|
|
cfg = yaml.safe_load(f)
|
|
|
|
if args.epochs is not None:
|
|
cfg["train"]["epochs"] = args.epochs
|
|
cfg["num_threads"] = args.num_threads
|
|
cfg["train"]["use_cuda"] = args.use_cuda
|
|
|
|
train(cfg)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|