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

Function test_parallel_embedding

mpu/tests/test_layers.py:31–106  ·  view source on GitHub ↗
(model_parallel_size)

Source from the content-addressed store, hash-verified

29
30
31def test_parallel_embedding(model_parallel_size):
32
33 if torch.distributed.get_rank() == 0:
34 print('> testing parallel embedding with model parallel size {} ...'.
35 format(model_parallel_size))
36
37 mpu.initialize_model_parallel(model_parallel_size)
38 model_parallel_size = mpu.get_model_parallel_world_size()
39
40 batch_size = 17
41 seq_length = 23
42 vocab_size = 48
43 hidden_size = 16
44 seed = 1236
45
46 set_random_seed(123)
47 input_data = torch.LongTensor(
48 size=(batch_size,seq_length)).random_(0, vocab_size).cuda()
49 loss_weight = torch.randn([batch_size, seq_length, hidden_size]).cuda()
50
51 set_random_seed(seed)
52 embedding_original = torch.nn.Embedding(vocab_size, hidden_size).cuda()
53
54 output = embedding_original(input_data)
55 loss_original = torch.mul(output, loss_weight).sum()
56 loss_original.backward()
57
58 set_random_seed(seed)
59 embedding_parallel = layers.ParallelEmbedding(
60 vocab_size, hidden_size, init_method=init.normal_).cuda()
61 output = embedding_parallel(input_data)
62 loss_parallel = torch.mul(output, loss_weight).sum()
63 loss_parallel.backward()
64
65 set_random_seed(seed)
66 embedding_vocab_parallel = layers.VocabParallelEmbedding(
67 vocab_size, hidden_size, init_method=init.normal_).cuda()
68 output = embedding_vocab_parallel(input_data)
69 loss_vocab_parallel = torch.mul(output, loss_weight).sum()
70 loss_vocab_parallel.backward()
71
72 torch.distributed.barrier()
73 error = loss_parallel.sub(loss_original).abs()
74 print(' error in loss (parallel) on global rank {}: {}'.format(
75 torch.distributed.get_rank(), error))
76 assert error < 1.0e-12, 'error: {}'.format(error)
77
78 torch.distributed.barrier()
79 error = loss_vocab_parallel.sub(loss_original).abs()
80 print(' error in loss (vocab parallel) on global rank {}: {}'.format(
81 torch.distributed.get_rank(), error))
82 assert error < 1.0e-12, 'error: {}'.format(error)
83
84 weight_grad_orig = torch.split(embedding_original.weight.grad,
85 hidden_size // model_parallel_size,
86 1)[mpu.get_model_parallel_rank()]
87 error = embedding_parallel.weight.grad.sub(weight_grad_orig).abs().max()
88 print(' error in grad (parallel) on global rank {}: {}'.format(

Callers 1

test_layers.pyFile · 0.85

Calls 2

set_random_seedFunction · 0.90
backwardMethod · 0.45

Tested by

no test coverage detected