MCPcopy Create free account
hub / github.com/CVCUDA/CV-CUDA / _add_efficientnms_to_model

Function _add_efficientnms_to_model

samples/common.py:461–562  ·  view source on GitHub ↗

Add TensorRT EfficientNMS plugin to ONNX model.

(
        model_proto: onnx.ModelProto,
        num_classes: int,
        max_detections: int,
    )

Source from the content-addressed store, hash-verified

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 )

Callers 1

export_retinanet_onnxFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected