| 456 | |
| 457 | |
| 458 | class BertPooler(nn.Module): |
| 459 | def __init__(self, config): |
| 460 | super().__init__() |
| 461 | self.dense = nn.Linear(config.hidden_size, config.hidden_size) |
| 462 | self.activation = nn.Tanh() |
| 463 | |
| 464 | def forward(self, hidden_states): |
| 465 | # We "pool" the model by simply taking the hidden state corresponding |
| 466 | # to the first token. |
| 467 | first_token_tensor = hidden_states[:, 0] |
| 468 | pooled_output = self.dense(first_token_tensor) |
| 469 | pooled_output = self.activation(pooled_output) |
| 470 | return pooled_output |
| 471 | |
| 472 | |
| 473 | class BertPredictionHeadTransform(nn.Module): |