Download src/train.py from imperat0/strike: direct link, hf CLI and curl.
- Browser
- Download file 1.61 kB
-
https://huggingface.co/spaces/imperat0/strike/resolve/main/src/train.py
- Command line
-
hf download hf://spaces/imperat0/strike/src/train.py
-
curl -L -o train.py https://huggingface.co/spaces/imperat0/strike/resolve/main/src/train.py
1.61 kB
| import argparse | |
| import shutil | |
| from pathlib import Path | |
| def main(): | |
| parser = argparse.ArgumentParser(description="Treino do detector de componentes de arquitetura") | |
| parser.add_argument("--data", type=str, default="../data/data.yaml", | |
| help="Caminho para o data.yaml do dataset") | |
| parser.add_argument("--epochs", type=int, default=100) | |
| parser.add_argument("--imgsz", type=int, default=960, | |
| help="Resolução de treino. 960 (acima do padrão 640) para preservar " | |
| "ícones pequenos em diagramas densos.") | |
| parser.add_argument("--model", type=str, default="yolov8n.pt", | |
| help="Checkpoint base do YOLOv8 (transfer learning)") | |
| parser.add_argument("--batch", type=int, default=8, | |
| help="Batch reduzido para acomodar imgsz=960 em GPUs/CPUs limitadas.") | |
| parser.add_argument("--out", type=str, default="../models/best.pt") | |
| args = parser.parse_args() | |
| from ultralytics import YOLO | |
| model = YOLO(args.model) | |
| results = model.train( | |
| data=args.data, | |
| epochs=args.epochs, | |
| imgsz=args.imgsz, | |
| batch=args.batch, | |
| project="runs_threat_modeling", | |
| name="detector_componentes", | |
| exist_ok=True, | |
| ) | |
| run_dir = Path(results.save_dir) | |
| best_ckpt = run_dir / "weights" / "best.pt" | |
| out_path = Path(args.out) | |
| out_path.parent.mkdir(parents=True, exist_ok=True) | |
| shutil.copy(best_ckpt, out_path) | |
| print(f"[train] Melhor checkpoint copiado para {out_path}") | |
| if __name__ == "__main__": | |
| main() | |