From ff1cac69b13f47c378819756cdcac2b44dca06b6 Mon Sep 17 00:00:00 2001 From: Raashish Aggarwal <94279692+raashish1601@users.noreply.github.com> Date: Wed, 7 Oct 2026 23:50:51 +0530 Subject: [PATCH 1/2] Make Sort and ArgSort axis a static Op property --- pytensor/link/jax/dispatch/sort.py | 6 +- pytensor/link/mlx/dispatch/sort.py | 26 ++---- pytensor/link/numba/dispatch/sort.py | 12 +-- pytensor/link/pytorch/dispatch/sort.py | 6 +- pytensor/tensor/sort.py | 117 ++++++++++--------------- tests/link/mlx/test_sort.py | 13 +-- tests/link/numba/test_sort.py | 13 +-- tests/tensor/test_sort.py | 83 +++++++++--------- 8 files changed, 117 insertions(+), 159 deletions(-) diff --git a/pytensor/link/jax/dispatch/sort.py b/pytensor/link/jax/dispatch/sort.py index 00a733389f..02b357eaa2 100644 --- a/pytensor/link/jax/dispatch/sort.py +++ b/pytensor/link/jax/dispatch/sort.py @@ -7,8 +7,9 @@ @jax_funcify.register(SortOp) def jax_funcify_Sort(op, **kwargs): stable = op.kind == "stable" + axis = op.axis - def sort(arr, axis): + def sort(arr): return jnp.sort(arr, axis=axis, stable=stable) return sort @@ -17,8 +18,9 @@ def sort(arr, axis): @jax_funcify.register(ArgSortOp) def jax_funcify_ArgSort(op, **kwargs): stable = op.kind == "stable" + axis = op.axis - def argsort(arr, axis): + def argsort(arr): return jnp.argsort(arr, axis=axis, stable=stable) return argsort diff --git a/pytensor/link/mlx/dispatch/sort.py b/pytensor/link/mlx/dispatch/sort.py index 231836c9a4..c5e84334cf 100644 --- a/pytensor/link/mlx/dispatch/sort.py +++ b/pytensor/link/mlx/dispatch/sort.py @@ -3,9 +3,6 @@ import mlx.core as mx from pytensor.link.mlx.dispatch.basic import convert_dtype_to_mlx, mlx_funcify -from pytensor.link.mlx.dispatch.tensor_basic import coerce_to_int -from pytensor.tensor.basic import get_scalar_constant_value -from pytensor.tensor.exceptions import NotScalarConstantError from pytensor.tensor.sort import ArgSortOp, SortOp @@ -18,24 +15,13 @@ def _warn_unsupported_kind(op, name): ) -def _static_axis(node): - try: - return int(get_scalar_constant_value(node.inputs[1])) - except NotScalarConstantError: - return None - - -def _resolve_axis(static_axis, axis): - return coerce_to_int(axis) if static_axis is None else static_axis - - @mlx_funcify.register(SortOp) def mlx_funcify_Sort(op, node, **kwargs): _warn_unsupported_kind(op, "sort") - static_axis = _static_axis(node) + axis = op.axis - def sort(x, axis): - return mx.sort(x, axis=_resolve_axis(static_axis, axis)) + def sort(x): + return mx.sort(x, axis=axis) return sort @@ -43,10 +29,10 @@ def sort(x, axis): @mlx_funcify.register(ArgSortOp) def mlx_funcify_ArgSort(op, node, **kwargs): _warn_unsupported_kind(op, "argsort") - static_axis = _static_axis(node) + axis = op.axis out_dtype = convert_dtype_to_mlx(node.outputs[0].dtype) - def argsort(x, axis): - return mx.argsort(x, axis=_resolve_axis(static_axis, axis)).astype(out_dtype) + def argsort(x): + return mx.argsort(x, axis=axis).astype(out_dtype) return argsort diff --git a/pytensor/link/numba/dispatch/sort.py b/pytensor/link/numba/dispatch/sort.py index 51cd411b5c..38d6a40c03 100644 --- a/pytensor/link/numba/dispatch/sort.py +++ b/pytensor/link/numba/dispatch/sort.py @@ -20,10 +20,10 @@ def numba_funcify_SortOp(op, node, **kwargs): UserWarning, ) - @numba_basic.numba_njit - def sort_f(a, axis): - axis = axis.item() + axis = op.axis + @numba_basic.numba_njit + def sort_f(a): a_swapped = np.swapaxes(a, axis, -1) a_sorted = np.sort(a_swapped) a_sorted_swapped = np.swapaxes(a_sorted, -1, axis) @@ -47,10 +47,10 @@ def numba_funcify_ArgSortOp(op, node, **kwargs): UserWarning, ) - @numba_basic.numba_njit - def argort_f(X, axis): - axis = axis.item() + axis = op.axis + @numba_basic.numba_njit + def argort_f(X): Y = np.swapaxes(X, axis, 0) result = np.empty_like(Y, dtype="int64") diff --git a/pytensor/link/pytorch/dispatch/sort.py b/pytensor/link/pytorch/dispatch/sort.py index 95e24c4fe3..242fc7a643 100644 --- a/pytensor/link/pytorch/dispatch/sort.py +++ b/pytensor/link/pytorch/dispatch/sort.py @@ -7,8 +7,9 @@ @pytorch_funcify.register(SortOp) def pytorch_funcify_Sort(op, **kwargs): stable = op.kind == "stable" + axis = op.axis - def sort(arr, axis): + def sort(arr): sorted, _ = torch.sort(arr, dim=axis, stable=stable) return sorted @@ -18,8 +19,9 @@ def sort(arr, axis): @pytorch_funcify.register(ArgSortOp) def pytorch_funcify_ArgSort(op, **kwargs): stable = op.kind == "stable" + axis = op.axis - def argsort(arr, axis): + def argsort(arr): return torch.argsort(arr, dim=axis, stable=stable) return argsort diff --git a/pytensor/tensor/sort.py b/pytensor/tensor/sort.py index c911be988d..32a85c4b19 100644 --- a/pytensor/tensor/sort.py +++ b/pytensor/tensor/sort.py @@ -1,12 +1,11 @@ import typing import numpy as np +from numpy.lib.array_utils import normalize_axis_index -from pytensor.gradient import grad_undefined from pytensor.graph.basic import Apply from pytensor.graph.op import Op -from pytensor.tensor.basic import arange, as_tensor_variable, switch -from pytensor.tensor.math import eq, ge +from pytensor.tensor.basic import _validate_axis_argument, arange, as_tensor_variable from pytensor.tensor.type import TensorType @@ -28,51 +27,46 @@ def _parse_sort_args(kind: KIND | None, order, stable: bool | None) -> KIND: return kind +def _validate_sort_axis(axis, op_name: str) -> int: + int_axis: int = _validate_axis_argument(axis, op_name) + if int_axis < 0: + raise ValueError(f"{op_name} axis must be non-negative, got {int_axis}.") + return int_axis + + class SortOp(Op): """ This class is a wrapper for numpy sort function. """ - __props__ = ("kind",) + __props__ = ("kind", "axis") - def __init__(self, kind: KIND): + def __init__(self, kind: KIND, axis: int): self.kind = kind + self.axis = _validate_sort_axis(axis, "Sort") - def make_node(self, input, axis=-1): + def make_node(self, input): input = as_tensor_variable(input) - axis = as_tensor_variable(axis, ndim=0, dtype=int) - if axis.type.numpy_dtype.kind != "i": - raise ValueError( - f"Sort axis must have an integer dtype, got {axis.type.dtype}" - ) + if self.axis >= input.type.ndim: + raise np.exceptions.AxisError(self.axis, input.type.ndim) out_type = input.type() - return Apply(self, [input, axis], [out_type]) + return Apply(self, [input], [out_type]) def perform(self, node, inputs, output_storage): - a, axis = inputs + [a] = inputs z = output_storage[0] - z[0] = np.sort(a, axis, self.kind) + z[0] = np.sort(a, self.axis, self.kind) def infer_shape(self, node, inputs_shapes): - assert node.inputs[0].ndim == node.outputs[0].ndim - assert inputs_shapes[1] == () return [inputs_shapes[0]] def pullback(self, inputs, outputs, output_grads): - a, axis = inputs - indices = self.__get_argsort_indices(a, axis) - inp_grad = output_grads[0][tuple(indices)] - axis_grad = grad_undefined( - self, - 1, - axis, - "The gradient of sort is not defined " - "with respect to the integer axes itself", - ) - return [inp_grad, axis_grad] + [a] = inputs + indices = self.__get_argsort_indices(a) + return [output_grads[0][tuple(indices)]] - def __get_expanded_dim(self, a, axis, i): + def __get_expanded_dim(self, a, i): index_shape = [1] * a.ndim index_shape[i] = a.shape[i] # it's a way to emulate @@ -80,7 +74,7 @@ def __get_expanded_dim(self, a, axis, i): index_val = arange(a.shape[i]).reshape(index_shape) return index_val - def __get_argsort_indices(self, a, axis): + def __get_argsort_indices(self, a): """ Calculates indices which can be used to reverse sorting operation of "a" tensor along "axis". @@ -94,19 +88,13 @@ def __get_argsort_indices(self, a, axis): # The goal is to get gradient wrt input from gradient # wrt sort(input, axis) - idx = argsort(a, axis, kind=self.kind) + idx = argsort(a, self.axis, kind=self.kind) # rev_idx is the reverse of previous argsort operation - rev_idx = argsort(idx, axis, kind=self.kind) - indices = [] - axis_data = switch(ge(axis.data, 0), axis.data, a.ndim + axis.data) - for i in range(a.ndim): - index_val = switch( - eq(i, axis_data), - rev_idx, - self.__get_expanded_dim(a, axis, i), - ) - indices.append(index_val) - return indices + rev_idx = argsort(idx, self.axis, kind=self.kind) + return [ + rev_idx if i == self.axis else self.__get_expanded_dim(a, i) + for i in range(a.ndim) + ] """ def pushforward(self, inputs, outputs, eval_points): @@ -129,9 +117,9 @@ def sort( ---------- a: TensorVariable Tensor to be sorted - axis: TensorVariable + axis: int, optional Axis along which to sort. If None, the array is flattened before - sorting. + sorting. Must be a constant. kind: {'quicksort', 'mergesort', 'heapsort' 'stable'}, optional Sorting algorithm. Default is 'quicksort' unless stable is defined. order: list, optional @@ -146,11 +134,12 @@ def sort( """ kind = _parse_sort_args(kind, order, stable) - + a = as_tensor_variable(a) if axis is None: a = a.flatten() axis = 0 - return SortOp(kind)(a, axis) + axis = normalize_axis_index(_validate_axis_argument(axis, "sort"), a.type.ndim) + return SortOp(kind, axis)(a) class ArgSortOp(Op): @@ -159,49 +148,37 @@ class ArgSortOp(Op): """ - __props__ = ("kind",) + __props__ = ("kind", "axis") - def __init__(self, kind: KIND): + def __init__(self, kind: KIND, axis: int): self.kind = kind + self.axis = _validate_sort_axis(axis, "ArgSort") - def make_node(self, input, axis=-1): + def make_node(self, input): input = as_tensor_variable(input) - axis = as_tensor_variable(axis, ndim=0, dtype=int) - if axis.type.numpy_dtype.kind != "i": - raise ValueError( - f"ArgSort axis must have an integer dtype, got {axis.type.dtype}" - ) + if self.axis >= input.type.ndim: + raise np.exceptions.AxisError(self.axis, input.type.ndim) return Apply( self, - [input, axis], + [input], [TensorType(dtype="int64", shape=input.type.shape)()], ) def perform(self, node, inputs, output_storage): - a, axis = inputs + [a] = inputs z = output_storage[0] z[0] = np.asarray( - np.argsort(a, axis, self.kind), + np.argsort(a, self.axis, self.kind), dtype=node.outputs[0].dtype, ) def infer_shape(self, node, inputs_shapes): - assert node.inputs[0].ndim == node.outputs[0].ndim - assert inputs_shapes[1] == () return [inputs_shapes[0]] def pullback(self, inputs, outputs, output_grads): # No grad defined for integers. - inp, axis = inputs - inp_grad = inp.zeros_like() - axis_grad = grad_undefined( - self, - 1, - axis, - "argsort is not defined for non-integer axes so" - " argsort(x, axis+eps) is undefined", - ) - return [inp_grad, axis_grad] + [inp] = inputs + return [inp.zeros_like()] """ def pushforward(self, inputs, outputs, eval_points): @@ -228,7 +205,9 @@ def argsort( """ kind = _parse_sort_args(kind, order, stable) + a = as_tensor_variable(a) if axis is None: a = a.flatten() axis = 0 - return ArgSortOp(kind)(a, axis) + axis = normalize_axis_index(_validate_axis_argument(axis, "argsort"), a.type.ndim) + return ArgSortOp(kind, axis)(a) diff --git a/tests/link/mlx/test_sort.py b/tests/link/mlx/test_sort.py index f86f57b021..ec66f612f4 100644 --- a/tests/link/mlx/test_sort.py +++ b/tests/link/mlx/test_sort.py @@ -2,8 +2,8 @@ import pytest from pytensor.tensor.sort import argsort, sort -from pytensor.tensor.type import iscalar, matrix, tensor3 -from tests.link.mlx.test_basic import compare_mlx_and_py, mlx_mode_no_compile +from pytensor.tensor.type import matrix, tensor3 +from tests.link.mlx.test_basic import compare_mlx_and_py @pytest.mark.parametrize("axis", [None, 0, -1, -2]) @@ -16,15 +16,6 @@ def test_sort(func, axis): assert np.asarray(res).dtype == out.dtype -@pytest.mark.parametrize("func", (sort, argsort)) -def test_sort_symbolic_axis(func): - x = matrix("x", shape=(2, 3), dtype="float64") - axis = iscalar("axis") - out = func(x, axis=axis) - arr = np.random.default_rng(0).permutation(np.arange(6.0)).reshape(2, 3) - compare_mlx_and_py([x, axis], [out], [arr, 1], mlx_mode=mlx_mode_no_compile) - - def test_sort_invalid_kind_warning(): x = matrix("x", shape=(2, 2), dtype="float64") z = sort(x, axis=-1, kind="mergesort") diff --git a/tests/link/numba/test_sort.py b/tests/link/numba/test_sort.py index 06861ebb98..324a486adb 100644 --- a/tests/link/numba/test_sort.py +++ b/tests/link/numba/test_sort.py @@ -4,7 +4,7 @@ import pytest from pytensor import tensor as pt -from pytensor.tensor.sort import ArgSortOp, SortOp +from pytensor.tensor.sort import argsort, sort from tests.link.numba.test_basic import compare_numba_and_py @@ -28,10 +28,7 @@ ) def test_Sort(x_test, axis, kind, exc): x = pt.as_tensor(x_test).type("x") - if axis: - g = SortOp(kind)(x, axis) - else: - g = SortOp(kind)(x) + g = sort(x, axis=axis, kind=kind) cm = contextlib.suppress() if not exc else pytest.warns(exc) @@ -62,11 +59,7 @@ def test_ArgSort(x_test, axis, kind, exc): np.random.shuffle(x_test) x_test = np.reshape(x_test, (5, 5, 5, 5)) x = pt.as_tensor(x_test).type("x") - - if axis: - g = ArgSortOp(kind)(x, axis) - else: - g = ArgSortOp(kind)(x) + g = argsort(x, axis=axis, kind=kind) cm = contextlib.suppress() if not exc else pytest.warns(exc) diff --git a/tests/tensor/test_sort.py b/tests/tensor/test_sort.py index 319fbfe560..d216e48ae7 100644 --- a/tests/tensor/test_sort.py +++ b/tests/tensor/test_sort.py @@ -2,16 +2,15 @@ import pytest import pytensor +from pytensor.tensor.basic import constant from pytensor.tensor.sort import ArgSortOp, SortOp, argsort, sort from pytensor.tensor.type import ( dmatrix, dvector, float_dtypes, - fscalar, integer_dtypes, lscalar, matrix, - scalar, ) from tests import unittest_tools as utt @@ -32,11 +31,15 @@ def setup_method(self): self.m_val = self.rng.random((3, 2)) self.v_val = self.rng.random(4) - def test_invalid_axis_dtype(self): - with pytest.raises( - ValueError, match="Sort axis must have an integer dtype, got float32" - ): - sort(dmatrix(), fscalar()) + def test_invalid_axis(self): + with pytest.raises(TypeError, match="axis of sort must be a constant integer"): + sort(dmatrix(), lscalar()) + with pytest.raises(np.exceptions.AxisError): + sort(dmatrix(), 2) + with pytest.raises(ValueError, match="Sort axis must be non-negative"): + SortOp("quicksort", -1) + with pytest.raises(np.exceptions.AxisError): + SortOp("quicksort", 2)(dmatrix()) def test1(self): a = dmatrix() @@ -46,11 +49,11 @@ def test1(self): def test2(self): a = dmatrix() - axis = scalar(dtype="int64") - w = sort(a, axis) - f = pytensor.function([a, axis], w) for axis_val in 0, 1: - gv = f(self.m_val, axis_val) + w = sort(a, constant(axis_val, dtype="int64")) + assert w.owner.op.axis == axis_val + f = pytensor.function([a], w) + gv = f(self.m_val) gt = np.sort(self.m_val, axis_val) utt.assert_allclose(gv, gt) @@ -64,21 +67,22 @@ def test3(self): def test4(self): a = dmatrix() - axis = scalar(dtype="int8") - l = sort(a, axis, "mergesort") - f = pytensor.function([a, axis], l) - for axis_val in 0, 1: - gv = f(self.m_val, np.array(axis_val, dtype="int8")) - gt = np.sort(self.m_val, np.array(axis_val, dtype="int8")) + for axis_val in -2, -1: + l = sort(a, np.int8(axis_val), "mergesort") + assert l.owner.op.axis == axis_val + 2 + f = pytensor.function([a], l) + gv = f(self.m_val) + gt = np.sort(self.m_val, axis_val) utt.assert_allclose(gv, gt) def test5(self): - a1 = SortOp("mergesort") - a2 = SortOp("quicksort") + a1 = SortOp("mergesort", 0) + a2 = SortOp("quicksort", 0) assert a1 != a2 - assert a1 == SortOp("mergesort") - assert a2 == SortOp("quicksort") + assert a1 == SortOp("mergesort", 0) + assert a2 == SortOp("quicksort", 0) + assert a1 != SortOp("mergesort", 1) def test_None(self): a = dmatrix() @@ -188,11 +192,11 @@ def test_argsort(): # Example 2 a = dmatrix() - axis = lscalar() - w = argsort(a, axis) - f = pytensor.function([a, axis], w) for axis_val in 0, 1: - gv = f(m_val, axis_val) + w = argsort(a, constant(axis_val, dtype="int64")) + assert w.owner.op.axis == axis_val + f = pytensor.function([a], w) + gv = f(m_val) gt = np.argsort(m_val, axis_val) utt.assert_allclose(gv, gt) @@ -206,20 +210,21 @@ def test_argsort(): # Example 4 a = dmatrix() - axis = scalar(dtype="int8") - l = argsort(a, axis, "mergesort") - f = pytensor.function([a, axis], l) - for axis_val in 0, 1: - gv = f(m_val, np.array(axis_val, dtype="int8")) - gt = np.argsort(m_val, np.array(axis_val, dtype="int8")) + for axis_val in -2, -1: + l = argsort(a, np.int8(axis_val), "mergesort") + assert l.owner.op.axis == axis_val + 2 + f = pytensor.function([a], l) + gv = f(m_val) + gt = np.argsort(m_val, axis_val) utt.assert_allclose(gv, gt) # Example 5 - a1 = ArgSortOp("mergesort") - a2 = ArgSortOp("quicksort") + a1 = ArgSortOp("mergesort", 0) + a2 = ArgSortOp("quicksort", 0) assert a1 != a2 - assert a1 == ArgSortOp("mergesort") - assert a2 == ArgSortOp("quicksort") + assert a1 == ArgSortOp("mergesort", 0) + assert a2 == ArgSortOp("quicksort", 0) + assert a1 != ArgSortOp("mergesort", 1) # Example 6: Testing axis=None a = dmatrix() @@ -229,10 +234,10 @@ def test_argsort(): gt = np.argsort(m_val, None) utt.assert_allclose(gv, gt) - with pytest.raises( - ValueError, match="ArgSort axis must have an integer dtype, got float32" - ): - argsort(dmatrix(), fscalar()) + with pytest.raises(TypeError, match="axis of argsort must be a constant integer"): + argsort(dmatrix(), lscalar()) + with pytest.raises(ValueError, match="ArgSort axis must be non-negative"): + ArgSortOp("quicksort", -1) def test_argsort_grad(): From 694d2c417e238f7745017ae91e61fe647de7b611 Mon Sep 17 00:00:00 2001 From: Raashish Aggarwal <94279692+raashish1601@users.noreply.github.com> Date: Fri, 9 Oct 2026 22:15:29 +0530 Subject: [PATCH 2/2] Rename argort_f and use int axes in sort tests --- pytensor/link/numba/dispatch/sort.py | 4 ++-- tests/tensor/test_sort.py | 11 ++++------- 2 files changed, 6 insertions(+), 9 deletions(-) diff --git a/pytensor/link/numba/dispatch/sort.py b/pytensor/link/numba/dispatch/sort.py index 38d6a40c03..54518e682a 100644 --- a/pytensor/link/numba/dispatch/sort.py +++ b/pytensor/link/numba/dispatch/sort.py @@ -50,7 +50,7 @@ def numba_funcify_ArgSortOp(op, node, **kwargs): axis = op.axis @numba_basic.numba_njit - def argort_f(X): + def argsort_f(X): Y = np.swapaxes(X, axis, 0) result = np.empty_like(Y, dtype="int64") @@ -62,4 +62,4 @@ def argort_f(X): result = np.swapaxes(result, 0, axis) return result - return argort_f + return argsort_f diff --git a/tests/tensor/test_sort.py b/tests/tensor/test_sort.py index d216e48ae7..88c0668097 100644 --- a/tests/tensor/test_sort.py +++ b/tests/tensor/test_sort.py @@ -40,6 +40,8 @@ def test_invalid_axis(self): SortOp("quicksort", -1) with pytest.raises(np.exceptions.AxisError): SortOp("quicksort", 2)(dmatrix()) + # A constant scalar is still accepted and normalized, as for join + assert sort(dmatrix(), constant(-1)).owner.op == SortOp("quicksort", 1) def test1(self): a = dmatrix() @@ -50,7 +52,7 @@ def test1(self): def test2(self): a = dmatrix() for axis_val in 0, 1: - w = sort(a, constant(axis_val, dtype="int64")) + w = sort(a, axis_val) assert w.owner.op.axis == axis_val f = pytensor.function([a], w) gv = f(self.m_val) @@ -193,7 +195,7 @@ def test_argsort(): # Example 2 a = dmatrix() for axis_val in 0, 1: - w = argsort(a, constant(axis_val, dtype="int64")) + w = argsort(a, axis_val) assert w.owner.op.axis == axis_val f = pytensor.function([a], w) gv = f(m_val) @@ -234,11 +236,6 @@ def test_argsort(): gt = np.argsort(m_val, None) utt.assert_allclose(gv, gt) - with pytest.raises(TypeError, match="axis of argsort must be a constant integer"): - argsort(dmatrix(), lscalar()) - with pytest.raises(ValueError, match="ArgSort axis must be non-negative"): - ArgSortOp("quicksort", -1) - def test_argsort_grad(): rng = np.random.default_rng(seed=utt.fetch_seed())