| 106 | # |data|: dictionary of the input data |
| 107 | |
| 108 | def preprocess_input(self, data): |
| 109 | # move to GPU and change data types |
| 110 | data['label'] = data['label'].long() |
| 111 | if self.use_gpu(): |
| 112 | data['label'] = data['label'].cuda() |
| 113 | data['instance'] = data['instance'].cuda() |
| 114 | data['image'] = data['image'].cuda() |
| 115 | |
| 116 | # create one-hot label map |
| 117 | label_map = data['label'] |
| 118 | bs, _, h, w = label_map.size() |
| 119 | nc = self.opt.label_nc + 1 if self.opt.contain_dontcare_label \ |
| 120 | else self.opt.label_nc |
| 121 | input_label = self.FloatTensor(bs, nc, h, w).zero_() |
| 122 | input_semantics = input_label.scatter_(1, label_map, 1.0) |
| 123 | |
| 124 | # concatenate instance map if it exists |
| 125 | if not self.opt.no_instance: |
| 126 | inst_map = data['instance'] |
| 127 | instance_edge_map = self.get_edges(inst_map) |
| 128 | input_semantics = torch.cat((input_semantics, instance_edge_map), dim=1) |
| 129 | |
| 130 | return input_semantics, data['image'] |
| 131 | |
| 132 | def compute_generator_loss(self, input_semantics, real_image): |
| 133 | G_losses = {} |