| 18 | |
| 19 | |
| 20 | def _split_channels(num_chan, num_groups, split_op='equal'): |
| 21 | if split_op == 'equal': |
| 22 | # split = [num_chan // num_groups for _ in range(num_groups)] |
| 23 | split = [(num_chan // num_groups) // 4 * 4 for _ in range(num_groups)] |
| 24 | split[0] += num_chan - sum(split) |
| 25 | elif split_op == 'exp': |
| 26 | split = [int((num_chan * math.pow(2, -i))) // 4 * 4 for i in range(1, num_groups + 1)] |
| 27 | split[0] += num_chan - sum(split) |
| 28 | else: |
| 29 | raise ValueError('Unknown split option: {}'.format(split_op)) |
| 30 | return split |
| 31 | |
| 32 | |
| 33 | class Conv2dSame(nn.Conv2d): |