diff --git a/src/pyrecest/backend_support/_pytorch_nonzero_scalar_contract.py b/src/pyrecest/backend_support/_pytorch_nonzero_scalar_contract.py index e9540cc43b..d7b6a30768 100644 --- a/src/pyrecest/backend_support/_pytorch_nonzero_scalar_contract.py +++ b/src/pyrecest/backend_support/_pytorch_nonzero_scalar_contract.py @@ -2,6 +2,8 @@ from __future__ import annotations +import numpy as np + _NONZERO_SCALAR_MESSAGE = ( "Calling nonzero on 0d arrays is not allowed. " "Use np.atleast_1d(scalar).nonzero() instead." @@ -9,7 +11,7 @@ def patch_pytorch_nonzero_scalar_contract() -> None: - """Make public and raw PyTorch ``nonzero`` reject 0-D inputs.""" + """Make public and raw PyTorch ``nonzero`` reject 0-D inputs and honor masks.""" try: import pyrecest._backend.pytorch as raw_pytorch # pylint: disable=import-outside-toplevel @@ -27,6 +29,8 @@ def patch_pytorch_nonzero_scalar_contract() -> None: return def nonzero(x): + if np.ma.isMaskedArray(x): + x = np.ma.filled(x, 0) values = x if torch.is_tensor(x) else torch.as_tensor(x) if values.ndim == 0: raise ValueError(_NONZERO_SCALAR_MESSAGE) @@ -40,4 +44,4 @@ def nonzero(x): backend.nonzero = nonzero -__all__ = ["patch_pytorch_nonzero_scalar_contract"] +__all__ = ["patch_pytorch_nonzero_scalar_contract"] \ No newline at end of file diff --git a/tests/backend_support/test_pytorch_nonzero_scalar_contract.py b/tests/backend_support/test_pytorch_nonzero_scalar_contract.py index 46eb0c46e0..e87e1c1d7a 100644 --- a/tests/backend_support/test_pytorch_nonzero_scalar_contract.py +++ b/tests/backend_support/test_pytorch_nonzero_scalar_contract.py @@ -5,13 +5,14 @@ @pytest.mark.backend_portable -def test_pytorch_nonzero_rejects_zero_dimensional_inputs(): +def test_pytorch_nonzero_rejects_scalars_and_honors_masks(): if importlib.util.find_spec("torch") is None: pytest.skip("PyTorch is not installed") result = run_backend_code( "pytorch", """ +import numpy as np import pyrecest._backend.pytorch as raw_pytorch import pyrecest.backend as backend import torch @@ -34,9 +35,17 @@ def assert_rejects_scalar(helper, value): assert rows.tolist() == [0, 1] assert columns.tolist() == [1, 0] + masked = np.ma.array( + [[0, 2], [3, 4]], + mask=[[False, True], [False, True]], + ) + rows, columns = nonzero(masked) + assert rows.tolist() == [1] + assert columns.tolist() == [0] + print("ok") """, ) assert result.returncode == 0, result.stderr - assert "ok" in result.stdout + assert "ok" in result.stdout \ No newline at end of file