Skip to content

BUG: torch: accept Python scalars in binary elementwise functions - #478

Open
raashish1601 wants to merge 1 commit into
data-apis:mainfrom
raashish1601:fix/271-torch-binary-scalars
Open

raashish1601 wants to merge 1 commit into
data-apis:mainfrom
raashish1601:fix/271-torch-binary-scalars

Conversation

@raashish1601

Copy link
Copy Markdown

Closes #271 (also covers #410).

torch.maximum, minimum, atan2, hypot, logaddexp, nextafter and logical_and/or/xor reject Python scalars, and equal, not_equal, less, less_equal, greater, greater_equal and copysign reject a scalar as the first argument. The 2024.12 standard allows a scalar in either position.

Following the idea in the issue, _two_arg now takes scalar_x1 / scalar_x2 flags. When set, a Python scalar in that position is converted to a 0-D tensor with result_type(scalar, other) on the other argument's device. The flags are only set where torch rejects the scalar, so calls that already work (e.g. add(x, 1)) take the same path as before. nextafter and the logical functions are now wrapped too and added to __all__.

from array_api_compat import torch as xp
xp.maximum(xp.asarray([1., 2.]), 1.5)   # was TypeError
xp.less(1.5, xp.asarray([1, 2]))        # was TypeError

Testing: removed the matching test_binary_with_scalars_* entries from torch-xfails.txt; those 28 tests pass locally against array-api-tests, along with the rest of test_operators_and_elementwise_functions.py and test_signatures.py. Added tests to tests/test_torch.py; they fail on main and pass with this change. ruff check . passes.

torch.maximum, minimum, atan2, hypot, logaddexp, nextafter and the
logical functions reject Python scalars, and the comparisons and
copysign reject a scalar as the first argument. Convert the scalar to a
0-D tensor of the result dtype in these positions.

Closes data-apis#271

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

several torch binary functions don't accept scalars for one argument

1 participant