(
self,
expected: Sequence[dict[str, Any]],
text: str,
max_len: int,
eos_id: int,
input_key: str = "text",
truncate: bool = False,
)
| 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 | |
| 252 | class SpeechTextInputTest(parameterized.TestCase, tf.test.TestCase): |
nothing calls this directly
no test coverage detected