(
self,
*,
actions_cfg: Optional[InstantiableConfig] = None,
**kwargs,
)
| 711 | |
| 712 | class MaskedLmInputTest(parameterized.TestCase, tf.test.TestCase): |
| 713 | def _mlm_processor_config( |
| 714 | self, |
| 715 | *, |
| 716 | actions_cfg: Optional[InstantiableConfig] = None, |
| 717 | **kwargs, |
| 718 | ): |
| 719 | def noop_normalizer() -> input_tf_data.DatasetToDatasetFn: |
| 720 | return lambda ds: ds |
| 721 | |
| 722 | mlm_mask_cfg = config_for_function(apply_mlm_mask).set( |
| 723 | whole_word_mask=True, |
| 724 | ) |
| 725 | if actions_cfg is not None: |
| 726 | mlm_mask_cfg.set(actions_cfg=actions_cfg) |
| 727 | |
| 728 | defaults = dict( |
| 729 | sentence_piece_vocab=config_for_class(seqio.SentencePieceVocabulary).set( |
| 730 | sentencepiece_model_file=t5_sentence_piece_vocab_file, |
| 731 | extra_ids=1, |
| 732 | ), |
| 733 | normalization=config_for_function(noop_normalizer), |
| 734 | apply_mlm_mask=mlm_mask_cfg, |
| 735 | mask_token="▁<extra_id_0>", |
| 736 | ) |
| 737 | defaults.update(kwargs) |
| 738 | return config_for_function(text_to_mlm_input).set(**defaults) |
| 739 | |
| 740 | @parameterized.product(is_training=[True, False], truncate=[True, False]) |
| 741 | @pytest.mark.skipif( |
no test coverage detected