(kernel_hw, stride_hw, padding_hw, tensor_format)
| 633 | |
| 634 | |
| 635 | def _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) |
no test coverage detected