(model_parallel_size)
| 24 | |
| 25 | |
| 26 | def test_initialize_model_parallel(model_parallel_size): |
| 27 | |
| 28 | if torch.distributed.get_rank() == 0: |
| 29 | print('> testing initialize_model_parallel with size {} ...'.format( |
| 30 | model_parallel_size)) |
| 31 | model_parallel_size_ = min(model_parallel_size, |
| 32 | torch.distributed.get_world_size()) |
| 33 | assert not mpu.model_parallel_is_initialized() |
| 34 | mpu.initialize_model_parallel(model_parallel_size_) |
| 35 | assert mpu.model_parallel_is_initialized() |
| 36 | |
| 37 | # Checks. |
| 38 | def check(group, world_size, rank): |
| 39 | assert world_size == torch.distributed.get_world_size(group=group) |
| 40 | assert rank == torch.distributed.get_rank(group=group) |
| 41 | |
| 42 | # Model parallel. |
| 43 | world_size = model_parallel_size_ |
| 44 | rank = torch.distributed.get_rank() % model_parallel_size_ |
| 45 | assert world_size == mpu.get_model_parallel_world_size() |
| 46 | assert rank == mpu.get_model_parallel_rank() |
| 47 | check(mpu.get_model_parallel_group(), world_size, rank) |
| 48 | |
| 49 | |
| 50 | # Data parallel. |
| 51 | world_size = torch.distributed.get_world_size() // model_parallel_size_ |
| 52 | rank = torch.distributed.get_rank() // model_parallel_size |
| 53 | assert world_size == mpu.get_data_parallel_world_size() |
| 54 | assert rank == mpu.get_data_parallel_rank() |
| 55 | check(mpu.get_data_parallel_group(), world_size, rank) |
| 56 | |
| 57 | # Reset groups |
| 58 | mpu.destroy_model_parallel() |
| 59 | |
| 60 | torch.distributed.barrier() |
| 61 | if torch.distributed.get_rank() == 0: |
| 62 | print('>> passed the test :-)') |
| 63 | |
| 64 | |
| 65 | def test_get_model_parallel_src_rank(model_parallel_size_): |
no test coverage detected