Support bounded high-resolution inference batches

This commit is contained in:
Jens
2026-07-27 14:50:34 +02:00
parent 100d9e220b
commit 31596d98f2
2 changed files with 10 additions and 2 deletions
@@ -145,6 +145,12 @@ def main() -> int:
parser.add_argument("--split", default="val")
parser.add_argument("--augment", action="store_true", help="Enable deterministic YOLO test-time augmentation.")
parser.add_argument("--imgsz", type=int, default=640)
parser.add_argument(
"--batch",
type=int,
default=16,
help="Inference batch size; lower this for high-resolution CUDA evaluation.",
)
parser.add_argument(
"--max-det",
type=int,
@@ -211,6 +217,7 @@ def main() -> int:
device=args.device,
augment=args.augment,
imgsz=args.imgsz,
batch=args.batch,
max_det=args.max_det,
iou=args.nms_iou,
verbose=False,
@@ -219,7 +226,7 @@ def main() -> int:
if args.additional_model:
additional_results = YOLO(str(args.additional_model)).predict(
image_paths, conf=min(args.thresholds), device=args.device, augment=args.augment,
imgsz=args.imgsz, max_det=args.max_det, iou=args.nms_iou, verbose=False,
imgsz=args.imgsz, batch=args.batch, max_det=args.max_det, iou=args.nms_iou, verbose=False,
)
observations: list[dict[str, Any]] = []
for result_index, (tile, result) in enumerate(zip(tiles, results, strict=True)):
@@ -315,6 +322,7 @@ def main() -> int:
"match_iou": args.match_iou,
"test_time_augmentation": args.augment,
"inference_imgsz": args.imgsz,
"inference_batch": args.batch,
"max_detections_per_tile": args.max_det,
"nms_iou": args.nms_iou,
"containment_nms": args.containment_nms,