diff --git a/scripts/assess_belgium_building_training_iteration.py b/scripts/assess_belgium_building_training_iteration.py index df1bfa07..5063576f 100644 --- a/scripts/assess_belgium_building_training_iteration.py +++ b/scripts/assess_belgium_building_training_iteration.py @@ -35,7 +35,7 @@ def select_calibration_threshold(report: dict[str, Any]) -> dict[str, Any]: INFERENCE_CONFIG_FIELDS = ( - "model", "match_iou", "test_time_augmentation", "inference_imgsz", + "model", "match_iou", "test_time_augmentation", "inference_imgsz", "inference_batch", "max_detections_per_tile", "nms_iou", "containment_nms", "box_scale", "box_offset_x", "box_offset_y", "additional_model", "ensemble_mode", "ensemble_match_iou", "proposal_classifier", "proposal_classifier_threshold", diff --git a/scripts/evaluate_belgium_building_candidate.py b/scripts/evaluate_belgium_building_candidate.py index 05cdff83..cb3efd45 100644 --- a/scripts/evaluate_belgium_building_candidate.py +++ b/scripts/evaluate_belgium_building_candidate.py @@ -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,