(model_parallel_size, num_att_heads_per_partition,
hidden_size_per_att_head, batch_size, sequence_length)
| 410 | print(' >> passed the test :-)') |
| 411 | |
| 412 | def parallel_transformer(model_parallel_size, num_att_heads_per_partition, |
| 413 | hidden_size_per_att_head, batch_size, sequence_length): |
| 414 | |
| 415 | mpu.initialize_model_parallel(model_parallel_size) |
| 416 | model_parallel_size = mpu.get_model_parallel_world_size() |
| 417 | |
| 418 | seed = 12345 |
| 419 | set_random_seed(seed) |
| 420 | |
| 421 | num_att_heads = num_att_heads_per_partition * \ |
| 422 | torch.distributed.get_world_size() |
| 423 | hidden_size = hidden_size_per_att_head * num_att_heads |
| 424 | intermediate_size = 4 * hidden_size |
| 425 | |
| 426 | # Network |
| 427 | identity_layer = IdentityLayer3D(batch_size, sequence_length, |
| 428 | hidden_size).cuda() |
| 429 | transformer_layer = mpu.BertParallelTransformerLayer( |
| 430 | hidden_size, intermediate_size, num_att_heads, 0.0, 0.0, |
| 431 | torch.nn.functional.relu, 1.0e-5).cuda() |
| 432 | |
| 433 | loss_weight = torch.randn([batch_size, sequence_length, hidden_size]).cuda() |
| 434 | attention_mask = torch.randn([batch_size, 1, 1, sequence_length]).cuda() |
| 435 | # Forward |
| 436 | input_ = identity_layer() |
| 437 | output = transformer_layer(input_, attention_mask) |
| 438 | loss = torch.mul(output, loss_weight).sum() |
| 439 | # Backward |
| 440 | loss.backward() |
| 441 | |
| 442 | rank = mpu.get_model_parallel_rank() |
| 443 | mpu.destroy_model_parallel() |
| 444 | return rank, hidden_size, model_parallel_size, loss, \ |
| 445 | transformer_layer, identity_layer |
| 446 | |
| 447 | |
| 448 | def test_parallel_transformer_layer(model_parallel_size): |
no test coverage detected