77import math
88from dataclasses import dataclass
99from pathlib import Path
10- from typing import List , Optional , Sequence , Union
10+ from typing import Dict , List , Optional , Sequence , Union
1111
1212import librosa
1313import lightning as L
@@ -101,22 +101,55 @@ def __init__(
101101 self .keep_seconds_before_silence = keep_seconds_before_silence
102102 self .keep_seconds_after_silence = keep_seconds_after_silence
103103
104+ self .piper_config : Optional [PiperConfig ] = None
105+ self .is_multispeaker = self .num_speakers > 1
106+
104107 def prepare_data (self ):
105108 self .cache_dir .mkdir (parents = True , exist_ok = True )
106109
110+ self .piper_config = PiperConfig (
111+ num_symbols = self .num_symbols ,
112+ num_speakers = self .num_speakers ,
113+ sample_rate = self .sample_rate ,
114+ espeak_voice = self .espeak_voice ,
115+ phoneme_id_map = DEFAULT_PHONEME_ID_MAP ,
116+ phoneme_type = PhonemeType .ESPEAK ,
117+ piper_version = "1.3.0" ,
118+ )
119+
120+ speaker_id_map : Dict [str , int ] = {}
121+ if self .is_multispeaker :
122+ # Generate speaker id map
123+ with open (self .csv_path , "r" , encoding = "utf-8" ) as csv_file :
124+ reader = csv .reader (csv_file , delimiter = "|" )
125+ for row in reader :
126+ assert (
127+ len (row ) >= 3
128+ ), "Expected CSV columns for multi-speaker metadata: wav|speaker|text"
129+ speaker_name = row [1 ]
130+ if speaker_name in speaker_id_map :
131+ continue
132+
133+ speaker_id_map [speaker_name ] = len (speaker_id_map )
134+
135+ assert (
136+ len (speaker_id_map ) <= self .num_speakers
137+ ), "More speakers in metadata than num_speakers"
138+
139+ if len (speaker_id_map ) != self .num_speakers :
140+ _LOGGER .warning (
141+ "Expected %s speakers in the dataset, got %s" ,
142+ self .num_speakers ,
143+ len (speaker_id_map ),
144+ )
145+
146+ self .piper_config .speaker_id_map = speaker_id_map
147+
107148 # Write config
108149 self .config_path .parent .mkdir (parents = True , exist_ok = True )
109150 with open (self .config_path , "w" , encoding = "utf-8" ) as config_file :
110151 json .dump (
111- PiperConfig (
112- num_symbols = self .num_symbols ,
113- num_speakers = self .num_speakers ,
114- sample_rate = self .sample_rate ,
115- espeak_voice = self .espeak_voice ,
116- phoneme_id_map = DEFAULT_PHONEME_ID_MAP ,
117- phoneme_type = PhonemeType .ESPEAK ,
118- piper_version = "1.3.0" ,
119- ).to_dict (),
152+ self .piper_config .to_dict (),
120153 config_file ,
121154 ensure_ascii = False ,
122155 indent = 2 ,
@@ -131,6 +164,14 @@ def prepare_data(self):
131164 reader = csv .reader (csv_file , delimiter = "|" )
132165 for row_number , row in enumerate (reader , start = 1 ):
133166 utt_id , text = row [0 ], row [- 1 ]
167+ speaker_id : Optional [int ] = None
168+ if self .is_multispeaker :
169+ assert (
170+ len (row ) >= 3
171+ ), "Expected CSV columns for multi-speaker metadata: wav|speaker|text"
172+ speaker_name = row [1 ]
173+ speaker_id = speaker_id_map [speaker_name ]
174+
134175 audio_path = self .audio_dir / utt_id
135176 if not audio_path .exists ():
136177 audio_path = self .audio_dir / f"{ utt_id } .wav"
@@ -139,8 +180,9 @@ def prepare_data(self):
139180 _LOGGER .warning ("Missing audio file: %s" , audio_path )
140181 continue
141182
142- cache_id = get_cache_id (row_number , text )
183+ cache_id = get_cache_id (row_number , text , speaker_id = speaker_id )
143184
185+ # text
144186 text_path = self .cache_dir / f"{ cache_id } .txt"
145187 if not text_path .exists ():
146188 text_path .write_text (text , encoding = "utf-8" )
@@ -232,12 +274,23 @@ def prepare_data(self):
232274 _LOGGER .info ("Processed %s utterance(s)" , num_utterances )
233275
234276 def setup (self , stage : str ) -> None :
277+ assert self .piper_config is not None
278+
235279 all_utts : list [CachedUtterance ] = []
280+ speaker_id_map = self .piper_config .speaker_id_map
236281
237282 with open (self .csv_path , "r" , encoding = "utf-8" ) as csv_file :
238283 reader = csv .reader (csv_file , delimiter = "|" )
239284 for row_number , row in enumerate (reader , start = 1 ):
240285 utt_id , text = row [0 ], row [- 1 ]
286+ speaker_id : Optional [int ] = None
287+ if self .is_multispeaker :
288+ assert (
289+ len (row ) >= 3
290+ ), "Expected CSV columns for multi-speaker metadata: wav|speaker|text"
291+ speaker_name = row [1 ]
292+ speaker_id = speaker_id_map [speaker_name ]
293+
241294 audio_path = self .audio_dir / utt_id
242295 if not audio_path .exists ():
243296 audio_path = self .audio_dir / f"{ utt_id } .wav"
@@ -246,7 +299,7 @@ def setup(self, stage: str) -> None:
246299 _LOGGER .warning ("Missing audio file: %s" , audio_path )
247300 continue
248301
249- cache_id = get_cache_id (row_number , text )
302+ cache_id = get_cache_id (row_number , text , speaker_id = speaker_id )
250303
251304 phoneme_ids_path = self .cache_dir / f"{ cache_id } .phonemes.pt"
252305 if not phoneme_ids_path :
@@ -286,6 +339,7 @@ def setup(self, stage: str) -> None:
286339 audio_norm_path = audio_norm_path ,
287340 audio_spec_path = audio_spec_path ,
288341 text = text ,
342+ speaker_id = speaker_id ,
289343 )
290344 )
291345
0 commit comments