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

Function shard_features

src_code/thumt1_code/thumt/utils/parallel.py:57–89  ·  view source on GitHub ↗
(features, device_list)

Source from the content-addressed store, hash-verified

55
56
57def shard_features(features, device_list):
58 num_datashards = len(device_list)
59 sharded_features = {}
60
61 with tf.device("/cpu:0"):
62 for k, v in six.iteritems(features):
63 v = tf.convert_to_tensor(v)
64
65 if not v.shape.as_list():
66 v = tf.expand_dims(v, axis=-1)
67 v = tf.tile(v, [num_datashards])
68
69 batch_size = tf.shape(v)[0]
70 size_splits = []
71
72 for i in range(num_datashards):
73 size_splits.append(
74 tf.cond(tf.greater(tf.mod(batch_size, num_datashards), i),
75 lambda: batch_size // num_datashards + 1,
76 lambda: batch_size // num_datashards)
77 )
78
79 sharded_features[k] = tf.split(v, size_splits, 0)
80
81 datashard_to_features = []
82
83 for d in range(num_datashards):
84 feat = {
85 k: v[d] for k, v in six.iteritems(sharded_features)
86 }
87 datashard_to_features.append(feat)
88
89 return datashard_to_features
90
91
92def parallel_model(model_fn, features, devices):

Callers 1

parallel_modelFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected