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

Function pooling_lower

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

Source from the content-addressed store, hash-verified

655
656@register_lower_rule(mops.Pooling)
657def pooling_lower(ctx, *args: Union[HLOTensor, Sequence[HLOTensor]]):
658 assert len(args) == 1, f"pooling should have only 1 input, but give {len(args)}"
659 assert len(ctx.vars_in) == 1 and len(ctx.vars_out) == 1
660 assert (
661 args[0].ndim == 4
662 ), f"pooling only support 4d tensor, but give {args[0].shape}"
663 opr = ctx.op
664 kernel, stride, padding = _get_pool_param(
665 (opr.window_h, opr.window_w),
666 (opr.stride_h, opr.stride_w),
667 (opr.pad_h, opr.pad_w),
668 opr.format,
669 )
670
671 oshape, _ = ctx.vars_out[0].shape, ctx.vars_out[0].dtype
672 if opr.mode == mops.AdaptivePooling.Mode.AVERAGE:
673 return avgpooling(
674 args[0], stride, kernel, padding, count_include_pad=True, oshape=oshape
675 )
676 elif opr.mode == mops.AdaptivePooling.Mode.AVERAGE_COUNT_EXCLUDE_PADDING:
677 return avgpooling(
678 args[0], stride, kernel, padding, count_include_pad=False, oshape=oshape
679 )
680 else:
681 assert (
682 opr.mode == mops.AdaptivePooling.Mode.MAX
683 ), f"unknown adaptive pooling mode {opr.mode}"
684 return maxpooling(args[0], stride, kernel, padding, oshape=oshape)
685
686
687@register_lower_rule("PoolingBackwardV1")

Callers

nothing calls this directly

Calls 2

_get_pool_paramFunction · 0.85
avgpoolingFunction · 0.85

Tested by

no test coverage detected