BERT-base: hidden_size=768, num_layers=12, num_heads=12, intermediate=3072 (110M params).
| 67 | |
| 68 | |
| 69 | class BertModel(nn.Module): |
| 70 | """ |
| 71 | BERT-base: hidden_size=768, num_layers=12, num_heads=12, intermediate=3072 (110M params). |
| 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 |
| 103 | positions = torch.arange(T, device=input_ids.device).unsqueeze(0) |
| 104 | token_types = torch.zeros_like(input_ids) |
| 105 | |
| 106 | x = self.word_embeddings(input_ids) + self.position_embeddings(positions) + self.token_type_embeddings(token_types) |
| 107 | x = self.embed_dropout(self.embed_norm(x)) |
| 108 | |
| 109 | for layer in self.layers: |
| 110 | x = layer(x) |
| 111 | |
| 112 | # Pooled output from [CLS] token |
| 113 | pooled = torch.tanh(self.pooler(x[:, 0])) |
| 114 | return pooled |
nothing calls this directly
no outgoing calls
no test coverage detected