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

Function parallel_transformer

mpu/tests/test_layers.py:412–445  ·  view source on GitHub ↗
(model_parallel_size, num_att_heads_per_partition,
                         hidden_size_per_att_head, batch_size, sequence_length)

Source from the content-addressed store, hash-verified

410 print(' >> passed the test :-)')
411
412def 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
448def test_parallel_transformer_layer(model_parallel_size):

Callers 1

Calls 3

set_random_seedFunction · 0.90
IdentityLayer3DClass · 0.85
backwardMethod · 0.45

Tested by

no test coverage detected