Builds configs for BERT model. Defaults are from the BERT-BASE model. Args: vocab_size: Vocab size. hidden_dim: Hidden dim. dropout_rate: Dropout rate. dtype: Model dtype. embedding_cfg: Optional embedding config. Defaults to a BERT-base Embedding.
(
*,
vocab_size: int,
hidden_dim: int = 768,
dropout_rate: float = 0.0,
dtype: jnp.dtype = jnp.float32,
embedding_cfg: Optional[Embedding.Config] = None,
stack_cfg: Optional[BaseStackedTransformerLayer.Config] = None,
head_cfg: Optional[BaseClassificationHead.Config] = None,
encoder_cfg: Optional[Encoder.Config] = None,
base_cfg: Optional[BertModel.Config] = None,
)
| 461 | |
| 462 | |
| 463 | def bert_model_config( |
| 464 | *, |
| 465 | vocab_size: int, |
| 466 | hidden_dim: int = 768, |
| 467 | dropout_rate: float = 0.0, |
| 468 | dtype: jnp.dtype = jnp.float32, |
| 469 | embedding_cfg: Optional[Embedding.Config] = None, |
| 470 | stack_cfg: Optional[BaseStackedTransformerLayer.Config] = None, |
| 471 | head_cfg: Optional[BaseClassificationHead.Config] = None, |
| 472 | encoder_cfg: Optional[Encoder.Config] = None, |
| 473 | base_cfg: Optional[BertModel.Config] = None, |
| 474 | ) -> BertModel.Config: |
| 475 | """Builds configs for BERT model. |
| 476 | |
| 477 | Defaults are from the BERT-BASE model. |
| 478 | |
| 479 | Args: |
| 480 | vocab_size: Vocab size. |
| 481 | hidden_dim: Hidden dim. |
| 482 | dropout_rate: Dropout rate. |
| 483 | dtype: Model dtype. |
| 484 | embedding_cfg: Optional embedding config. Defaults to a BERT-base Embedding. |
| 485 | stack_cfg: Optional transformer stack config. Defaults to a StackedTransformerLayer. |
| 486 | head_cfg: Optional head config. Defaults to a BertLMHead.Config. |
| 487 | encoder_cfg: Optional encoder config. Defaults to a Encoder.Config. |
| 488 | base_cfg: Optional base config. Will be cloned. |
| 489 | |
| 490 | Returns: |
| 491 | The stack config. |
| 492 | """ |
| 493 | base_cfg = base_cfg.clone() if base_cfg else BertModel.default_config() |
| 494 | |
| 495 | encoder_cfg = encoder_cfg or Encoder.default_config().set( |
| 496 | dim=hidden_dim, |
| 497 | vocab_size=vocab_size, |
| 498 | dropout_rate=dropout_rate, |
| 499 | emb=embedding_cfg or bert_embedding_config(), |
| 500 | transformer=stack_cfg or bert_transformer_config(), |
| 501 | pad_token_id=0, |
| 502 | ) |
| 503 | head_cfg = head_cfg or bert_lm_head_config( |
| 504 | base_cfg=base_cfg.head, # pylint: disable=no-member |
| 505 | vocab_size=vocab_size, |
| 506 | ) |
| 507 | base_cfg.set( |
| 508 | dtype=dtype, vocab_size=vocab_size, dim=hidden_dim, encoder=encoder_cfg, head=head_cfg |
| 509 | ) |
| 510 | set_layer_norm_eps_recursively(base_cfg, bert_layer_norm_epsilon(dtype=dtype)) |
| 511 | return base_cfg |
| 512 | |
| 513 | |
| 514 | def bert_layer_norm_epsilon(dtype=None): |