MCPcopy Create free account
hub / github.com/RightNow-AI/autokernel / BertModel

Class BertModel

models/bert_base.py:69–114  ·  view source on GitHub ↗

BERT-base: hidden_size=768, num_layers=12, num_heads=12, intermediate=3072 (110M params).

Source from the content-addressed store, hash-verified

67
68
69class 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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected