(modeldir, segmodel_arch, segvocab, epoch=None)
| 507 | return segmodel |
| 508 | |
| 509 | def load_segmentation_model(modeldir, segmodel_arch, segvocab, epoch=None): |
| 510 | # Load csv of class names |
| 511 | segmodel_dir = 'dataset/segmodel/%s-%s-%s' % ((segvocab,) + segmodel_arch) |
| 512 | with open(os.path.join(segmodel_dir, 'labels.json')) as f: |
| 513 | labeldata = EasyDict(json.load(f)) |
| 514 | # Automatically pick the last epoch available. |
| 515 | if epoch is None: |
| 516 | choices = [os.path.basename(n)[14:-4] for n in |
| 517 | glob.glob(os.path.join(segmodel_dir, 'encoder_epoch_*.pth'))] |
| 518 | epoch = max([int(c) for c in choices if c.isdigit()]) |
| 519 | # Create a segmentation model |
| 520 | segbuilder = segmodel_module.ModelBuilder() |
| 521 | # example segmodel_arch = ('resnet101', 'upernet') |
| 522 | seg_encoder = segbuilder.build_encoder( |
| 523 | arch=segmodel_arch[0], |
| 524 | fc_dim=2048, |
| 525 | weights=os.path.join(segmodel_dir, 'encoder_epoch_%d.pth' % epoch)) |
| 526 | seg_decoder = segbuilder.build_decoder( |
| 527 | arch=segmodel_arch[1], |
| 528 | fc_dim=2048, inference=True, num_class=len(labeldata.labels), |
| 529 | weights=os.path.join(segmodel_dir, 'decoder_epoch_%d.pth' % epoch)) |
| 530 | segmodel = segmodel_module.SegmentationModule(seg_encoder, seg_decoder, |
| 531 | torch.nn.NLLLoss(ignore_index=-1)) |
| 532 | segmodel.categories = [cat.name for cat in labeldata.categories] |
| 533 | segmodel.labels = [label.name for label in labeldata.labels] |
| 534 | categories = OrderedDict() |
| 535 | label_category = numpy.zeros(len(segmodel.labels), dtype=int) |
| 536 | for i, label in enumerate(labeldata.labels): |
| 537 | label_category[i] = segmodel.categories.index(label.category) |
| 538 | segmodel.meta = labeldata |
| 539 | segmodel.eval() |
| 540 | return segmodel |
| 541 | |
| 542 | def ensure_upp_segmenter_downloaded(directory): |
| 543 | baseurl = 'http://netdissect.csail.mit.edu/data/segmodel' |
no test coverage detected