105 lines
3.4 KiB
Python
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()
|