Add TensorRT EfficientNMS plugin to ONNX model.
(
model_proto: onnx.ModelProto,
num_classes: int,
max_detections: int,
)
| 459 | ) |
| 460 | |
| 461 | def _add_efficientnms_to_model( |
| 462 | model_proto: onnx.ModelProto, |
| 463 | num_classes: int, |
| 464 | max_detections: int, |
| 465 | ) -> None: |
| 466 | """Add TensorRT EfficientNMS plugin to ONNX model.""" |
| 467 | import onnx |
| 468 | import onnx.helper as helper |
| 469 | |
| 470 | graph = model_proto.graph |
| 471 | |
| 472 | # The raw_output is [B, anchors, num_classes+4] |
| 473 | # We need to split it into boxes [B, anchors, 4] and scores [B, anchors, num_classes] |
| 474 | |
| 475 | # Create constants for slicing |
| 476 | # Slice for classes: [:, :, :num_classes] |
| 477 | start_cls = helper.make_tensor("start_cls", onnx.TensorProto.INT64, [1], [0]) |
| 478 | end_cls = helper.make_tensor( |
| 479 | "end_cls", onnx.TensorProto.INT64, [1], [num_classes] |
| 480 | ) |
| 481 | axes_cls = helper.make_tensor("axes_cls", onnx.TensorProto.INT64, [1], [2]) |
| 482 | |
| 483 | # Slice for boxes: [:, :, num_classes:] |
| 484 | start_box = helper.make_tensor( |
| 485 | "start_box", onnx.TensorProto.INT64, [1], [num_classes] |
| 486 | ) |
| 487 | end_box = helper.make_tensor( |
| 488 | "end_box", onnx.TensorProto.INT64, [1], [num_classes + 4] |
| 489 | ) |
| 490 | axes_box = helper.make_tensor("axes_box", onnx.TensorProto.INT64, [1], [2]) |
| 491 | |
| 492 | # Add constant tensors to graph |
| 493 | graph.initializer.extend( |
| 494 | [start_cls, end_cls, axes_cls, start_box, end_box, axes_box] |
| 495 | ) |
| 496 | |
| 497 | # Create slice nodes to split raw_output |
| 498 | slice_cls = helper.make_node( |
| 499 | "Slice", |
| 500 | inputs=["raw_output", "start_cls", "end_cls", "axes_cls"], |
| 501 | outputs=["class_logits"], |
| 502 | name="slice_class_logits", |
| 503 | ) |
| 504 | |
| 505 | slice_box = helper.make_node( |
| 506 | "Slice", |
| 507 | inputs=["raw_output", "start_box", "end_box", "axes_box"], |
| 508 | outputs=["box_predictions"], |
| 509 | name="slice_boxes", |
| 510 | ) |
| 511 | |
| 512 | # Apply sigmoid to class logits to get scores |
| 513 | sigmoid_node = helper.make_node( |
| 514 | "Sigmoid", |
| 515 | inputs=["class_logits"], |
| 516 | outputs=["class_scores"], |
| 517 | name="sigmoid_scores", |
| 518 | ) |
no outgoing calls
no test coverage detected