2929from nemo .collections .speechlm2 .parts .precision import fp32_precision
3030from nemo .collections .tts .data .text_to_speech_dataset_lhotse import setup_tokenizers
3131from 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)
3640from nemo .core .classes .common import safe_instantiate
3741from 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 ])
0 commit comments