| 24 | return modelG, modelD, flowNet |
| 25 | |
| 26 | class myModel(nn.Module): |
| 27 | def __init__(self, opt, model): |
| 28 | super(myModel, self).__init__() |
| 29 | self.opt = opt |
| 30 | self.module = model |
| 31 | self.model = nn.DataParallel(model, device_ids=opt.gpu_ids) |
| 32 | self.bs_per_gpu = int(np.ceil(float(opt.batchSize) / len(opt.gpu_ids))) # batch size for each GPU |
| 33 | self.pad_bs = self.bs_per_gpu * len(opt.gpu_ids) - opt.batchSize |
| 34 | |
| 35 | def forward(self, *inputs, **kwargs): |
| 36 | inputs = self.add_dummy_to_tensor(inputs, self.pad_bs) |
| 37 | outputs = self.model(*inputs, **kwargs, dummy_bs=self.pad_bs) |
| 38 | if self.pad_bs == self.bs_per_gpu: # gpu 0 does 0 batch but still returns 1 batch |
| 39 | return self.remove_dummy_from_tensor(outputs, 1) |
| 40 | return outputs |
| 41 | |
| 42 | def add_dummy_to_tensor(self, tensors, add_size=0): |
| 43 | if add_size == 0 or tensors is None: return tensors |
| 44 | if type(tensors) == list or type(tensors) == tuple: |
| 45 | return [self.add_dummy_to_tensor(tensor, add_size) for tensor in tensors] |
| 46 | |
| 47 | if isinstance(tensors, torch.Tensor): |
| 48 | dummy = torch.zeros_like(tensors)[:add_size] |
| 49 | tensors = torch.cat([dummy, tensors]) |
| 50 | return tensors |
| 51 | |
| 52 | def remove_dummy_from_tensor(self, tensors, remove_size=0): |
| 53 | if remove_size == 0 or tensors is None: return tensors |
| 54 | if type(tensors) == list or type(tensors) == tuple: |
| 55 | return [self.remove_dummy_from_tensor(tensor, remove_size) for tensor in tensors] |
| 56 | |
| 57 | if isinstance(tensors, torch.Tensor): |
| 58 | tensors = tensors[remove_size:] |
| 59 | return tensors |
| 60 | |
| 61 | def create_model(opt): |
| 62 | print(opt.model) |