Build language model and return along with the key to save.
(config, num_tokentypes, add_pooler,
encoder_attn_mask_type,
add_encoder=True,
add_decoder=False,
decoder_attn_mask_type=AttnMaskType.causal,
pre_process=True, post_process=True)
| 48 | |
| 49 | |
| 50 | def get_language_model(config, num_tokentypes, add_pooler, |
| 51 | encoder_attn_mask_type, |
| 52 | add_encoder=True, |
| 53 | add_decoder=False, |
| 54 | decoder_attn_mask_type=AttnMaskType.causal, |
| 55 | pre_process=True, post_process=True): |
| 56 | """Build language model and return along with the key to save.""" |
| 57 | args = get_args() |
| 58 | if config.init_method is None: |
| 59 | config.init_method = init_method_normal(config.init_method_std) |
| 60 | |
| 61 | if config.output_layer_init_method is None: |
| 62 | config.output_layer_init_method = scaled_init_method_normal(config.init_method_std, |
| 63 | config.num_layers) |
| 64 | |
| 65 | # Language model. |
| 66 | language_model = TransformerLanguageModel( |
| 67 | config, |
| 68 | encoder_attn_mask_type, |
| 69 | num_tokentypes=num_tokentypes, |
| 70 | add_encoder=add_encoder, |
| 71 | add_decoder=add_decoder, |
| 72 | decoder_attn_mask_type=decoder_attn_mask_type, |
| 73 | add_pooler=add_pooler, |
| 74 | pre_process=pre_process, |
| 75 | post_process=post_process |
| 76 | ) |
| 77 | # key used for checkpoints. |
| 78 | language_model_key = 'language_model' |
| 79 | |
| 80 | return language_model, language_model_key |
| 81 | |
| 82 | |
| 83 | class Pooler(MegatronModule): |
no test coverage detected