(self, texts: list[str], max_len: int, batch_size: int = 1)
| 692 | """Tests `text_to_lm_eval_input`.""" |
| 693 | |
| 694 | def _source(self, texts: list[str], max_len: int, batch_size: int = 1): |
| 695 | vocab_cls = with_regex_mapping( |
| 696 | seqio.SentencePieceVocabulary, |
| 697 | encode_mapping=[("\n", "<n>")], |
| 698 | decode_mapping=[("<n>", "\n")], |
| 699 | ) |
| 700 | vocab = vocab_cls( |
| 701 | sentencepiece_model_file=t5_sentence_piece_vocab_file, |
| 702 | ) |
| 703 | ds = fake_grain_source([{"text": text} for text in texts]) |
| 704 | ds = text_to_lm_eval_input(ds, vocab=vocab, max_len=max_len, stride=2) |
| 705 | if batch_size > 1: |
| 706 | ds = ds.batch(batch_size=batch_size) |
| 707 | ds = maybe_to_iter_dataset( |
| 708 | ds, |
| 709 | read_options=grain.ReadOptions(num_threads=1, prefetch_buffer_size=2), |
| 710 | ) |
| 711 | return ds |
| 712 | |
| 713 | @parameterized.parameters( |
| 714 | "How long is a piece of string?", |
no test coverage detected