Skip to content
Merged
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
6 changes: 4 additions & 2 deletions pytensor/link/jax/dispatch/sort.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
26 changes: 6 additions & 20 deletions pytensor/link/mlx/dispatch/sort.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand All @@ -18,35 +15,24 @@ 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


@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
14 changes: 7 additions & 7 deletions pytensor/link/numba/dispatch/sort.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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 argsort_f(X):
Y = np.swapaxes(X, axis, 0)
result = np.empty_like(Y, dtype="int64")

Expand All @@ -62,4 +62,4 @@ def argort_f(X, axis):
result = np.swapaxes(result, 0, axis)
return result

return argort_f
return argsort_f
6 changes: 4 additions & 2 deletions pytensor/link/pytorch/dispatch/sort.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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
117 changes: 48 additions & 69 deletions pytensor/tensor/sort.py
Original file line number Diff line number Diff line change
@@ -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


Expand All @@ -28,59 +27,54 @@ 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
# numpy.ogrid[0: a.shape[0], 0: a.shape[1], 0: a.shape[2]]
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".
Expand All @@ -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):
Expand All @@ -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
Expand All @@ -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):
Expand All @@ -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):
Expand All @@ -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)
13 changes: 2 additions & 11 deletions tests/link/mlx/test_sort.py
Original file line number Diff line number Diff line change
Expand Up @@ -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])
Expand All @@ -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")
Expand Down
Loading
Loading