| 39 | r"""Embedding layer of the PanGUAlpha Model""" |
| 40 | |
| 41 | def __init__(self, config): |
| 42 | super(EmbeddingLayer, self).__init__() |
| 43 | # Only for the pipeline mode, the embedding needs to be row sliced. |
| 44 | dp = config.parallel_config.embedding_dp_mp_config.data_parallel |
| 45 | mp = config.parallel_config.embedding_dp_mp_config.model_parallel |
| 46 | self.word_embedding = VocabEmbedding( |
| 47 | vocab_size=(config.vocab_size // 1024 + 1) * 1024, |
| 48 | embedding_size=config.hidden_size, |
| 49 | param_init=initializer( |
| 50 | "normal", |
| 51 | [(config.vocab_size // 1024 + 1) * 1024, config.hidden_size], |
| 52 | # dtype=config.param_init_type, |
| 53 | dtype=mstype.float32, |
| 54 | ), |
| 55 | parallel_config=config.parallel_config.embedding_dp_mp_config, |
| 56 | ) |
| 57 | self.word_embedding.gather.shard(((mp, 1), (dp, 1))) |
| 58 | # self.word_embedding.embedding_table.parallel_optimizer = True |
| 59 | copied_parallel_config = copy.deepcopy(config.parallel_config) |
| 60 | copied_parallel_config.vocab_emb_dp = True |
| 61 | self.position_embedding = VocabEmbedding( |
| 62 | vocab_size=config.seq_length, |
| 63 | embedding_size=config.hidden_size, |
| 64 | param_init=initializer( |
| 65 | "normal", |
| 66 | [config.seq_length, config.hidden_size], |
| 67 | # dtype=config.param_init_type, |
| 68 | dtype=mstype.float32, |
| 69 | ), |
| 70 | parallel_config=copied_parallel_config.embedding_dp_mp_config, |
| 71 | ) |
| 72 | self.split = P.Split(1, 2).shard( |
| 73 | ((config.parallel_config.data_parallel, 1, 1),) |
| 74 | ) |
| 75 | self.print = P.Print() |
| 76 | self.add = P.Add().shard( |
| 77 | ( |
| 78 | (config.parallel_config.data_parallel, 1, 1), |
| 79 | (config.parallel_config.data_parallel, 1, 1), |
| 80 | ) |
| 81 | ) |
| 82 | # self.dropout = nn.Dropout(1 - config.dropout_rate) |
| 83 | self.dropout = _Dropout(1 - config.dropout_rate) |
| 84 | self.dropout.shard(((config.parallel_config.data_parallel, 1, 1),)) |
| 85 | # self.dropout.dropout.shard( |
| 86 | # ((config.parallel_config.data_parallel, 1, 1),) |
| 87 | # ) |
| 88 | self.is_first_iteration = True |
| 89 | self.use_past = config.use_past |
| 90 | self.batch_size = config.batch_size |
| 91 | |
| 92 | def construct( |
| 93 | self, input_ids, input_position, init_reset, batch_valid_length |