diff --git a/py/torch_tensorrt/dynamo/conversion/impl/topk.py b/py/torch_tensorrt/dynamo/conversion/impl/topk.py index 638cbf599e..3dac11d9a8 100644 --- a/py/torch_tensorrt/dynamo/conversion/impl/topk.py +++ b/py/torch_tensorrt/dynamo/conversion/impl/topk.py @@ -194,26 +194,47 @@ def topk( sorted: Optional[bool], return_indices: bool = True, ) -> Union[TRTTensor, Tuple[TRTTensor, TRTTensor]]: + # Resolve dim on the original shape before any potential reshape. + positive_dim = get_positive_dim(dim, len(input.shape)) + + # ITopKLayer requires >= 2 dims; reshape (N,) -> (N, 1), run topk, then squeeze back. + squeezed_1d = len(input.shape) == 1 + if squeezed_1d: + input = impl.shuffle.reshape( + ctx, target, source_ir, f"{name}_unsqueeze1d", input, (*input.shape, 1) + ) + 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))), + get_axes_for_reduce_op(positive_dim), ) else: topk_layer = ctx.net.add_topk( input, trt.TopKOperation.MIN, k, - get_axes_for_reduce_op(get_positive_dim(dim, len(input.shape))), + get_axes_for_reduce_op(positive_dim), ) # 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 squeezed_1d: + values = impl.squeeze.squeeze( + ctx, target, source_ir, f"{name}_squeeze_values", values, 1 + ) + indices = impl.squeeze.squeeze( + ctx, target, source_ir, f"{name}_squeeze_indices", 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_topk_aten.py b/tests/py/dynamo/conversion/test_topk_aten.py index 2f85388548..562a2d2af2 100644 --- a/tests/py/dynamo/conversion/test_topk_aten.py +++ b/tests/py/dynamo/conversion/test_topk_aten.py @@ -34,5 +34,15 @@ def forward(self, x): ) +class TestTopk1DConverter(DispatchTestCase): + def test_topk_1d_values_and_indices_match_eager(self): + class Topk(nn.Module): + def forward(self, x): + return torch.ops.aten.topk.default(x, 3) + + inputs = [torch.tensor([4.0, 1.0, 7.0, 2.0, 9.0, 3.0])] + self.run_test(Topk(), inputs, enable_passes=True) + + if __name__ == "__main__": run_tests()