(
*, vocab_size: int, source_len: int, target_len: int, remat_spec: Optional[RematSpec] = None
)
| 50 | |
| 51 | |
| 52 | def _model_config( |
| 53 | *, vocab_size: int, source_len: int, target_len: int, remat_spec: Optional[RematSpec] = None |
| 54 | ) -> EncoderDecoderModel.Config: |
| 55 | hidden_dim = 12 |
| 56 | num_heads = 4 |
| 57 | encoder = Encoder.default_config().set( |
| 58 | dim=hidden_dim, |
| 59 | vocab_size=vocab_size, |
| 60 | emb=bert_embedding_config(type_vocab_size=1, max_position_embeddings=source_len), |
| 61 | transformer=bert_transformer_config(num_layers=2, num_heads=num_heads), |
| 62 | pad_token_id=0, |
| 63 | ) |
| 64 | set_layer_norm_eps_recursively(encoder, 1e-8) |
| 65 | |
| 66 | decoder = gpt_decoder_config( |
| 67 | stack_cfg=StackedTransformerLayer.default_config(), |
| 68 | num_layers=2, |
| 69 | hidden_dim=hidden_dim, |
| 70 | num_heads=num_heads, |
| 71 | vocab_size=vocab_size, |
| 72 | activation_function="nn.relu", |
| 73 | max_position_embeddings=target_len, |
| 74 | layer_remat=remat_spec, |
| 75 | ) |
| 76 | set_decoder_cross_attention_config(decoder_cfg=decoder, num_heads=num_heads) |
| 77 | return EncoderDecoderModel.default_config().set(decoder=decoder, encoder=encoder) |
| 78 | |
| 79 | |
| 80 | class TestEncoderDecoder(TestCase): |
no test coverage detected