Skip to content

Commit 6fb8173

Browse files
authored
feat(data): add indexed resumable multimodal loading and exact token budgets (#16193)
* feat(data): add exact indexed multimodal runtime Add the final indexed JSONL, native tar, WDSv2, and ShareGPT routing runtime together with exact audio-token estimation, resumable packed sampling, explicit audio/manifest failure policies, and focused runtime tests. This is a data-only reconstruction from the original stacked work. SALM packed tensor shaping and model execution remain outside this topic. Co-authored-by: Kunal Dhawan <kunaldhawan97@gmail.com> Original-Commit: 8ef2ff0 Original-Commit: 08f4621 Original-Commit: 210a5b6bb6fc4e5eb11ca92d004820bc38cf836b Original-Commit: 3355af1c17fcb2425a3120848b30117470315538 Original-Commit: 303f36e6308bb67d80a882cff3c427325ea52d7d Original-Commit: e1f23096783973258534783d1c2dfcd6fe6f1302 Original-Commit: 95a77c3 Original-Commit: 64ed195 Original-Commit: 7176d60 Original-Commit: 8979b44 Original-Commit: 524c694 Original-Commit: acfdf7d Original-Commit: 3d7a2b8 Original-Commit: a8ea945 Original-Commit: 28a0a3a Original-Commit: fd7a491 Original-Commit: 958955b Original-Commit: 36f2503 Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * feat(data): add indexed data validation and operations Add index construction and conversion, atomic record validation and publication, ShareGPT route creation and relocation, full dataloader validation, materialized-SFT token accounting, and packed checkpoint progress analysis. The full validator intentionally uses dense strict SALM collation in this topic; packed SALM batch construction is reconstructed in the dependent training topic. Co-authored-by: Kunal Dhawan <kunaldhawan97@gmail.com> Original-Commit: 8ef2ff0 Original-Commit: 2791228 Original-Commit: 64ed195 Original-Commit: 9c15fb2 Original-Commit: 262b771 Original-Commit: 524c694 Original-Commit: acfdf7d Original-Commit: 8c7da42 Original-Commit: a591a8b Original-Commit: 958955b Original-Commit: bbfbc82 Original-Commit: a9e80e0 Original-Commit: 36f2503 Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * style(data): satisfy formatting and lint checks Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * build: pin Lhotse main for indexed data APIs Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * fix(data): close datastore readers after caching Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * fix(data): address indexed loading review Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * fix(data): preserve and override conversation system prompts Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * fix(data): preserve indexed tar byte ranges Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * fix(data): defer indexed tar validation to audio loading Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * build: pin Lhotse 2.0.0a5 Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * fix(data): embed native tar routing in index packs Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * build: pin Lhotse 2.0.0a6 Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * fix(data): preserve fixed-bucket caps in packed sampling Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * perf(data): parallelize native tar route construction Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * fix(data): diagnose invalid multimodal conversation turns Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * docs(data): document packed caps and route workers Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> * test(data): add sampling rate to tarred fixture Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com> --------- Signed-off-by: Piotr Żelasko <pzelasko@nvidia.com>
1 parent ea2c046 commit 6fb8173

79 files changed

Lines changed: 20936 additions & 847 deletions

File tree

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

docs/source/dataloaders.rst

Lines changed: 193 additions & 9 deletions
Large diffs are not rendered by default.

docs/source/speechlm2/datasets.rst

Lines changed: 579 additions & 1 deletion
Large diffs are not rendered by default.

nemo/collections/asr/data/audio_to_diar_label_lhotse.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -23,11 +23,12 @@
2323
get_hidden_length_from_sample_length,
2424
speaker_to_target,
2525
)
26+
from nemo.collections.common.data.lhotse.audio_loading import LhotseAudioLoadingDatasetMixin
2627
from nemo.core.neural_types import AudioSignal, LabelsType, LengthsType, NeuralType
2728
from nemo.utils import logging
2829

2930

