(model_parallel_size)
| 446 | |
| 447 | |
| 448 | def test_parallel_transformer_layer(model_parallel_size): |
| 449 | |
| 450 | if torch.distributed.get_rank() == 0: |
| 451 | print('> testing ParallelTransformerLayer with model parallel ' |
| 452 | 'size: {}'.format(model_parallel_size)) |
| 453 | |
| 454 | num_att_heads_per_partition = 3 |
| 455 | hidden_size_per_att_head = 7 |
| 456 | batch_size = 5 |
| 457 | sequence_length = 13 |
| 458 | |
| 459 | rank_1, hidden_size_1, model_parallel_size_1, loss_1, \ |
| 460 | transformer_layer_1, identity_layer_1 = parallel_transformer( |
| 461 | 1, num_att_heads_per_partition, |
| 462 | hidden_size_per_att_head, batch_size, sequence_length) |
| 463 | |
| 464 | rank, hidden_size, model_parallel_size, loss, \ |
| 465 | transformer_layer, identity_layer = parallel_transformer( |
| 466 | model_parallel_size, num_att_heads_per_partition, |
| 467 | hidden_size_per_att_head, batch_size, sequence_length) |
| 468 | |
| 469 | error = loss_1.sub(loss).abs().max() |
| 470 | torch.distributed.barrier() |
| 471 | print(' loss error on global rank {}: {}'.format( |
| 472 | torch.distributed.get_rank(), error)) |
| 473 | assert error < 5.0e-5, 'error: {}'.format(error) |
| 474 | |
| 475 | error = identity_layer_1.weight.grad.sub( |
| 476 | identity_layer.weight.grad).abs().max() |
| 477 | torch.distributed.barrier() |
| 478 | print(' input gradient error on global rank {}: {}'.format( |
| 479 | torch.distributed.get_rank(), error)) |
| 480 | assert error < 5.0e-5, 'error: {}'.format(error) |
| 481 | |
| 482 | torch.distributed.barrier() |
| 483 | if torch.distributed.get_rank() == 0: |
| 484 | print(' >> passed the test :-)') |
| 485 | |
| 486 | |
| 487 | if __name__ == '__main__': |
no test coverage detected