MCPcopy Create free account
hub / github.com/NVIDIA/vid2vid / myModel

Class myModel

models/models.py:26–59  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

24 return modelG, modelD, flowNet
25
26class 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
61def create_model(opt):
62 print(opt.model)

Callers 1

wrap_modelFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected