(model_parallel_size)
| 29 | |
| 30 | |
| 31 | def 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( |
no test coverage detected