50 lines
1.5 KiB
Python
50 lines
1.5 KiB
Python
from argparse import ArgumentParser
|
|
from pathlib import Path
|
|
|
|
from ultralytics import YOLO
|
|
|
|
|
|
def parse_args():
|
|
parser = ArgumentParser(description="Train YOLOv8 on this dataset.")
|
|
parser.add_argument("--model", default="yolov8n.pt", help="Model or checkpoint path.")
|
|
parser.add_argument("--epochs", type=int, default=100)
|
|
parser.add_argument("--imgsz", type=int, default=640)
|
|
parser.add_argument("--batch", type=int, default=8, help="Use -1 for automatic batch size.")
|
|
parser.add_argument("--device", default="0", help="CUDA device such as 0, or cpu.")
|
|
parser.add_argument("--workers", type=int, default=4)
|
|
parser.add_argument("--patience", type=int, default=30)
|
|
parser.add_argument("--name", default="yolov8n_acne")
|
|
parser.add_argument("--resume", action="store_true", help="Resume from --model checkpoint.")
|
|
return parser.parse_args()
|
|
|
|
|
|
def main():
|
|
args = parse_args()
|
|
root = Path(__file__).resolve().parent
|
|
data = root / "data.yaml"
|
|
|
|
if not data.is_file():
|
|
raise FileNotFoundError(f"Dataset config not found: {data}")
|
|
|
|
model = YOLO(args.model)
|
|
model.train(
|
|
data=str(data),
|
|
epochs=args.epochs,
|
|
imgsz=args.imgsz,
|
|
batch=args.batch,
|
|
device=args.device,
|
|
workers=args.workers,
|
|
patience=args.patience,
|
|
project=str(root / "runs"),
|
|
name=args.name,
|
|
pretrained=True,
|
|
cache=False,
|
|
amp=True,
|
|
plots=True,
|
|
resume=args.resume,
|
|
)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|