| 72 | """ |
| 73 | |
| 74 | def __init__( |
| 75 | self, |
| 76 | vocab_size: int = 30522, |
| 77 | hidden_size: int = 768, |
| 78 | num_layers: int = 12, |
| 79 | num_heads: int = 12, |
| 80 | intermediate_size: int = 3072, |
| 81 | max_seq_len: int = 512, |
| 82 | dropout: float = 0.0, |
| 83 | ): |
| 84 | super().__init__() |
| 85 | self.word_embeddings = nn.Embedding(vocab_size, hidden_size) |
| 86 | self.position_embeddings = nn.Embedding(max_seq_len, hidden_size) |
| 87 | self.token_type_embeddings = nn.Embedding(2, hidden_size) |
| 88 | self.embed_norm = nn.LayerNorm(hidden_size) |
| 89 | self.embed_dropout = nn.Dropout(dropout) |
| 90 | |
| 91 | self.layers = nn.ModuleList([ |
| 92 | BertLayer(hidden_size, num_heads, intermediate_size, dropout) |
| 93 | for _ in range(num_layers) |
| 94 | ]) |
| 95 | |
| 96 | self.pooler = nn.Linear(hidden_size, hidden_size) |
| 97 | |
| 98 | n_params = sum(p.numel() for p in self.parameters()) |
| 99 | print(f"BertModel: {n_params / 1e6:.1f}M parameters") |
| 100 | |
| 101 | def forward(self, input_ids: torch.Tensor) -> torch.Tensor: |
| 102 | B, T = input_ids.shape |