MCPcopy Create free account
hub / github.com/XL2248/MSCTD / data_parallelism

Function data_parallelism

src_code/thumt1_code/thumt/utils/parallel.py:23–54  ·  view source on GitHub ↗
(devices, fn, *args, **kwargs)

Source from the content-addressed store, hash-verified

21
22# Data-level parallelism
23def data_parallelism(devices, fn, *args, **kwargs):
24 num_worker = len(devices)
25 devices = ["gpu:%d" % d for d in devices]
26
27 # Replicate args and kwargs
28 if args:
29 new_args = [_maybe_repeat(arg, num_worker) for arg in args]
30 # Transpose
31 new_args = [list(x) for x in zip(*new_args)]
32 else:
33 new_args = [[] for _ in range(num_worker)]
34
35 new_kwargs = [{} for _ in range(num_worker)]
36
37 for k, v in six.iteritems(kwargs):
38 vals = _maybe_repeat(v, num_worker)
39
40 for i in range(num_worker):
41 new_kwargs[i][k] = vals[i]
42
43 fns = _maybe_repeat(fn, num_worker)
44
45 # Now make the parallel call.
46 outputs = []
47
48 for i in range(num_worker):
49 with tf.variable_scope(tf.get_variable_scope(), reuse=(i != 0)):
50 with tf.name_scope("parallel_%d" % i):
51 with tf.device(devices[i]):
52 outputs.append(fns[i](*new_args[i], **new_kwargs[i]))
53
54 return outputs
55
56
57def shard_features(features, device_list):

Callers 1

parallel_modelFunction · 0.70

Calls 1

_maybe_repeatFunction · 0.70

Tested by

no test coverage detected