Skip to content
Open
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
51 changes: 34 additions & 17 deletions src/array_api_compat/torch/_aliases.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,9 +64,15 @@
_promotion_table.update({(a, a): a for a in _array_api_dtypes})


def _two_arg(f):
def _two_arg(f, *, scalar_x1=False, scalar_x2=False):
# scalar_x1, scalar_x2: convert a Python scalar in this position to a tensor,
# for torch functions which do not accept Python scalars there.
@_wraps(f)
def _f(x1, x2, /, **kwargs):
if scalar_x1:
x1 = _scalar_to_tensor(x1, x2)
if scalar_x2:
x2 = _scalar_to_tensor(x2, x1)
x1, x2 = _fix_promotion(x1, x2)
return f(x1, x2, **kwargs)
if _f.__doc__ is None:
Expand Down Expand Up @@ -97,6 +103,12 @@ def _fix_promotion(x1, x2, only_scalar=True):
_py_scalars = (bool, int, float, complex)


def _scalar_to_tensor(x, other):
if isinstance(x, _py_scalars) and isinstance(other, torch.Tensor):
return torch.asarray(x, dtype=result_type(x, other), device=other.device)
return x


def result_type(*arrays_and_dtypes: Array | DType | complex) -> DType:
num = len(arrays_and_dtypes)

