88
99from dataclasses import dataclass
1010from functools import partial
11- from typing import Final , List , Tuple , final
12-
13- from typing_extensions import override
11+ from typing import Final , final , List , Tuple
1412
1513from fairseq2 .data import CollateOptionsOverride , Collater , DataPipelineBuilder
1614from fairseq2 .data .parquet import NamedColumns
2624from fairseq2 .logging import log
2725from fairseq2 .models .seq2seq import Seq2SeqBatch
2826
27+ from typing_extensions import override
28+
2929PARQUET_ASR_DATASET_FAMILY : Final = "generic_parquet_asr"
3030
3131
@@ -154,6 +154,11 @@ def build_parquet_audio_text_reading(
154154 builder , options , GenericSpeechDataset .rename_feature
155155 )
156156
157+ # Tokenizer langauge name as well
158+ builder = GenericAsrParquetDataset .add_lang_tokenization_pipeline (
159+ builder , tokenizer
160+ )
161+
157162 # Collate bucketed examples into a batch.
158163 text_collate_opts = CollateOptionsOverride (
159164 "text" , pad_value = tokenizer .vocab_info .pad_idx
@@ -170,8 +175,27 @@ def build_parquet_audio_text_reading(
170175 # Prefetch `num_prefetch` batches in background.
171176 builder .prefetch (options .num_prefetch )
172177
173- # Wrap examples with `Seq2SeqBatch`.
178+ builder = builder .map (
179+ lambda x : x .to (gang .device ),
180+ selector = "lang_tokens.seqs,lang_tokens.seq_lens" ,
181+ )
174182
183+ # Wrap examples with `Seq2SeqBatch`.
175184 builder = builder .map (partial (GenericAsrDataset .to_batch , device = gang .device ))
176185
177186 return builder
187+
188+ @staticmethod
189+ def add_lang_tokenization_pipeline (
190+ builder : DataPipelineBuilder ,
191+ tokenizer : TextTokenizer ,
192+ ) -> DataPipelineBuilder :
193+ text_encoder = tokenizer .create_encoder ()
194+
195+ def tokenizer_lang (batch ):
196+ for bb in batch :
197+ bb ["lang_tokens" ] = text_encoder (bb ["lang" ].lower ())
198+ return batch
199+
200+ builder .map (tokenizer_lang )
201+ return builder
0 commit comments