| 27 | |
| 28 | |
| 29 | def 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 | |
| 81 | if __name__ == '__main__': |