Skip to content

Commit 0dd2203

Browse files
committed
Add language ID
1 parent 7c206da commit 0dd2203

1 file changed

Lines changed: 28 additions & 4 deletions

File tree

src/fairseq2/datasets/asr_parquet.py

Lines changed: 28 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -8,9 +8,7 @@
88

99
from dataclasses import dataclass
1010
from 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

1513
from fairseq2.data import CollateOptionsOverride, Collater, DataPipelineBuilder
1614
from fairseq2.data.parquet import NamedColumns
@@ -26,6 +24,8 @@
2624
from fairseq2.logging import log
2725
from fairseq2.models.seq2seq import Seq2SeqBatch
2826

27+
from typing_extensions import override
28+
2929
PARQUET_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

Comments
 (0)