(devices, fn, *args, **kwargs)
| 21 | |
| 22 | # Data-level parallelism |
| 23 | def 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 | |
| 57 | def shard_features(features, device_list): |
no test coverage detected