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

Function _get_pool_param

imperative/python/megengine/xla/rules/nn.py:635–653  ·  view source on GitHub ↗
(kernel_hw, stride_hw, padding_hw, tensor_format)

Source from the content-addressed store, hash-verified

633
634
635def _get_pool_param(kernel_hw, stride_hw, padding_hw, tensor_format):
636 assert len(kernel_hw) == 2 and len(stride_hw) == 2 and len(padding_hw) == 2
637 # for backward, the tensor format is str
638 if not isinstance(tensor_format, str):
639 tensor_format = str(tensor_format)
640
641 stride, kernel, padding = None, None, None
642 if tensor_format in str(mops.AdaptivePooling.Format.NCHW):
643 stride = (1, 1, *stride_hw)
644 kernel = (1, 1, *kernel_hw)
645 padding = (0, 0, *padding_hw)
646 elif tensor_format in str(mops.AdaptivePooling.Format.NHWC):
647 stride = (1, *stride_hw, 1)
648 kernel = (1, *kernel_hw, 1)
649 padding = (0, *padding_hw, 0)
650 else:
651 assert False, f"adaptive pooling only nchw or nhwc, get {tensor_format}"
652
653 return kernel, stride, padding
654
655
656@register_lower_rule(mops.Pooling)

Callers 2

pooling_lowerFunction · 0.85
pooling_backward_lowerFunction · 0.85

Calls 1

strFunction · 0.85

Tested by

no test coverage detected