From e5d5ae55f99cabf1a5bedb963224fb4e1b74978e Mon Sep 17 00:00:00 2001 From: chenyu Date: Sun, 15 Jun 2025 21:21:15 -0700 Subject: [PATCH] smaller inputs for test_sort and test_topk (#10829) --- test/test_ops.py | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/test/test_ops.py b/test/test_ops.py index bd159ab326..f3a887c549 100644 --- a/test/test_ops.py +++ b/test/test_ops.py @@ -1093,8 +1093,8 @@ class TestOps(unittest.TestCase): def test_sort(self): for dim in [-1, 0, 1]: for descending in [True, False]: - helper_test_op([(8,45,6)], lambda x: x.sort(dim, descending).values, lambda x: x.sort(dim, descending)[0], forward_only=True) - helper_test_op([(8,45,6)], lambda x: x.sort(dim, descending).indices.type(torch.int32), lambda x: x.sort(dim, descending)[1], + helper_test_op([(8,8,6)], lambda x: x.sort(dim, descending).values, lambda x: x.sort(dim, descending)[0], forward_only=True) + helper_test_op([(8,8,6)], lambda x: x.sort(dim, descending).indices.type(torch.int32), lambda x: x.sort(dim, descending)[1], forward_only=True) # repeated values helper_test_op(None, lambda x: x.sort(stable=True).values, lambda x: x.sort()[0], forward_only=True, vals=[[0, 1] * 9]) @@ -1110,12 +1110,12 @@ class TestOps(unittest.TestCase): for dim in [0, 1, -1]: for largest in [True, False]: for sorted_ in [True]: # TODO support False - helper_test_op([(10,12,6)], - lambda x: x.topk(5, dim, largest, sorted_).values, - lambda x: x.topk(5, dim, largest, sorted_)[0], forward_only=True) - helper_test_op([(10,12,6)], - lambda x: x.topk(5, dim, largest, sorted_).indices.type(torch.int32), - lambda x: x.topk(5, dim, largest, sorted_)[1], forward_only=True) + helper_test_op([(6,5,4)], + lambda x: x.topk(4, dim, largest, sorted_).values, + lambda x: x.topk(4, dim, largest, sorted_)[0], forward_only=True) + helper_test_op([(5,5,4)], + lambda x: x.topk(4, dim, largest, sorted_).indices.type(torch.int32), + lambda x: x.topk(4, dim, largest, sorted_)[1], forward_only=True) # repeated values value, indices = Tensor([1, 1, 0, 1, 0, 1, 0, 0, 1, 0, 0, 0, 1, 0]).topk(3) np.testing.assert_equal(value.numpy(), [1, 1, 1])