first commit
This commit is contained in:
35
scripts/train.py
Normal file
35
scripts/train.py
Normal file
@@ -0,0 +1,35 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user