| 36 | class BertEmbedding(Module): |
| 37 | |
| 38 | def __init__(self, |
| 39 | vocab_size, |
| 40 | hidden_size, |
| 41 | max_position_embeddings, |
| 42 | type_vocab_size, |
| 43 | dtype=None): |
| 44 | super().__init__() |
| 45 | self.vocab_embedding = Embedding(vocab_size, hidden_size, dtype=dtype) |
| 46 | self.position_embedding = Embedding(max_position_embeddings, |
| 47 | hidden_size, |
| 48 | dtype=dtype) |
| 49 | self.token_embedding = Embedding(type_vocab_size, |
| 50 | hidden_size, |
| 51 | dtype=dtype) |
| 52 | self.max_position_embeddings = max_position_embeddings |
| 53 | |
| 54 | self.embedding_ln = LayerNorm(normalized_shape=hidden_size, dtype=dtype) |
| 55 | |
| 56 | def forward(self, input_ids, position_ids, token_type_ids): |
| 57 | x = self.vocab_embedding(input_ids) |