Support bounded high-resolution inference batches
This commit is contained in:
@@ -35,7 +35,7 @@ def select_calibration_threshold(report: dict[str, Any]) -> dict[str, Any]:
|
|||||||
|
|
||||||
|
|
||||||
INFERENCE_CONFIG_FIELDS = (
|
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",
|
"max_detections_per_tile", "nms_iou", "containment_nms", "box_scale",
|
||||||
"box_offset_x", "box_offset_y", "additional_model", "ensemble_mode",
|
"box_offset_x", "box_offset_y", "additional_model", "ensemble_mode",
|
||||||
"ensemble_match_iou", "proposal_classifier", "proposal_classifier_threshold",
|
"ensemble_match_iou", "proposal_classifier", "proposal_classifier_threshold",
|
||||||
|
|||||||
@@ -145,6 +145,12 @@ def main() -> int:
|
|||||||
parser.add_argument("--split", default="val")
|
parser.add_argument("--split", default="val")
|
||||||
parser.add_argument("--augment", action="store_true", help="Enable deterministic YOLO test-time augmentation.")
|
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("--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(
|
parser.add_argument(
|
||||||
"--max-det",
|
"--max-det",
|
||||||
type=int,
|
type=int,
|
||||||
@@ -211,6 +217,7 @@ def main() -> int:
|
|||||||
device=args.device,
|
device=args.device,
|
||||||
augment=args.augment,
|
augment=args.augment,
|
||||||
imgsz=args.imgsz,
|
imgsz=args.imgsz,
|
||||||
|
batch=args.batch,
|
||||||
max_det=args.max_det,
|
max_det=args.max_det,
|
||||||
iou=args.nms_iou,
|
iou=args.nms_iou,
|
||||||
verbose=False,
|
verbose=False,
|
||||||
@@ -219,7 +226,7 @@ def main() -> int:
|
|||||||
if args.additional_model:
|
if args.additional_model:
|
||||||
additional_results = YOLO(str(args.additional_model)).predict(
|
additional_results = YOLO(str(args.additional_model)).predict(
|
||||||
image_paths, conf=min(args.thresholds), device=args.device, augment=args.augment,
|
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]] = []
|
observations: list[dict[str, Any]] = []
|
||||||
for result_index, (tile, result) in enumerate(zip(tiles, results, strict=True)):
|
for result_index, (tile, result) in enumerate(zip(tiles, results, strict=True)):
|
||||||
@@ -315,6 +322,7 @@ def main() -> int:
|
|||||||
"match_iou": args.match_iou,
|
"match_iou": args.match_iou,
|
||||||
"test_time_augmentation": args.augment,
|
"test_time_augmentation": args.augment,
|
||||||
"inference_imgsz": args.imgsz,
|
"inference_imgsz": args.imgsz,
|
||||||
|
"inference_batch": args.batch,
|
||||||
"max_detections_per_tile": args.max_det,
|
"max_detections_per_tile": args.max_det,
|
||||||
"nms_iou": args.nms_iou,
|
"nms_iou": args.nms_iou,
|
||||||
"containment_nms": args.containment_nms,
|
"containment_nms": args.containment_nms,
|
||||||
|
|||||||
Reference in New Issue
Block a user