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

Function test_initialize_model_parallel

mpu/tests/test_initialize.py:26–62  ·  view source on GitHub ↗
(model_parallel_size)

Source from the content-addressed store, hash-verified

24
25
26def 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
65def test_get_model_parallel_src_rank(model_parallel_size_):

Callers 1

test_initialize.pyFile · 0.85

Calls 1

checkFunction · 0.85

Tested by

no test coverage detected