30-
class LhotseAudioToSpeechE2ESpkDiarDataset(torch.utils.data.Dataset):
31+
class LhotseAudioToSpeechE2ESpkDiarDataset(LhotseAudioLoadingDatasetMixin, torch.utils.data.Dataset):
3132
"""
3233
This dataset is a Lhotse version of diarization dataset in audio_to_diar_label.py.
3334
Unlike native NeMo datasets, Lhotse dataset defines only the mapping from
@@ -94,7 +95,7 @@ def __getitem__(self, cuts) -> Tuple[torch.Tensor, ...]:
9495
speaker_activities.append(speaker_activity)
9596

9697
cuts = type(cuts).from_cuts(mono_cuts)
97-
audio, audio_lens, cuts = self.load_audio(cuts)
98+
audio, audio_lens, cuts = self.load_audio_with_cuts(cuts)
9899
max_num_speakers = max(1, max(activity.shape[1] for activity in speaker_activities))
99100
speaker_activities = [
100101
torch.nn.functional.pad(activity, (0, max_num_speakers - activity.shape[1]))

nemo/collections/asr/data/audio_to_eou_label_lhotse.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -26,6 +26,7 @@
2626

2727
from nemo.collections.asr.parts.preprocessing.perturb import process_augmentations
2828
from nemo.collections.asr.parts.preprocessing.segment import AudioSegment
29+
from nemo.collections.common.data.lhotse.audio_loading import LhotseAudioLoadingDatasetMixin
2930
from nemo.collections.common.tokenizers.aggregate_tokenizer import TokenizerWrapper
3031
from nemo.collections.common.tokenizers.tokenizer_spec import TokenizerSpec
3132
from nemo.core.neural_types import AudioSignal, LabelsType, LengthsType, NeuralType
@@ -73,7 +74,7 @@ class RandomPaddingConfig:
7374
post_pad_duration: float = 3.0 # amount of right-padding when pad_distribution='constant'
7475

7576

76-
class LhotseSpeechToTextBpeEOUDataset(torch.utils.data.Dataset):
77+
class LhotseSpeechToTextBpeEOUDataset(LhotseAudioLoadingDatasetMixin, torch.utils.data.Dataset):
7778
"""
7879
This dataset processes the audio data and the corresponding text data to generate the ASR labels,
7980
along with EOU labels for each frame. The audios used in this dataset should only contain speech with
@@ -206,7 +207,7 @@ def _check_special_tokens(self, tokenizer: TokenizerSpec):
206207
)
207208

208209
def __getitem__(self, cuts: CutSet) -> AudioToTextEOUBatch:
209-
audio, audio_lens, cuts = self.load_audio(cuts)
210+
audio, audio_lens, cuts = self.load_audio_with_cuts(cuts)
210211
audio_signals = []
211212
audio_lengths = []
212213
eou_targets = []

nemo/collections/asr/data/audio_to_text_lhotse.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -21,12 +21,13 @@
2121
from lhotse.dataset import AudioSamples
2222
from lhotse.dataset.collation import collate_vectors
2323

24+
from nemo.collections.common.data.lhotse.audio_loading import LhotseAudioLoadingDatasetMixin
2425
from nemo.collections.common.tokenizers.aggregate_tokenizer import TokenizerWrapper
2526
from nemo.collections.common.tokenizers.tokenizer_spec import TokenizerSpec
2627
from nemo.core.neural_types import AudioSignal, LabelsType, LengthsType, NeuralType
2728

2829

