(num_classes=1000, pretrained=False, distillation=False, fuse=False, pretrained_cfg=None, model_cfg=EfficientViT_m5)
| 155 | |
| 156 | @register_model |
| 157 | def EfficientViT_M5(num_classes=1000, pretrained=False, distillation=False, fuse=False, pretrained_cfg=None, model_cfg=EfficientViT_m5): |
| 158 | model = EfficientViT(num_classes=num_classes, distillation=distillation, **model_cfg) |
| 159 | if pretrained: |
| 160 | pretrained = _checkpoint_url_format.format(pretrained) |
| 161 | checkpoint = torch.hub.load_state_dict_from_url( |
| 162 | pretrained, map_location='cpu') |
| 163 | d = checkpoint['model'] |
| 164 | D = model.state_dict() |
| 165 | for k in d.keys(): |
| 166 | if D[k].shape != d[k].shape: |
| 167 | d[k] = d[k][:, :, None, None] |
| 168 | model.load_state_dict(d) |
| 169 | if fuse: |
| 170 | replace_batchnorm(model) |
| 171 | return model |
| 172 | |
| 173 | def replace_batchnorm(net): |
| 174 | for child_name, child in net.named_children(): |
nothing calls this directly
no test coverage detected