MCPcopy Create free account
hub / github.com/FreeformRobotics/OTS / SegmentationModule

Class SegmentationModule

models/models.py:21–47  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

19
20
21class SegmentationModule(SegmentationModuleBase):
22 def __init__(self, net_enc, net_dec, crit, deep_sup_scale=None):
23 super(SegmentationModule, self).__init__()
24 self.encoder = net_enc
25 self.decoder = net_dec
26 self.crit = crit
27 self.deep_sup_scale = deep_sup_scale
28
29 def forward(self, feed_dict, *, segSize=None):
30 # training
31 if segSize is None:
32 if self.deep_sup_scale is not None: # use deep supervision technique
33 (pred, pred_deepsup) = self.decoder(self.encoder(feed_dict['img_data'], return_feature_maps=True))
34 else:
35 pred = self.decoder(self.encoder(feed_dict['img_data'], return_feature_maps=True))
36
37 loss = self.crit(pred, feed_dict['seg_label'])
38 if self.deep_sup_scale is not None:
39 loss_deepsup = self.crit(pred_deepsup, feed_dict['seg_label'])
40 loss = loss + loss_deepsup * self.deep_sup_scale
41
42 acc = self.pixel_acc(pred, feed_dict['seg_label'])
43 return loss, acc
44 # inference
45 else:
46 pred = self.decoder(self.encoder(feed_dict['img_data'], return_feature_maps=True), segSize=segSize)
47 return pred
48
49
50class ModelBuilder:

Callers 1

load_seg_moduleFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected