(
self,
inputs: Sequence[dict],
expected: Sequence[dict],
max_speech_len: int,
max_text_len: int,
truncate: bool,
eos_id: int,
)
| 384 | ), |
| 385 | ) |
| 386 | def test_asr_input( |
| 387 | self, |
| 388 | inputs: Sequence[dict], |
| 389 | expected: Sequence[dict], |
| 390 | max_speech_len: int, |
| 391 | max_text_len: int, |
| 392 | truncate: bool, |
| 393 | eos_id: int, |
| 394 | ): |
| 395 | spm_file = os.path.join(tokenizers_dir, "librispeech_unigram_1024.model") |
| 396 | vocab_cfg = config_for_class(seqio.SentencePieceVocabulary).set( |
| 397 | sentencepiece_model_file=spm_file |
| 398 | ) |
| 399 | vocab = vocab_cfg.instantiate() |
| 400 | |
| 401 | source = input_fake.fake_source( |
| 402 | is_training=False, |
| 403 | examples=inputs, |
| 404 | spec=dict( |
| 405 | speech=tf.TensorSpec(shape=(None,), dtype=tf.int16), |
| 406 | text=tf.TensorSpec(shape=(), dtype=tf.string), |
| 407 | ), |
| 408 | ) |
| 409 | processor = input_tf_data.chain( |
| 410 | input_asr.speech_input(max_len=max_speech_len, normalize_by_scale=2.0), |
| 411 | input_asr.text_input( |
| 412 | max_len=max_text_len, vocab=vocab_cfg, truncate=truncate, eos_id=eos_id |
| 413 | ), |
| 414 | input_asr.make_autoregressive_inputs(vocab=vocab_cfg, bos_id=eos_id), |
| 415 | ) |
| 416 | actual = list(processor(source())) |
| 417 | self.assertEqual(len(expected), len(actual)) |
| 418 | for i, expect in enumerate(expected): |
| 419 | # Compare text fields separately from assertAllClose. |
| 420 | self.assertEqual(actual[i].pop("text"), expect.pop("text")) |
| 421 | for k, v in expect.items(): |
| 422 | if k in ["input_ids", "target_labels"]: |
| 423 | v = [vocab.tokenizer.PieceToId(x) if isinstance(x, str) else x for x in v] |
| 424 | expect[k] = tf.constant(v) |
| 425 | tf.nest.map_structure(self.assertAllClose, expected, actual) |
| 426 | |
| 427 | |
| 428 | class FilterTest(parameterized.TestCase, tf.test.TestCase): |
nothing calls this directly
no test coverage detected