MCPcopy Create free account
hub / github.com/apple/axlearn / bert_model_config

Function bert_model_config

axlearn/common/bert.py:463–511  ·  view source on GitHub ↗

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,
)

Source from the content-addressed store, hash-verified

461
462
463def 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
514def bert_layer_norm_epsilon(dtype=None):

Calls 8

bert_embedding_configFunction · 0.85
bert_transformer_configFunction · 0.85
bert_lm_head_configFunction · 0.85
bert_layer_norm_epsilonFunction · 0.85
cloneMethod · 0.80
default_configMethod · 0.45
setMethod · 0.45