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

Function test_get_model_parallel_src_rank

mpu/tests/test_initialize.py:65–85  ·  view source on GitHub ↗
(model_parallel_size_)

Source from the content-addressed store, hash-verified

63
64
65def 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
88if __name__ == '__main__':

Callers 1

test_initialize.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected