Files
VideoDetect/src/export_face_detector.py
T
2026-08-10 15:22:22 -04:00

105 lines
3.4 KiB
Python

"""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()