| 350 | |
| 351 | @register_lower_rule(mops.Argsort) |
| 352 | def argsort_lower(ctx, *args: Union[HLOTensor, Sequence[HLOTensor]]): |
| 353 | assert ( |
| 354 | len(args) == 1 and len(ctx.vars_in) == 1 and len(ctx.vars_out) == 2 |
| 355 | ), f"{len(args)}, {len(ctx.vars_in)}, {len(ctx.vars_out)}" |
| 356 | assert ctx.op.order in [ |
| 357 | mops.Argsort.Order.DESCENDING, |
| 358 | mops.Argsort.Order.ASCENDING, |
| 359 | ], f"{ctx.op.order}" |
| 360 | descending = ctx.op.order == mops.Argsort.Order.DESCENDING |
| 361 | axis = args[0].ndim - 1 # megengine only support sort in the last dimension |
| 362 | return argsort(args[0], axis, descending, is_stable=True) |
| 363 | |
| 364 | |
| 365 | @register_lower_rule("ArgsortBackward") |