29-
class LhotseSpeechToTextBpeDataset(torch.utils.data.Dataset):
30+
class LhotseSpeechToTextBpeDataset(LhotseAudioLoadingDatasetMixin, torch.utils.data.Dataset):
3031
"""
3132
This dataset is based on BPE datasets from audio_to_text.py.
3233
Unlike native NeMo datasets, Lhotse dataset defines only the mapping from
@@ -65,7 +66,7 @@ def __init__(self, tokenizer: TokenizerSpec, return_cuts: bool = False):
6566
self.return_cuts = return_cuts
6667

6768
def __getitem__(self, cuts) -> Tuple[torch.Tensor, ...]:
68-
audio, audio_lens, cuts = self.load_audio(cuts)
69+
audio, audio_lens, cuts = self.load_audio_with_cuts(cuts)
6970
tokens = [
7071
torch.cat(
7172
[

nemo/collections/asr/data/audio_to_text_lhotse_prompt.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -21,12 +21,13 @@
2121
from lhotse.dataset import AudioSamples
2222
from lhotse.dataset.collation import collate_matrices, collate_vectors
2323

24+
from nemo.collections.common.data.lhotse.audio_loading import LhotseAudioLoadingDatasetMixin
2425
from nemo.collections.common.tokenizers.aggregate_tokenizer import AggregateTokenizer
2526
from nemo.collections.common.tokenizers.tokenizer_spec import TokenizerSpec
2627
from nemo.core.neural_types import AudioSignal, LabelsType, LengthsType, NeuralType
2728

2829

29-
class LhotseSpeechToTextBpeDatasetWithPrompt(torch.utils.data.Dataset):
30+
class LhotseSpeechToTextBpeDatasetWithPrompt(LhotseAudioLoadingDatasetMixin, torch.utils.data.Dataset):
3031
"""
3132
Dataset class for speech-to-text with prompt vectors.
3233
Supports both language ID and custom prompts.
@@ -115,7 +116,7 @@ def get_hidden_length_from_sample_length(self, num_samples: int) -> int:
115116
return int(hidden_length)
116117

117118
def __getitem__(self, cuts) -> Tuple[torch.Tensor, ...]:
118-
audio, audio_lens, cuts = self.load_audio(cuts)
119+
audio, audio_lens, cuts = self.load_audio_with_cuts(cuts)
119120
tokens = [torch.as_tensor(self.tokenizer(c.supervisions[0].text, c.supervisions[0].language)) for c in cuts]
120121

121122
# Create prompt targets

nemo/collections/asr/data/audio_to_text_lhotse_prompt_index.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -26,13 +26,14 @@
2626
from lhotse.dataset import AudioSamples
2727
from lhotse.dataset.collation import collate_vectors
2828

29+
from nemo.collections.common.data.lhotse.audio_loading import LhotseAudioLoadingDatasetMixin
2930
from nemo.collections.common.tokenizers.aggregate_tokenizer import TokenizerWrapper
3031
from nemo.collections.common.tokenizers.tokenizer_spec import TokenizerSpec
3132
from nemo.core.neural_types import AudioSignal, LabelsType, LengthsType, NeuralType
3233
from nemo.utils import logging
3334

3435

