Skip to content

Make Sort and ArgSort axis a static Op property - #2462

Open
raashish1601 wants to merge 2 commits into
pymc-devs:mainfrom
raashish1601:feature/2426-static-sort-axis
Open

raashish1601 wants to merge 2 commits into
pymc-devs:mainfrom
raashish1601:feature/2426-static-sort-axis

Conversation

@raashish1601

Copy link
Copy Markdown

Description

SortOp and ArgSortOp now take axis in the constructor and store it as an Op property, the same way Join and Split do since #2144. The Op checks the axis is non-negative and within the input's ndim.

pt.sort and pt.argsort keep their signatures. They accept a Python int or a constant scalar, normalize negative axes, and raise a TypeError for a symbolic axis, using the same helper as join/split.

With a static axis the JAX, PyTorch, Numba and MLX dispatches just read op.axis, the MLX constant-axis lookup goes away, and the gradient no longer builds a switch over every axis.

Tests: updated tests/tensor/test_sort.py (constant-variable axes, negative axes, invalid axes, Op equality with axis), the Numba tests now go through sort/argsort, and I removed the MLX symbolic-axis test since that's no longer supported. Ran the tensor, JAX, PyTorch and Numba sort tests locally. I couldn't run the MLX tests (no Apple hardware).

Related Issue

Checklist

Type of change

  • New feature / enhancement
  • Bug fix
  • Documentation
  • Maintenance
  • Other (please specify):

Comment thread pytensor/link/numba/dispatch/sort.py Outdated
axis = op.axis

@numba_basic.numba_njit
def argort_f(X):

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

pre-existing typo, should be argsort_f not argort_f

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed in 694d2c4, it is argsort_f now.

Comment thread tests/tensor/test_sort.py Outdated
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"))

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

hmm? why are we allowing pytensor constant axis?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

sort/argsort use the same _validate_axis_argument as join/split, which still accepts a constant scalar (test_symbolic_axis_rejected checks join(constant(-1), ...)), so code that passed a constant keeps working. The tests only used a constant because they used to build a symbolic scalar. In 694d2c4 they pass plain ints, and one line in test_invalid_axis checks the constant case like the join test does. If you would rather have sort only take Python ints, I can drop that.

Comment thread tests/tensor/test_sort.py Outdated
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"):

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This check seems redundant with the some of the ones above?

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, TestSort.test_invalid_axis already covers it through the same helpers. Removed in 694d2c4.

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.

Make Sort/ArgSort axis a static Op property

2 participants