import argparse import numpy as np from tqdm import tqdm import time import os import threading from queue import Queue, Empty from pathlib import Path from superpoint_pruning.distillation.utils import load_grayscale_image, rescale_image from superpoint_pruning.paths import DEFAULT_IMAGE_DIR QUEUE_SENTINEL = object() def benchmark_tensorrt( img_dir: str, model_path: str, max_keypoints: int = 512, descriptor_dim: int = 256, num_images: int = 1000, start_index: int = 1000, num_loader_threads: int = 2, queue_size: int = 32, image_size: tuple = (640, 480), ): try: from superpoint_pruning.evaluation.onnx_helper import SP_ONNXClassifierWrapper except ImportError: raise ImportError( "ONNX helper requires pycuda and TensorRT utilities installed." ) trt_model = SP_ONNXClassifierWrapper( str(model_path), max_keypoints=max_keypoints, descriptor_dim=descriptor_dim ) times = [] index_queue = Queue() data_queue = Queue(maxsize=max(1, queue_size)) errors = [] errors_lock = threading.Lock() stop_event = threading.Event() processed_images = 0 for img_index in range(start_index, start_index + num_images): index_queue.put(img_index) effective_loader_threads = max(1, num_loader_threads) def loader_worker(): try: while not stop_event.is_set(): try: img_index = index_queue.get_nowait() except Empty: break img_path = os.path.join(img_dir, f"image_0_{img_index}.png") original = load_grayscale_image(img_path) img, scale = rescale_image(original, new_size=image_size) img = img[None, None].astype(np.float32) data_queue.put((img, scale)) except Exception as exc: with errors_lock: errors.append(exc) stop_event.set() finally: data_queue.put(QUEUE_SENTINEL) loader_threads = [ threading.Thread(target=loader_worker, daemon=True) for _ in range(effective_loader_threads) ] for thread in loader_threads: thread.start() try: finished_loaders = 0 with tqdm(total=num_images, desc="Inferencing", unit="img") as pbar: while finished_loaders < effective_loader_threads: if stop_event.is_set() and data_queue.empty(): break try: item = data_queue.get(timeout=0.1) except Empty: continue if item is QUEUE_SENTINEL: finished_loaders += 1 continue img, scale = item scale = np.array(scale, dtype=np.float32) start = time.perf_counter() keypoints, _, descriptors = trt_model.predict(img) keypoints = (keypoints.astype(np.float32) + 0.5) / scale[None] - 0.5 stop = time.perf_counter() times.append(stop - start) processed_images += 1 pbar.update(1) for thread in loader_threads: thread.join() if errors: raise RuntimeError(f"Benchmark failed in loader thread: {errors[0]}") if processed_images != num_images: raise RuntimeError( f"Processed {processed_images}/{num_images} images before stopping." ) print(f"Average time: {np.array(times).mean()}") finally: trt_model.close() def add_parser_args(parser: argparse.ArgumentParser) -> None: parser.add_argument("--img-dir", type=Path, default=DEFAULT_IMAGE_DIR) parser.add_argument( "--model-path", type=Path, required=True, help="Path to the TensorRT SuperPoint engine.", ) parser.add_argument( "--max-keypoints", type=int, default=512, help="Maximum number of keypoints from input.", ) parser.add_argument( "--descriptor-dim", type=int, default=256, help="Descriptor dimension." ) parser.add_argument( "--num-images", type=int, default=1000, help="Number of images to benchmark." ) parser.add_argument( "--start-index", type=int, default=1000, help="Start index of the images." ) parser.add_argument( "--num-loader-threads", type=int, default=2, help="Number of producer threads for loading and preprocessing images.", ) parser.add_argument( "--queue-size", type=int, default=64, help="Max number of preprocessed images buffered for inference.", ) parser.add_argument("--width", type=int, default=640, help="Image width.") parser.add_argument("--height", type=int, default=480, help="Image height.") def main(args: argparse.Namespace) -> None: benchmark_tensorrt( img_dir=args.img_dir, model_path=args.model_path, max_keypoints=args.max_keypoints, descriptor_dim=args.descriptor_dim, num_images=args.num_images, start_index=args.start_index, num_loader_threads=args.num_loader_threads, queue_size=args.queue_size, image_size=(args.width, args.height), ) if __name__ == "__main__": parser = argparse.ArgumentParser( description="Benchmark a TensorRT SuperPoint engine." ) add_parser_args(parser) main(parser.parse_args())