| 19 | |
| 20 | |
| 21 | class Model(ImageNetModel): |
| 22 | def __init__(self, depth, mode='resnet'): |
| 23 | self.mode = mode |
| 24 | basicblock = getattr(resnet_model, mode + '_basicblock', None) |
| 25 | bottleneck = getattr(resnet_model, mode + '_bottleneck', None) |
| 26 | self.num_blocks, self.block_func = { |
| 27 | 18: ([2, 2, 2, 2], basicblock), |
| 28 | 34: ([3, 4, 6, 3], basicblock), |
| 29 | 50: ([3, 4, 6, 3], bottleneck), |
| 30 | 101: ([3, 4, 23, 3], bottleneck), |
| 31 | 152: ([3, 8, 36, 3], bottleneck) |
| 32 | }[depth] |
| 33 | assert self.block_func is not None, \ |
| 34 | "(mode={}, depth={}) not implemented!".format(mode, depth) |
| 35 | |
| 36 | def get_logits(self, image): |
| 37 | with argscope([Conv2D, MaxPooling, GlobalAvgPooling, BatchNorm], data_format=self.data_format): |
| 38 | return resnet_backbone( |
| 39 | image, self.num_blocks, |
| 40 | preact_group if self.mode == 'preact' else resnet_group, self.block_func) |
| 41 | |
| 42 | |
| 43 | def get_config(model): |
no outgoing calls
no test coverage detected
searching dependent graphs…