Skip to content

Commit 6799bf8

Browse files
add phoneme control augmentation to multiturn dataloader (#16117)
* add phoneme control augmentation to multiturn dataloader Signed-off-by: paarthneekhara <paarth.n@gmail.com> * Apply suggestions from code review Co-authored-by: Jason <jasoli@nvidia.com> Signed-off-by: Jason <jasoli@nvidia.com> --------- Signed-off-by: paarthneekhara <paarth.n@gmail.com> Signed-off-by: Jason <jasoli@nvidia.com> Co-authored-by: Jason <jasoli@nvidia.com>
1 parent 3998328 commit 6799bf8

4 files changed

Lines changed: 171 additions & 9 deletions

File tree

examples/tts/conf/magpietts/easy_magpietts_lhotse_multiturn.yaml

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -93,6 +93,14 @@ model:
9393
phoneme_corruption_timestep_ratio: 0.15
9494
phoneme_corruption_unk_mode_prob: 0.5
9595
phoneme_corruption_type: "repeat_skip_unk" # "repeat_skip_unk" or "complete_channel"
96+
enable_phoneme_text_input: false # Enable inline IPA spans in text as <bop>...<eop>; requires phoneme_tokenizer.
97+
partial_phoneme_text_prob: 0.0 # Training-only probability of using precomputed IPA alignments.
98+
partial_phoneme_portion_min: 0.25 # Minimum portion of aligned words to phonemize in a selected sample.
99+
partial_phoneme_portion_max: 0.75 # Maximum portion of aligned words to phonemize in a selected sample.
100+
phoneme_text_bop_marker: "<bop>"
101+
phoneme_text_eop_marker: "<eop>"
102+
ignore_phoneme_languages: [] # Languages for which missing IPA phoneme fields are allowed during training.
103+
add_language_to_context_text: false # Prefix context text with language metadata for multilingual conditioning.
96104
phoneme_turn_dropout_batch_prob: 0.0 # prob of applying turn dropout to a sample
97105
phoneme_turn_dropout_turn_prob: 0.0 # prob of dropping each phoneme turn within a sample
98106
phoneme_turn_max_words_to_drop: 0 # turns with <= this many words keep phoneme tokens as pad_id

nemo/collections/tts/data/text_to_speech_dataset_lhotse_multiturn.py

Lines changed: 99 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -29,9 +29,13 @@
2929
from nemo.collections.speechlm2.parts.precision import fp32_precision
3030
from nemo.collections.tts.data.text_to_speech_dataset_lhotse import setup_tokenizers
3131
from nemo.collections.tts.parts.utils.tts_dataset_utils import (
32+
_sample_probability_range,
3233
beta_binomial_prior_distribution,
34+
has_phoneme_text_spans,
3335
normalize_volume,
36+
partially_phonemize_text,
3437
stack_tensors,
38+
tokenize_text_with_phoneme_spans,
3539
)
3640
from nemo.core.classes.common import safe_instantiate
3741
from nemo.utils import logging
@@ -150,6 +154,13 @@ def __init__(
150154
text_context_remapping_prob: float = 0.0,
151155
phoneme_tokenizer_config: DictConfig = None,
152156
ignore_phoneme_languages: List[str] = None,
157+
enable_phoneme_text_input: bool = False,
158+
text_phoneme_token_offset: int = None,
159+
partial_phoneme_text_prob: float = 0.0,
160+
partial_phoneme_portion_min: float = 0.25,
161+
partial_phoneme_portion_max: float = 0.75,
162+
phoneme_text_bop_marker: str = "<bop>",
163+
phoneme_text_eop_marker: str = "<eop>",
153164
add_language_to_context_text: bool = False,
154165
source_sample_rate: int = 16000,
155166
input_roles: List[str] = ["user", "User"],
@@ -183,6 +194,13 @@ def __init__(
183194
self.text_context_remapping_prob = text_context_remapping_prob
184195
self.phoneme_tokenizer_config = phoneme_tokenizer_config
185196
self.ignore_phoneme_languages = ignore_phoneme_languages or []
197+
self.enable_phoneme_text_input = enable_phoneme_text_input
198+
self.text_phoneme_token_offset = text_phoneme_token_offset
199+
self.partial_phoneme_text_prob = partial_phoneme_text_prob
200+
self.partial_phoneme_portion_min = partial_phoneme_portion_min
201+
self.partial_phoneme_portion_max = partial_phoneme_portion_max
202+
self.phoneme_text_bop_marker = phoneme_text_bop_marker
203+
self.phoneme_text_eop_marker = phoneme_text_eop_marker
186204
self.add_language_to_context_text = add_language_to_context_text
187205

188206
self.source_sample_rate = source_sample_rate
@@ -313,6 +331,16 @@ def _collate_text_channels(self, cuts: CutSet, batch_tokenizer_names: list[str])
313331
eos_id=self.eos_id,
314332
bos_id=self.bos_id,
315333
interruption_token_id=self.interruption_token_id,
334+
phoneme_tokenizer=self.phoneme_tokenizer,
335+
enable_phoneme_text_input=self.enable_phoneme_text_input,
336+
text_phoneme_token_offset=self.text_phoneme_token_offset,
337+
partial_phoneme_text_prob=self.partial_phoneme_text_prob,
338+
partial_phoneme_portion_min=self.partial_phoneme_portion_min,
339+
partial_phoneme_portion_max=self.partial_phoneme_portion_max,
340+
phoneme_text_bop_marker=self.phoneme_text_bop_marker,
341+
phoneme_text_eop_marker=self.phoneme_text_eop_marker,
342+
ignore_phoneme_languages=self.ignore_phoneme_languages,
343+
apply_partial_phoneme_text=self.dataset_type == 'train',
316344
)
317345
source_tokens, source_token_lens = collate_token_channel(
318346
cuts,
@@ -768,6 +796,16 @@ def collate_token_channel(
768796
eos_id: int = None,
769797
bos_id: int = None,
770798
interruption_token_id: int = None,
799+
phoneme_tokenizer=None,
800+
enable_phoneme_text_input: bool = False,
801+
text_phoneme_token_offset: int = None,
802+
partial_phoneme_text_prob: float = 0.0,
803+
partial_phoneme_portion_min: float = 0.25,
804+
partial_phoneme_portion_max: float = 0.75,
805+
phoneme_text_bop_marker: str = "<bop>",
806+
phoneme_text_eop_marker: str = "<eop>",
807+
ignore_phoneme_languages: list[str] = None,
808+
apply_partial_phoneme_text: bool = False,
771809
) -> tuple[torch.Tensor, torch.Tensor]:
772810
"""Build and collate token channels aligned to the audio frame grid."""
773811
tokens = []
@@ -786,6 +824,16 @@ def collate_token_channel(
786824
interruption_token_id,
787825
add_text_bos,
788826
tok_name,
827+
phoneme_tokenizer,
828+
enable_phoneme_text_input,
829+
text_phoneme_token_offset,
830+
partial_phoneme_text_prob,
831+
partial_phoneme_portion_min,
832+
partial_phoneme_portion_max,
833+
phoneme_text_bop_marker,
834+
phoneme_text_eop_marker,
835+
ignore_phoneme_languages,
836+
apply_partial_phoneme_text,
789837
)
790838
)
791839
token_lens = torch.tensor([len(tt) for tt in tokens])
@@ -831,22 +879,64 @@ def build_token_channel(
831879
interruption_token_id: int = -4,
832880
add_text_bos: bool = True,
833881
tokenizer_name: str = "english_phoneme",
882+
phoneme_tokenizer=None,
883+
enable_phoneme_text_input: bool = False,
884+
text_phoneme_token_offset: int = None,
885+
partial_phoneme_text_prob: float = 0.0,
886+
partial_phoneme_portion_min: float = 0.25,
887+
partial_phoneme_portion_max: float = 0.75,
888+
phoneme_text_bop_marker: str = "<bop>",
889+
phoneme_text_eop_marker: str = "<eop>",
890+
ignore_phoneme_languages: list[str] = None,
891+
apply_partial_phoneme_text: bool = False,
834892
) -> torch.Tensor:
835893

836894
total = compute_num_frames(cut.duration, frame_length, cut.sampling_rate)
837895
tokens = torch.ones(total, dtype=torch.long) * pad_id
838896

839897
for supervision in cut.supervisions:
840898
if supervision.speaker in roles:
841-
text = supervision.text
842-
843-
if hasattr(tokenizer, "encode"):
844-
try:
845-
raw_ids = tokenizer.encode(text=text, tokenizer_name=tokenizer_name)
846-
except TypeError:
847-
raw_ids = tokenizer.encode(text)
848-
else:
849-
raw_ids = tokenizer.text_to_ids(text)
899+
# TODO: Current multi-turn datasets do not contain the normalized_text field so check will always default to the else branch. Will need to evaluate whether it makes sense to keep this in future.
900+
# This code path is used for both multi-turn and single-turn datasets.
901+
text = supervision.normalized_text if supervision.has_custom("normalized_text") else supervision.text
902+
text_for_tokens = text
903+
language = cut.lang if cut.has_custom("lang") else supervision.language
904+
if (
905+
apply_partial_phoneme_text
906+
and enable_phoneme_text_input
907+
and partial_phoneme_text_prob > 0.0
908+
and language not in (ignore_phoneme_languages or [])
909+
and supervision.has_custom("ipa_alignment")
910+
and not has_phoneme_text_spans(
911+
text,
912+
bop_marker=phoneme_text_bop_marker,
913+
eop_marker=phoneme_text_eop_marker,
914+
)
915+
and random.random() < partial_phoneme_text_prob
916+
):
917+
sampled_portion = _sample_probability_range(
918+
"partial_phoneme_portion",
919+
partial_phoneme_portion_min,
920+
partial_phoneme_portion_max,
921+
)
922+
text_for_tokens = partially_phonemize_text(
923+
text=text,
924+
ipa_alignment=supervision.ipa_alignment,
925+
partial_phoneme_portion=sampled_portion,
926+
full_ipa_text=_get_supervision_ipa_text(supervision),
927+
bop_marker=phoneme_text_bop_marker,
928+
eop_marker=phoneme_text_eop_marker,
929+
)
930+
raw_ids = tokenize_text_with_phoneme_spans(
931+
text_tokenizer=tokenizer,
932+
text_str=text_for_tokens,
933+
tokenizer_name=tokenizer_name,
934+
enable_phoneme_text_input=enable_phoneme_text_input,
935+
phoneme_tokenizer=phoneme_tokenizer,
936+
text_phoneme_token_offset=text_phoneme_token_offset,
937+
bop_marker=phoneme_text_bop_marker,
938+
eop_marker=phoneme_text_eop_marker,
939+
)
850940

851941
if add_text_bos:
852942
text_ids = torch.as_tensor([bos_id] + raw_ids + [eos_id])

nemo/collections/tts/models/easy_magpietts.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1975,6 +1975,13 @@ def get_lhotse_dataloader(self, dataset_cfg, mode='train') -> torch.utils.data.D
19751975
tokenizer_config=self.cfg.text_tokenizers,
19761976
phoneme_tokenizer_config=self.cfg.get("phoneme_tokenizer", None),
19771977
ignore_phoneme_languages=self.cfg.get("ignore_phoneme_languages", []),
1978+
enable_phoneme_text_input=self.enable_phoneme_text_input,
1979+
text_phoneme_token_offset=self.text_phoneme_token_offset,
1980+
partial_phoneme_text_prob=self.partial_phoneme_text_prob if mode == 'train' else 0.0,
1981+
partial_phoneme_portion_min=self.partial_phoneme_portion_min,
1982+
partial_phoneme_portion_max=self.partial_phoneme_portion_max,
1983+
phoneme_text_bop_marker=self.phoneme_text_bop_marker,
1984+
phoneme_text_eop_marker=self.phoneme_text_eop_marker,
19781985
add_language_to_context_text=self.add_language_to_context_text,
19791986
source_sample_rate=self.sample_rate,
19801987
input_roles=["user", "User"],

tests/collections/tts/data/test_magpietts_dataset_lhotse.py

Lines changed: 57 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,23 @@
4747
pass
4848

4949

50+
class _FakeIPATokenizer:
51+
pad = 0
52+
bos_token_id = 1
53+
eos_token_id = 2
54+
55+
def encode(self, text):
56+
return [7, 8] if text else []
57+
58+
59+
class _FakeTextTokenizer:
60+
tokens = list(range(100))
61+
pad = 0
62+
63+
def encode(self, text, tokenizer_name):
64+
return [10 + len(text)]
65+
66+
5067
def _seed_everything():
5168
random.seed(42)
5269
np.random.seed(42)
@@ -140,6 +157,7 @@ def _multiturn_cutset():
140157
text="hello",
141158
language="en",
142159
speaker="assistant",
160+
custom={"ipa": "həloʊ", "ipa_alignment": [[0, 5, "hello", "həloʊ"]]},
143161
),
144162
SupervisionSegment(
145163
id="turn-user-1",
@@ -214,6 +232,45 @@ def test_single_turn_dataset_uses_bpe_and_cached_codes(self):
214232
assert batch["context_text_tokens_lens"].item() > 0
215233
assert batch["has_text_context"].tolist() == [True]
216234

235+
def test_multiturn_pronunciation_control_only_changes_target_turns(self):
236+
_seed_everything()
237+
kwargs = _dataset_kwargs()
238+
kwargs.update(
239+
{
240+
"codec_model_input_sample_rate": CODEC_MODEL_INPUT_SAMPLE_RATE,
241+
"frame_stacking_factor": FRAME_STACKING_FACTOR,
242+
"source_sample_rate": SAMPLE_RATE,
243+
"input_roles": ["user"],
244+
"output_roles": ["assistant"],
245+
"add_text_bos": False,
246+
"use_text_conditioning_tokenizer": False,
247+
"enable_phoneme_text_input": True,
248+
"partial_phoneme_text_prob": 1.0,
249+
"partial_phoneme_portion_min": 1.0,
250+
"partial_phoneme_portion_max": 1.0,
251+
}
252+
)
253+
dataset = MagpieTTSLhotseMultiturnDataset(**kwargs)
254+
dataset.text_tokenizer = _FakeTextTokenizer()
255+
dataset.phoneme_tokenizer = _FakeIPATokenizer()
256+
dataset.bos_id = len(dataset.text_tokenizer.tokens)
257+
dataset.eos_id = dataset.bos_id + 1
258+
dataset.cfg_unk_token_id = dataset.bos_id + 2
259+
dataset.interruption_token_id = dataset.bos_id + 3
260+
dataset.pad_id = dataset.text_tokenizer.pad
261+
dataset.text_phoneme_token_offset = dataset.bos_id + 4
262+
263+
batch = dataset[_multiturn_cutset()]
264+
265+
target_tokens = batch["text"][0]
266+
source_tokens = batch["source_tokens"][0]
267+
assert target_tokens[15:17].tolist() == [
268+
dataset.text_phoneme_token_offset + 7,
269+
dataset.text_phoneme_token_offset + 8,
270+
]
271+
assert torch.all(source_tokens < dataset.text_phoneme_token_offset)
272+
assert target_tokens[25].item() == dataset.interruption_token_id
273+
217274
def test_multiturn_dataset_uses_bpe_and_cached_codes(self):
218275
_seed_everything()
219276
kwargs = _dataset_kwargs()

0 commit comments

Comments
 (0)