MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / argsort_lower

Function argsort_lower

imperative/python/megengine/xla/rules/math.py:352–362  ·  view source on GitHub ↗
(ctx, *args: Union[HLOTensor, Sequence[HLOTensor]])

Source from the content-addressed store, hash-verified

350
351@register_lower_rule(mops.Argsort)
352def 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")

Callers

nothing calls this directly

Calls 1

argsortFunction · 0.70

Tested by

no test coverage detected