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

Function test_parallel_transformer_layer

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

Source from the content-addressed store, hash-verified

446
447
448def 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
487if __name__ == '__main__':

Callers 1

test_layers.pyFile · 0.85

Calls 1

parallel_transformerFunction · 0.85

Tested by

no test coverage detected