| 54 | |
| 55 | |
| 56 | def create_pool2d(pool_type, kernel_size, stride=None, **kwargs): |
| 57 | stride = stride or kernel_size |
| 58 | padding = kwargs.pop('padding', '') |
| 59 | padding, is_dynamic = get_padding_value(padding, kernel_size, stride=stride, **kwargs) |
| 60 | if is_dynamic: |
| 61 | if pool_type == 'avg': |
| 62 | return AvgPool2dSame(kernel_size, stride=stride, **kwargs) |
| 63 | elif pool_type == 'max': |
| 64 | return MaxPool2dSame(kernel_size, stride=stride, **kwargs) |
| 65 | else: |
| 66 | assert False, f'Unsupported pool type {pool_type}' |
| 67 | else: |
| 68 | if pool_type == 'avg': |
| 69 | return nn.AvgPool2d(kernel_size, stride=stride, padding=padding, **kwargs) |
| 70 | elif pool_type == 'max': |
| 71 | return nn.MaxPool2d(kernel_size, stride=stride, padding=padding, **kwargs) |
| 72 | else: |
| 73 | assert False, f'Unsupported pool type {pool_type}' |