Outflock / trainObbModel.py
aseylys's picture
Upload 4 files
194f0f9 verified
Raw
History Blame Contribute Delete
4.99 kB
import argparse
import random
import shutil
from pathlib import Path
import torch
import yaml
from ultralytics import YOLO
classes = {
0: "Body",
1: "Lens",
2: "System",
}
# ~/Outflock
REPO_DIR = Path(__file__).parent.parent
IMAGES_DIR = REPO_DIR / "train/data/images"
LABELS_DIR = REPO_DIR / "train/data/labels"
DATASET_DIR = REPO_DIR / "train/datasets/camera_obb"
RUNS_DIR = REPO_DIR / "train/runs"
ONNX_DIR = REPO_DIR / "model"
valRatio = 0.2
seed = 42
def hasValidLabel(imagePath: Path) -> bool:
labelPath = LABELS_DIR / f"{imagePath.stem}.txt"
if not labelPath.exists():
return False
lines = labelPath.read_text().strip().splitlines()
if not lines:
return False
for line in lines:
parts = line.split()
if len(parts) != 9:
return False
classId = int(parts[0])
if classId not in classes:
return False
coords = [float(value) for value in parts[1:]]
if any(value < 0 or value > 1 for value in coords):
return False
return True
def copyExample(imagePath: Path, split: str) -> None:
labelPath = LABELS_DIR / f"{imagePath.stem}.txt"
imageOut = DATASET_DIR / "images" / split / imagePath.name
labelOut = DATASET_DIR / "labels" / split / labelPath.name
imageOut.parent.mkdir(parents=True, exist_ok=True)
labelOut.parent.mkdir(parents=True, exist_ok=True)
shutil.copy2(imagePath, imageOut)
shutil.copy2(labelPath, labelOut)
def prepareDataset() -> Path:
if DATASET_DIR.exists():
shutil.rmtree(DATASET_DIR)
imagePaths = sorted(
path
for path in IMAGES_DIR.iterdir()
if path.suffix.lower() in {".jpg", ".jpeg", ".png", ".webp"}
and hasValidLabel(path)
)
random.Random(seed).shuffle(imagePaths)
valCount = max(1, int(len(imagePaths) * valRatio))
valImages = set(imagePaths[:valCount])
trainImages = imagePaths[valCount:]
for imagePath in trainImages:
copyExample(imagePath, "train")
for imagePath in valImages:
copyExample(imagePath, "val")
dataYaml = DATASET_DIR / "data.yaml"
dataYaml.write_text(
yaml.safe_dump(
{
"path": str(DATASET_DIR.resolve()),
"train": "images/train",
"val": "images/val",
"names": classes,
},
sort_keys=False,
)
)
print(f"Prepared {len(trainImages)} train and {len(valImages)} val images")
return dataYaml
def trainModel(dataYaml: Path, onnx: bool = False) -> None:
print(f"CUDA available: {torch.cuda.is_available()}")
if torch.cuda.is_available():
print(f"GPU: {torch.cuda.get_device_name(0)}")
model = YOLO("yolov8m-obb.pt")
model.train(
data=str(dataYaml),
epochs=50,
imgsz=960,
project=str(RUNS_DIR),
name="flockOBB",
task="obb",
batch=16,
device=0,
workers=8,
patience=20,
pretrained=True,
optimizer="auto",
amp=True,
# Detection-specific defaults worth making explicit.
single_cls=False,
rect=False,
cache=False,
# Augmentation. Conservative for real camera detection.
degrees=5,
translate=0.08,
scale=0.4,
shear=0.0,
perspective=0.0005,
flipud=0.0,
fliplr=0.5,
mosaic=0.7,
mixup=0.05,
copy_paste=0.0,
)
if not onnx:
return
# Optional: export to ONNX and weights for OpenCV later.
bestWeights = Path(model.trainer.best)
if not bestWeights.exists():
saveDir = Path(model.trainer.save_dir)
bestWeights = saveDir / "weights" / "best.pt"
if not bestWeights.exists():
raise FileNotFoundError(f"Could not find trained best weights at {bestWeights}")
print(f"Best weights saved to: {bestWeights}")
exportModel = YOLO(str(bestWeights))
onnxPath = Path(
exportModel.export(
format="onnx",
imgsz=960,
opset=12,
simplify=True,
dynamic=False,
)
)
ONNX_DIR.mkdir(parents=True, exist_ok=True)
targetPath = ONNX_DIR / onnxPath.name
shutil.copy2(onnxPath, targetPath)
print(f"ONNX model copied to: {targetPath}")
def parseArgs() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument(
"--onnx",
action="store_true",
help="Export the trained best weights to ONNX after training.",
)
parser.add_argument(
"--clean",
action="store_true",
help="Clean Non-ONNX model directories.",
)
return parser.parse_args()
if __name__ == "__main__":
args = parseArgs()
dataYaml = prepareDataset()
trainModel(dataYaml, onnx=args.onnx)
if args.clean:
shutil.rmtree(RUNS_DIR)
print(f"Removed Non-ONNX model directory: {RUNS_DIR}")