| 63 | |
| 64 | |
| 65 | def test_get_model_parallel_src_rank(model_parallel_size_): |
| 66 | |
| 67 | if torch.distributed.get_rank() == 0: |
| 68 | print('> testing get_model_parallel_src_rank with size {} ...'.format( |
| 69 | model_parallel_size_)) |
| 70 | model_parallel_size = min(model_parallel_size_, |
| 71 | torch.distributed.get_world_size()) |
| 72 | assert not mpu.model_parallel_is_initialized() |
| 73 | mpu.initialize_model_parallel(model_parallel_size) |
| 74 | assert mpu.model_parallel_is_initialized() |
| 75 | |
| 76 | # Checks |
| 77 | src_rank = torch.distributed.get_rank() - mpu.get_model_parallel_rank() |
| 78 | assert mpu.get_model_parallel_src_rank() == src_rank |
| 79 | |
| 80 | # Reset groups |
| 81 | mpu.destroy_model_parallel() |
| 82 | |
| 83 | torch.distributed.barrier() |
| 84 | if torch.distributed.get_rank() == 0: |
| 85 | print('>> passed the test :-)') |
| 86 | |
| 87 | |
| 88 | if __name__ == '__main__': |