[MPS] Add nonzero mps support (#91616) Adds nonzero support for mps: **Pseudocode**: ``` // // inputTensor = [1, 0, 0, 3] // inputNonZero = [1, 0, 0, 1] (input != 0) // scan = [1, 1, 1, 2] (prefix sum) // maskedIndices = [0, -1, -1, 1] (select) // coordinates = [0, 1, 2, 3] (coordinateAlongAxis) // scatterResult = [0, 3] (scatter) ``` Pull Request resolved: https://github.com/pytorch/pytorch/pull/91616 Approved by: https://github.com/razarmehr
diff --git a/test/test_mps.py b/test/test_mps.py index d3308d9..00dadbe 100644 --- a/test/test_mps.py +++ b/test/test_mps.py
@@ -6687,6 +6687,116 @@ supported_dtypes = [torch.float32, torch.float16, torch.int64, torch.int32, torch.int16, torch.uint8] supported_np_dtypes = [np.float32, np.float16, np.int64, np.int32, np.int16, np.uint8] + def test_nonzero_no_warning(self): + device = "mps" + t = torch.randn((2, 2), device=device) + with warnings.catch_warnings(record=True) as w: + warnings.simplefilter("always") + torch.nonzero(t) + t.nonzero() + self.assertEqual(len(w), 0) + + def test_nonzero(self): + def helper(dtype): + device = "mps" + shapes = [ + torch.Size((12,)), + torch.Size((12, 1)), + torch.Size((1, 12)), + torch.Size((6, 2)), + torch.Size((3, 2, 2)), + torch.Size((5, 5, 5)), + ] + + def gen_nontrivial_input(shape, dtype, device): + if dtype != torch.bfloat16: + return torch.randint(2, shape, device=device, dtype=dtype) + else: + # windows does not work for bfloat16 randing + return torch.randint(2, shape, device=device, dtype=torch.float).to(dtype) + + for shape in shapes: + tensor = gen_nontrivial_input(shape, dtype, device) + dst1 = torch.nonzero(tensor, as_tuple=False) + dst2 = tensor.nonzero(as_tuple=False) + dst3 = torch.empty([], dtype=torch.long, device=device) + dst3 = dst3.resize_(0) + torch.nonzero(tensor, out=dst3) + np_array = tensor.cpu().numpy() if dtype != torch.bfloat16 else tensor.float().cpu().numpy() + np_result = torch.from_numpy(np.stack(np_array.nonzero())).t() + self.assertEqual(dst1.cpu(), np_result, atol=0, rtol=0) + self.assertEqual(dst2.cpu(), np_result, atol=0, rtol=0) + self.assertEqual(dst3.cpu(), np_result, atol=0, rtol=0) + tup1 = torch.nonzero(tensor, as_tuple=True) + tup2 = tensor.nonzero(as_tuple=True) + tup1 = torch.stack(tup1).t().cpu() + tup2 = torch.stack(tup2).t().cpu() + self.assertEqual(tup1, np_result, atol=0, rtol=0) + self.assertEqual(tup2, np_result, atol=0, rtol=0) + [helper(dtype) for dtype in self.supported_dtypes] + + def test_nonzero_astuple_out(self): + device = "mps" + t = torch.randn((3, 3, 3), device=device) + out = torch.empty([], dtype=torch.long, device=device) + out = out.resize_(0) + + with self.assertRaises(RuntimeError): + torch.nonzero(t, as_tuple=True, out=out) + + self.assertEqual(torch.nonzero(t, as_tuple=False, out=out), torch.nonzero(t, out=out)) + + # Verifies that JIT script cannot handle the as_tuple kwarg + # See Issue https://github.com/pytorch/pytorch/issues/45499. + def _foo(t): + tuple_result = torch.nonzero(t, as_tuple=True) + nontuple_result = torch.nonzero(t, as_tuple=False) + out = torch.empty_like(nontuple_result) + torch.nonzero(t, as_tuple=False, out=out) + return tuple_result, nontuple_result, out + + with self.assertRaises(RuntimeError): + scripted_foo = torch.jit.script(_foo) + + # Verifies that JIT tracing works fine + traced_foo = torch.jit.trace(_foo, t) + traced_tuple, traced_nontuple, traced_out = traced_foo(t) + expected_tuple = torch.nonzero(t, as_tuple=True) + expected_nontuple = torch.nonzero(t) + + self.assertEqual(traced_tuple, expected_tuple) + self.assertEqual(traced_nontuple, expected_nontuple) + self.assertEqual(traced_out, expected_nontuple) + + def test_nonzero_discontiguous(self): + device = "mps" + shape = (4, 4) + tensor = torch.randint(2, shape, device=device) + tensor_nc = torch.empty(shape[0], shape[1] * 2, device=device)[:, ::2].copy_(tensor) + dst1 = tensor.nonzero(as_tuple=False) + dst2 = tensor_nc.nonzero(as_tuple=False) + self.assertEqual(dst1, dst2, atol=0, rtol=0) + dst3 = torch.empty_like(dst1) + data_ptr = dst3.data_ptr() + # expect dst3 storage to be reused + torch.nonzero(tensor, out=dst3) + self.assertEqual(data_ptr, dst3.data_ptr()) + self.assertEqual(dst1, dst3, atol=0, rtol=0) + # discontiguous out + dst4 = torch.empty(dst1.size(0), dst1.size(1) * 2, dtype=torch.long, device=device)[:, ::2] + data_ptr = dst4.data_ptr() + strides = dst4.stride() + torch.nonzero(tensor, out=dst4) + self.assertEqual(data_ptr, dst4.data_ptr()) + self.assertEqual(dst1, dst4, atol=0, rtol=0) + self.assertEqual(strides, dst4.stride()) + + def test_nonzero_non_diff(self): + device = "mps" + x = torch.randn(10, requires_grad=True) + nz = x.nonzero() + self.assertFalse(nz.requires_grad) + def test_masked_select(self): x = torch.randn(3, 4) x_mps = x.to("mps") @@ -7841,7 +7951,8 @@ 'vsplit': ['b8', 'f16', 'f32', 'i16', 'i32', 'i64', 'u8'], 'vstack': ['b8', 'f16', 'f32', 'i16', 'i32', 'i64', 'u8'], 'zero_': ['b8', 'f16', 'f32', 'i16', 'i32', 'i64', 'u8'], - 'where': ['f16', 'f32', 'i16', 'i32', 'i64', 'u8'] + 'where': ['f16', 'f32', 'i16', 'i32', 'i64', 'u8'], + 'nonzero': ['f32', 'i16', 'i32', 'i64'] } @@ -8066,6 +8177,8 @@ 'slice_scatter': [torch.uint8], 'square': [torch.bool, torch.int16, torch.int32, torch.int64, torch.uint8], # moved from section below + # count_nonzero returns wrong results for these dtypes + 'nonzero': [torch.uint8, torch.float16], # ALLOW_LIST doesn't know about variants 'nn.functional.padconstant': None, @@ -8141,7 +8254,6 @@ 'eq': None, 'mul': None, 'cartesian_prod': None, - 'nonzero': None, 'bool': None, 'inner': None, 'dstack': None,