(checkpoint, obj_model, classifier)
| 132 | |
| 133 | |
| 134 | def load_checkpoint(checkpoint, obj_model, classifier): |
| 135 | checkpoint = torch.load(checkpoint) |
| 136 | if classifier: |
| 137 | classifier.load_state_dict(checkpoint['classifier_state_dict']) |
| 138 | if obj_model: |
| 139 | obj_model.load_state_dict(checkpoint['obj_state_dict']) |
| 140 | return obj_model, classifier |
| 141 | |
| 142 | |
| 143 | def img2onehot(img_name, one_hot): |