"""Export a YOLOv8n face-detection model to ONNX and build a TensorRT engine.""" import argparse import json import logging from pathlib import Path logger = logging.getLogger(__name__) def export_onnx(output_dir: Path, input_size: int = 640) -> Path: """Export YOLOv8n to ONNX with the required input shape.""" output_dir.mkdir(parents=True, exist_ok=True) onnx_path = output_dir / "face_detector.onnx" metadata_path = output_dir / "model.json" try: from ultralytics import YOLO model = YOLO("yolov8n.pt") model.export( format="onnx", imgsz=input_size, half=False, simplify=True, dynamic=False, ) exported = Path("yolov8n.onnx") if exported.exists(): exported.rename(onnx_path) logger.info("Exported ONNX model to %s", onnx_path) except Exception as exc: logger.error("Failed to export ONNX model: %s", exc) raise metadata = { "input_shape": [1, 3, input_size, input_size], "mean": [0.485, 0.456, 0.406], "std": [0.229, 0.224, 0.225], "confidence_threshold": 0.25, "iou_threshold": 0.45, } metadata_path.write_text(json.dumps(metadata, indent=2)) return onnx_path def build_tensorrt_engine(onnx_path: Path, output_dir: Path, max_batch_size: int = 32) -> Path: """Build and serialize a TensorRT FP32 engine from an ONNX file.""" engine_path = output_dir / "face_detector.trt" try: import tensorrt as trt logger = trt.Logger(trt.Logger.WARNING) builder = trt.Builder(logger) network = builder.create_network(1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)) parser = trt.OnnxParser(network, logger) onnx_data = onnx_path.read_bytes() if not parser.parse(onnx_data): for error in range(parser.num_errors): logger.error(parser.get_error(error)) raise RuntimeError("ONNX parsing failed") config = builder.create_builder_config() config.max_workspace_size = 4 << 30 # 4GB config.set_flag(trt.BuilderFlag.FP32) profile = builder.create_optimization_profile() input_name = network.get_input(0).name profile.set_shape( input_name, (1, 3, 640, 640), (max_batch_size // 2, 3, 640, 640), (max_batch_size, 3, 640, 640), ) config.add_optimization_profile(profile) engine = builder.build_engine(network, config) if engine is None: raise RuntimeError("TensorRT engine build failed") engine_path.write_bytes(engine.serialize()) logger.info("Serialized TensorRT engine to %s", engine_path) return engine_path except Exception as exc: logger.error("Failed to build TensorRT engine: %s", exc) raise def main(): logging.basicConfig(level=logging.INFO) parser = argparse.ArgumentParser(description="Export YOLOv8n face detector to ONNX/TensorRT") parser.add_argument("--output-dir", type=Path, default=Path("models/face_detector")) parser.add_argument("--input-size", type=int, default=640) parser.add_argument("--max-batch-size", type=int, default=32) args = parser.parse_args() onnx_path = export_onnx(args.output_dir, args.input_size) build_tensorrt_engine(onnx_path, args.output_dir, args.max_batch_size) if __name__ == "__main__": main()