From afde52ff892a13bf472750177c5ce1f58882f2ee Mon Sep 17 00:00:00 2001 From: Joseph Loftin Date: Tue, 18 Aug 2026 19:53:24 +0000 Subject: [PATCH] Fix 1d topk --- .../dynamo/conversion/impl/topk.py | 41 ++++++++++++------- tests/py/dynamo/conversion/test_sort_aten.py | 4 ++ tests/py/dynamo/conversion/test_topk_aten.py | 4 ++ 3 files changed, 34 insertions(+), 15 deletions(-) diff --git a/py/torch_tensorrt/dynamo/conversion/impl/topk.py b/py/torch_tensorrt/dynamo/conversion/impl/topk.py index 638cbf599e..5ee106068b 100644 --- a/py/torch_tensorrt/dynamo/conversion/impl/topk.py +++ b/py/torch_tensorrt/dynamo/conversion/impl/topk.py @@ -194,26 +194,37 @@ def topk( sorted: Optional[bool], return_indices: bool = True, ) -> Union[TRTTensor, Tuple[TRTTensor, TRTTensor]]: - if largest: - topk_layer = ctx.net.add_topk( - input, - trt.TopKOperation.MAX, - k, - get_axes_for_reduce_op(get_positive_dim(dim, len(input.shape))), - ) - else: - topk_layer = ctx.net.add_topk( - input, - trt.TopKOperation.MIN, - k, - get_axes_for_reduce_op(get_positive_dim(dim, len(input.shape))), + # ITopKLayer requires rank >= 2 (same broadcast as argmax_argmin above). + is_rank_1 = len(input.shape) == 1 + if is_rank_1: + input = impl.shuffle.reshape( + ctx, target, source_ir, f"{name}_broadcast", input, (*input.shape, 1) ) + dim = 0 + + topk_layer = ctx.net.add_topk( + input, + trt.TopKOperation.MAX if largest else trt.TopKOperation.MIN, + k, + get_axes_for_reduce_op(get_positive_dim(dim, len(input.shape))), + ) # TensorRT ITopKLayer does not have a sorted flag, it is always returning the sorted topk elements # so here no matter sorted is True or False the returned the topk Tensor object is always sorted set_layer_name(topk_layer, target, f"{name}_topk", source_ir) + values = topk_layer.get_output(0) + indices = topk_layer.get_output(1) + + if is_rank_1: + values = impl.shuffle.reshape( + ctx, target, source_ir, f"{name}_values_1d", values, (-1,) + ) + indices = impl.shuffle.reshape( + ctx, target, source_ir, f"{name}_indices_1d", indices, (-1,) + ) + if return_indices: - return topk_layer.get_output(0), topk_layer.get_output(1) + return values, indices else: - return topk_layer.get_output(0) + return values diff --git a/tests/py/dynamo/conversion/test_sort_aten.py b/tests/py/dynamo/conversion/test_sort_aten.py index 5f1258c6ac..697f5da981 100644 --- a/tests/py/dynamo/conversion/test_sort_aten.py +++ b/tests/py/dynamo/conversion/test_sort_aten.py @@ -17,6 +17,10 @@ class TestSortConverter(DispatchTestCase): ((1, 5, 2, 1), -1, True), ((1, 2, 5, 3), -2, False), ((6, 2, 1, 3), -4, True), + # ITopKLayer requires rank >= 2; converter broadcasts rank-1 + ((64,), 0, False), + ((64,), 0, True), + ((16,), -1, True), ] ) def test_sort(self, input_shape, dim, descending): diff --git a/tests/py/dynamo/conversion/test_topk_aten.py b/tests/py/dynamo/conversion/test_topk_aten.py index 2f85388548..f29d6efb19 100644 --- a/tests/py/dynamo/conversion/test_topk_aten.py +++ b/tests/py/dynamo/conversion/test_topk_aten.py @@ -19,6 +19,10 @@ class TestSortConverter(DispatchTestCase): ((6, 4), 2, 1, False, False), # default dim:-1 largest:True, sorted:True ((3, 5, 12), 3), + # ITopKLayer requires rank >= 2; converter broadcasts rank-1 + ((64,), 64, 0, True, True), + ((64,), 8, 0, False, True), + ((16,), 4, -1, True, True), ] ) def test_topk(self, input_shape, k, dim=-1, largest=True, sorted=True):