Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -2,14 +2,16 @@

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."
)


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
Expand All @@ -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)
Expand All @@ -40,4 +44,4 @@ def nonzero(x):
backend.nonzero = nonzero


__all__ = ["patch_pytorch_nonzero_scalar_contract"]
__all__ = ["patch_pytorch_nonzero_scalar_contract"]
13 changes: 11 additions & 2 deletions tests/backend_support/test_pytorch_nonzero_scalar_contract.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Loading