Expand Down Expand Up @@ -161,29 +173,33 @@ def can_cast(from_: DType | Array, to: DType, /) -> bool:
# Two-arg elementwise functions
# These require a wrapper to do the correct type promotion on 0-D tensors
add = _two_arg(torch.add)
atan2 = _two_arg(torch.atan2)
atan2 = _two_arg(torch.atan2, scalar_x1=True, scalar_x2=True)
bitwise_and = _two_arg(torch.bitwise_and)
bitwise_left_shift = _two_arg(torch.bitwise_left_shift)
bitwise_or = _two_arg(torch.bitwise_or)
bitwise_right_shift = _two_arg(torch.bitwise_right_shift)
bitwise_xor = _two_arg(torch.bitwise_xor)
copysign = _two_arg(torch.copysign)
copysign = _two_arg(torch.copysign, scalar_x1=True)
divide = _two_arg(torch.divide)
# Also a rename. torch.equal does not broadcast
equal = _two_arg(torch.eq)
equal = _two_arg(torch.eq, scalar_x1=True)
floor_divide = _two_arg(torch.floor_divide)
greater = _two_arg(torch.greater)
greater_equal = _two_arg(torch.greater_equal)
hypot = _two_arg(torch.hypot)
less = _two_arg(torch.less)
less_equal = _two_arg(torch.less_equal)
logaddexp = _two_arg(torch.logaddexp)
# logical functions are not included here because they only accept bool in the
# spec, so type promotion is irrelevant.
maximum = _two_arg(torch.maximum)
minimum = _two_arg(torch.minimum)
greater = _two_arg(torch.greater, scalar_x1=True)
greater_equal = _two_arg(torch.greater_equal, scalar_x1=True)
hypot = _two_arg(torch.hypot, scalar_x1=True, scalar_x2=True)
less = _two_arg(torch.less, scalar_x1=True)
less_equal = _two_arg(torch.less_equal, scalar_x1=True)
logaddexp = _two_arg(torch.logaddexp, scalar_x1=True, scalar_x2=True)
# logical functions only accept bool in the spec, so type promotion is
# irrelevant, but torch does not accept Python scalars for them.
logical_and = _two_arg(torch.logical_and, scalar_x1=True, scalar_x2=True)
logical_or = _two_arg(torch.logical_or, scalar_x1=True, scalar_x2=True)
logical_xor = _two_arg(torch.logical_xor, scalar_x1=True, scalar_x2=True)
maximum = _two_arg(torch.maximum, scalar_x1=True, scalar_x2=True)
minimum = _two_arg(torch.minimum, scalar_x1=True, scalar_x2=True)
multiply = _two_arg(torch.multiply)
not_equal = _two_arg(torch.not_equal)
nextafter = _two_arg(torch.nextafter, scalar_x1=True, scalar_x2=True)
not_equal = _two_arg(torch.not_equal, scalar_x1=True)
pow = _two_arg(torch.pow)
remainder = _two_arg(torch.remainder)
subtract = _two_arg(torch.subtract)
Expand Down Expand Up @@ -951,8 +967,9 @@ def meshgrid(*arrays: Array, indexing: Literal['xy', 'ij'] = 'xy') -> tuple[Arra
'bitwise_right_shift', 'bitwise_xor', 'copysign', 'count_nonzero',
'diff', 'divide', 'round',
'equal', 'floor_divide', 'greater', 'greater_equal', 'hypot',
'less', 'less_equal', 'logaddexp', 'maximum', 'minimum',
'multiply', 'not_equal', 'pow', 'remainder', 'subtract', 'max',
'less', 'less_equal', 'logaddexp', 'logical_and', 'logical_or',
'logical_xor', 'maximum', 'minimum', 'multiply', 'nextafter',
'not_equal', 'pow', 'remainder', 'subtract', 'max',
'min', 'clip', 'unstack', 'cumulative_sum', 'cumulative_prod', 'sort',
'argsort', 'prod', 'sum', 'any', 'all', 'mean', 'std', 'var', 'concat',
'squeeze', 'broadcast_to', 'flip', 'roll', 'nonzero', 'where', 'reshape',
Expand Down
42 changes: 42 additions & 0 deletions tests/test_torch.py
Original file line number Diff line number Diff line change
Expand Up @@ -193,3 +193,45 @@ def foo(x):
x = torch.arange(3)
y = bar(x)
assert xp.all(y == x**2)


@pytest.mark.parametrize(
"func_name",
[
"atan2", "copysign", "hypot", "logaddexp", "maximum", "minimum",
"nextafter", "equal", "not_equal", "less", "less_equal", "greater",
"greater_equal",
]
)
def test_binary_python_scalars(func_name):
# https://github.com/data-apis/array-api-compat/issues/271
func = getattr(xp, func_name)
x = xp.asarray([-1.5, 0.0, 2.0], dtype=xp.float32)
s = 0.5
y = xp.full(x.shape, s, dtype=xp.float32)

assert xp.all(func(x, s) == func(x, y))
assert xp.all(func(s, x) == func(y, x))
assert func(x, s).dtype == func(s, x).dtype == func(x, y).dtype


@pytest.mark.parametrize("func_name", ["logical_and", "logical_or", "logical_xor"])
def test_logical_python_scalars(func_name):
func = getattr(xp, func_name)
x = xp.asarray([True, False])
for s in [True, False]:
y = xp.full(x.shape, s)
assert xp.all(func(x, s) == func(x, y))
assert xp.all(func(s, x) == func(y, x))


def test_binary_python_scalars_promotion():
# Python scalars take the dtype of the array argument
x = xp.asarray([1, 5], dtype=xp.int8)
assert xp.maximum(x, 3).dtype == xp.int8
assert xp.minimum(3, x).dtype == xp.int8
assert xp.all(xp.maximum(3, x) == xp.asarray([3, 5], dtype=xp.int8))

x = xp.asarray([1., 5.], dtype=xp.float64)
assert xp.maximum(3.0, x).dtype == xp.float64
assert xp.all(xp.less(3, x) == xp.asarray([False, True]))
23 changes: 0 additions & 23 deletions torch-xfails.txt
Original file line number Diff line number Diff line change
Expand Up @@ -147,26 +147,3 @@ array_api_tests/test_signatures.py::test_array_method_signature[__dlpack__]

# broadcast_shapes emits a RuntimeError where the spec says ValueError
array_api_tests/test_data_type_functions.py::TestBroadcastShapes::test_error


# 2024.12 support: binary functions reject python scalar arguments
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_real[atan2]
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_real[copysign]
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_real[divide]
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_real[hypot]
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_real[logaddexp]
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_real[maximum]
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_real[minimum]
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_real[nextafter]

# https://github.com/pytorch/pytorch/issues/149815
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_real[equal]
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_real[not_equal]
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_real[less]
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_real[less_equal]
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_real[greater]
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_real[greater_equal]

array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_bool[logical_and]
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_bool[logical_or]
array_api_tests/test_operators_and_elementwise_functions.py::test_binary_with_scalars_bool[logical_xor]
Loading