Repository navigation
Make Sort and ArgSort axis a static Op property - #2462
raashish1601 wants to merge 2 commits into
Conversation
| axis = op.axis | ||
|
|
||
| @numba_basic.numba_njit | ||
| def argort_f(X): |
There was a problem hiding this comment.
pre-existing typo, should be argsort_f not argort_f
| 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")) |
There was a problem hiding this comment.
hmm? why are we allowing pytensor constant axis?
There was a problem hiding this comment.
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.
| 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"): |
There was a problem hiding this comment.
This check seems redundant with the some of the ones above?
There was a problem hiding this comment.
Yes, TestSort.test_invalid_axis already covers it through the same helpers. Removed in 694d2c4.
Description
SortOpandArgSortOpnow takeaxisin the constructor and store it as an Op property, the same wayJoinandSplitdo since #2144. The Op checks the axis is non-negative and within the input's ndim.pt.sortandpt.argsortkeep their signatures. They accept a Python int or a constant scalar, normalize negative axes, and raise aTypeErrorfor a symbolic axis, using the same helper asjoin/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 aswitchover 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 throughsort/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