MCPcopy Create free account
hub / github.com/binbinjiang/CVT-SLR / GpuDataParallel

Class GpuDataParallel

utils/device.py:6–56  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

4
5
6class GpuDataParallel(object):
7 def __init__(self):
8 self.gpu_list = []
9 self.output_device = None
10
11 def set_device(self, device):
12 device = str(device)
13 if device != 'None':
14 self.gpu_list = [i for i in range(len(device.split(',')))]
15 os.environ["CUDA_VISIBLE_DEVICES"] = device
16 output_device = self.gpu_list[0]
17 self.occupy_gpu(self.gpu_list)
18 self.output_device = output_device if len(self.gpu_list) > 0 else "cpu"
19
20 def model_to_device(self, model):
21 # model = convert_model(model)
22 model = model.to(self.output_device)
23 if len(self.gpu_list) > 1:
24 model = nn.DataParallel(
25 model,
26 device_ids=self.gpu_list,
27 output_device=self.output_device)
28 return model
29
30 def data_to_device(self, data):
31 if isinstance(data, torch.FloatTensor):
32 return data.to(self.output_device)
33 elif isinstance(data, torch.DoubleTensor):
34 return data.float().to(self.output_device)
35 elif isinstance(data, torch.ByteTensor):
36 return data.long().to(self.output_device)
37 elif isinstance(data, torch.LongTensor):
38 return data.to(self.output_device)
39 elif isinstance(data, list) or isinstance(data, tuple):
40 return [self.data_to_device(d) for d in data]
41 else:
42 raise ValueError(data.shape, "Unknown Dtype: {}".format(data.dtype))
43
44 def criterion_to_device(self, loss):
45 return loss.to(self.output_device)
46
47 def occupy_gpu(self, gpus=None):
48 """
49 make program appear on nvidia-smi.
50 """
51 if len(gpus) == 0:
52 torch.zeros(1).cuda()
53 else:
54 gpus = [gpus] if isinstance(gpus, int) else list(gpus)
55 for g in gpus:
56 torch.zeros(1).cuda(g)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected