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

Method _mlm_processor_config

axlearn/common/input_mlm_test.py:713–738  ·  view source on GitHub ↗
(
        self,
        *,
        actions_cfg: Optional[InstantiableConfig] = None,
        **kwargs,
    )

Source from the content-addressed store, hash-verified

711
712class 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(

Callers 2

test_fake_text_dataMethod · 0.95
test_filteringMethod · 0.95

Calls 4

config_for_functionFunction · 0.90
config_for_classFunction · 0.90
setMethod · 0.45
updateMethod · 0.45

Tested by

no test coverage detected