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

Method test_text_input

axlearn/audio/input_asr_test.py:218–249  ·  view source on GitHub ↗
(
        self,
        expected: Sequence[dict[str, Any]],
        text: str,
        max_len: int,
        eos_id: int,
        input_key: str = "text",
        truncate: bool = False,
    )

Source from the content-addressed store, hash-verified

216 ),
217 )
218 def test_text_input(
219 self,
220 expected: Sequence[dict[str, Any]],
221 text: str,
222 max_len: int,
223 eos_id: int,
224 input_key: str = "text",
225 truncate: bool = False,
226 ):
227 spm_file = os.path.join(tokenizers_dir, "librispeech_unigram_1024.model")
228 vocab_cfg = config_for_class(seqio.SentencePieceVocabulary).set(
229 sentencepiece_model_file=spm_file
230 )
231 vocab = vocab_cfg.instantiate()
232 processor = input_tf_data.chain(
233 input_asr.text_input(
234 max_len=max_len,
235 vocab=vocab_cfg,
236 input_key=input_key,
237 truncate=truncate,
238 eos_id=eos_id,
239 ),
240 input_asr.make_autoregressive_inputs(vocab=vocab_cfg, bos_id=eos_id),
241 )
242 source = input_fake.fake_source(examples=[{input_key: text}], is_training=False)
243 actual = list(processor(source()))
244 for ex in expected:
245 for k, v in ex.items():
246 if k in ["input_ids", "target_labels"]:
247 v = [vocab.tokenizer.PieceToId(x) if isinstance(x, str) else x for x in v]
248 ex[k] = tf.constant(v)
249 tf.nest.map_structure(self.assertAllEqual, expected, actual)
250
251
252class SpeechTextInputTest(parameterized.TestCase, tf.test.TestCase):

Callers

nothing calls this directly

Calls 5

config_for_classFunction · 0.90
joinMethod · 0.80
itemsMethod · 0.80
setMethod · 0.45
instantiateMethod · 0.45

Tested by

no test coverage detected