Uh oh!
There was an error while loading. Please reload this page.
ENH: add partition and argpartition functions - #449
Conversation
partition and argpartition functionslucascolley
commented
Oct 1, 2025
see array-api-extra/tests/test_funcs.py Line 1183 in ca20f03 |
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
lucascolley
left a comment
There was a problem hiding this comment.
thanks @cakedev0 !
Looks like there is also a merge conflict now.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
| kth += 1 # HACK: we use a non-specified behavior of torch.topk: | ||
| # in `a_left`, the element in the last position is the max | ||
| a_left, indices = xp.topk(a, kth, dim=-1, largest=False, sorted=False) |
There was a problem hiding this comment.
hmmm, I would rather not rely on undocumented behaviour. Is there an alternative?
There was a problem hiding this comment.
Fair ^^
Three options:
- add an
assert a_left.max() == a_left[k] - We can just re-run the same logic with
kth=1andlargest=True. Impact on perfs is probably 10 to 100% slower depending on the input. But it doens't add a lot of logic - We can do a
if a_left.max() != a_left[k]: swap_max_with_last_element(a_left, axis=-1)=> requires to implementswap_max_with_last_element(and the equivalent for argsort).
I vote for 1 because I'm lazy but I like perf :p
There was a problem hiding this comment.
Edit: wait I need to rethink something about numpy.partition specs...
There was a problem hiding this comment.
So! I rewrote entirely this section, it now relies on torch.kthvalue and is very aligned with numpy's behavior.
On a side note: the description of the behavior of the partition function in numpy is fairly blurry when the k-th element has duplicates... In practice, numpy does a tree-way partitioning: <, == and >. I reproduced this behavior in my new torch implementation, but jax doesn't (I tried to test the tree-way partitioning and jax fails it...).
I will maybe open an issue on numpy to ask for some clarification.
There was a problem hiding this comment.
On a side note: the description of the behavior of the partition function in numpy is fairly blurry when the k-th element has duplicates... In practice, numpy does a tree-way partitioning: <, == and >. I reproduced this behavior in my new torch implementation, but jax doesn't (I tried to test the tree-way partitioning and jax fails it...).
It might be worth contributing this consideration to the array API spec discussion:
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
lucascolley
left a comment
There was a problem hiding this comment.
thanks @cakedev0, looks close!
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.
Co-authored-by: Lucas Colley <lucas.colley8@gmail.com>
cakedev0
commented
Oct 3, 2025
Thanks for the reactive and helpful reviews @lucascolley Sorry for the numpy docs style details, I'll make sure to read the doc about this carefully before opening another PR 😉 |
lucascolley
left a comment
There was a problem hiding this comment.
thanks @cakedev0, let's merge it! And thanks for taking a look too @ogrisel .
Would be great to follow-up on #449 (comment).
Are you interested in taking over gh-341 next?
Uh oh!
There was an error while loading. Please reload this page.
ogrisel
commented
Oct 6, 2025
@lucascolley I see that this as been milestoned for 0.9.1, but those are new functions. Wouln'd it make sense to make them part of a new 0.10.0 release instead? |
lucascolley
commented
Oct 6, 2025
we use https://jacobtomlinson.dev/effver/, so no. Unless you see any problems with that! |
ogrisel
commented
Oct 6, 2025
Ok, as you wish. |
Closes#448
Supports nd-arrays.
The crux of this PR is the torch part: transforming
kthvalueoutputs to a partition/argpartition. The rest is mostly wrappers and checks.