35-
class LhotseSpeechToTextBpeDatasetWithPromptIndex(torch.utils.data.Dataset):
36+
class LhotseSpeechToTextBpeDatasetWithPromptIndex(LhotseAudioLoadingDatasetMixin, torch.utils.data.Dataset):
3637
"""
3738
Simplified dataset class for speech-to-text with prompt support.
3839
@@ -136,7 +137,7 @@ def _get_prompt_index_for_cut(self, cut) -> int:
136137
return self._get_prompt_index(cut.supervisions[0].language)
137138

138139
def __getitem__(self, cuts) -> Tuple[torch.Tensor, ...]:
139-
audio, audio_lens, cuts = self.load_audio(cuts)
140+
audio, audio_lens, cuts = self.load_audio_with_cuts(cuts)
140141
tokens = [torch.as_tensor(self.tokenizer(c.supervisions[0].text, c.supervisions[0].language)) for c in cuts]
141142

142143
# Get prompt indices (just the language ID per sample, NOT full tensors)

nemo/collections/asr/data/audio_to_text_lhotse_prompted.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@
2323

2424
from nemo.collections.asr.data.audio_to_text_lhotse import _make_audio_samples
2525
from nemo.collections.common.data import apply_prompt_format_fn
26+
from nemo.collections.common.data.lhotse.audio_loading import LhotseAudioLoadingDatasetMixin
2627
from nemo.collections.common.prompts import PromptFormatter
2728
from nemo.collections.common.tokenizers import TokenizerSpec
2829

@@ -48,7 +49,7 @@ def get_decoder_inputs_outputs(self) -> tuple[torch.Tensor, torch.Tensor]:
4849
return self.prompted_transcript[:, :-1], self.prompted_transcript[:, 1:]
4950

5051

51-
class PromptedAudioToTextLhotseDataset(torch.utils.data.Dataset):
52+
class PromptedAudioToTextLhotseDataset(LhotseAudioLoadingDatasetMixin, torch.utils.data.Dataset):
5253
"""
5354
This dataset is based on :class:`~nemo.collections.asr.data.audio_to_text_lhotse.LhotseSpeechToTextBpeDataset`.
5455
It is a Lhotse-style dataset that converts a mini-batch of Cuts into tensors.
@@ -97,7 +98,7 @@ def __init__(
9798

9899
def __getitem__(self, cuts: CutSet) -> PromptedAudioToTextMiniBatch:
99100
# Load the audio's from AIS and add them to the CutSet
100-
audio, audio_lens, cuts = self.load_audio(cuts)
101+
audio, audio_lens, cuts = self.load_audio_with_cuts(cuts)
101102

102103
# Will work if batch_size is set to 1.
103104
if self.enable_chunking:

nemo/collections/asr/data/audio_to_text_lhotse_speaker.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -22,11 +22,12 @@
2222

2323
from nemo.collections.asr.data.audio_to_text_lhotse import TokenizerWrapper
2424
from nemo.collections.asr.parts.utils.asr_multispeaker_utils import speaker_to_target
25+
from nemo.collections.common.data.lhotse.audio_loading import LhotseAudioLoadingDatasetMixin
2526
from nemo.collections.common.tokenizers.tokenizer_spec import TokenizerSpec
2627
from nemo.core.neural_types import AudioSignal, LabelsType, LengthsType, NeuralType
2728

2829

29-
class LhotseSpeechToTextSpkBpeDataset(torch.utils.data.Dataset):
30+
class LhotseSpeechToTextSpkBpeDataset(LhotseAudioLoadingDatasetMixin, torch.utils.data.Dataset):
3031
"""
3132
This dataset is based on BPE datasets from audio_to_text.py. It has the same functionality of LhotseSpeechToTextBpeDataset but also yield speaker target tensor.
3233
Unlike native NeMo datasets, Lhotse dataset defines only the mapping from
@@ -60,7 +61,7 @@ def __init__(self, cfg, tokenizer: TokenizerSpec):
6061

6162
def __getitem__(self, cuts) -> Tuple[torch.Tensor, ...]:
6263

63-
audio, audio_lens, cuts = self.load_audio(cuts)
64+
audio, audio_lens, cuts = self.load_audio_with_cuts(cuts)
6465

6566
tokens = []
6667
spk_targets = []

nemo/collections/asr/data/ssl_dataset.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@
3232
from nemo.collections.asr.parts.preprocessing.segment import AudioSegment
3333
from nemo.collections.asr.parts.utils.manifest_utils import read_manifest
3434
from nemo.collections.common.data.dataset import ConcatDataset
35+
from nemo.collections.common.data.lhotse.audio_loading import LhotseAudioLoadingDatasetMixin
3536
from nemo.collections.common.parts.preprocessing.manifest import get_full_path
3637
from nemo.core.classes import Serialization
3738
from nemo.utils import logging
@@ -438,7 +439,7 @@ def _collate_fn(self, batch: List[AudioNoiseItem]) -> AudioNoiseBatch:
438439
return _audio_noise_collate_fn(batch, self.batch_augmentor)
439440

440441

441-
class LhotseAudioNoiseDataset(torch.utils.data.Dataset):
442+
class LhotseAudioNoiseDataset(LhotseAudioLoadingDatasetMixin, torch.utils.data.Dataset):
442443
def __init__(self, noise_manifest: str | None = None, batch_augmentor_cfg: DictConfig = None):
443444
super().__init__()
444445

@@ -453,7 +454,7 @@ def __init__(self, noise_manifest: str | None = None, batch_augmentor_cfg: DictC
453454

454455
def __getitem__(self, cuts):
455456

456-
audios, audio_lens, cuts = self.load_audio(cuts)
457+
audios, audio_lens, cuts = self.load_audio_with_cuts(cuts)
457458
if len(self.noise_data) > 0:
458459
sampled_noises = [sample_noise(self.noise_data, cut.sampling_rate, cut.num_samples) for cut in cuts]
459460
sampled_noises, sampled_noises_lens = zip(*sampled_noises)

0 commit comments

Comments
 (0)