MCPcopy Create free account
hub / github.com/THUDM/GLM / test_boradcast_data

Function test_boradcast_data

mpu/tests/test_data.py:29–78  ·  view source on GitHub ↗
(model_parallel_size)

Source from the content-addressed store, hash-verified

27
28
29def test_boradcast_data(model_parallel_size):
30
31 if torch.distributed.get_rank() == 0:
32 print('> testing boradcast_data with model parallel size {} ...'.
33 format(model_parallel_size))
34
35 mpu.initialize_model_parallel(model_parallel_size)
36 torch.manual_seed(1234 + mpu.get_data_parallel_rank())
37 model_parallel_size = mpu.get_model_parallel_world_size()
38
39 key_size_t = {'key1': [7, 11],
40 'key2': [8, 2, 1],
41 'key3': [13],
42 'key4': [5, 1, 2],
43 'key5': [5, 12]}
44 keys = list(key_size_t.keys())
45
46 data = {}
47 data_t = {}
48 for key in key_size_t:
49 data[key] = torch.LongTensor(size=key_size_t[key]).random_(0, 1000)
50 data_t[key] = data[key].clone()
51 data['keyX'] = torch.FloatTensor(size=(5, )).random_(0, 1000)
52 data_t['keyX'] = data['keyX'].clone()
53 if mpu.get_model_parallel_rank() != 0:
54 data = None
55
56 data_utils._check_data_types(keys, data_t, torch.int64)
57 key_size, key_numel, \
58 total_numel = data_utils._build_key_size_numel_dictionaries(keys, data)
59 for key in keys:
60 assert key_size[key] == key_size_t[key]
61 total_numel_t = 0
62 for key in keys:
63 target_size = functools.reduce(operator.mul, key_size_t[key], 1)
64 assert key_numel[key] == target_size
65 total_numel_t += target_size
66 assert total_numel == total_numel_t
67
68 data_b = data_utils.broadcast_data(keys, data, torch.int64)
69 for key in keys:
70 tensor = data_t[key].cuda()
71 assert data_b[key].sub(tensor).abs().max() == 0
72
73 # Reset groups
74 mpu.destroy_model_parallel()
75
76 torch.distributed.barrier()
77 if torch.distributed.get_rank() == 0:
78 print('>> passed the test :-)')
79
80
81if __name__ == '__main__':

Callers 1

test_data.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected