Skip to content

Commit dbf4cc2

Browse files
committed
Fixes for multispeaker
1 parent d297536 commit dbf4cc2

3 files changed

Lines changed: 76 additions & 15 deletions

File tree

src/piper/train/vits/dataset.py

Lines changed: 66 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
import math
88
from dataclasses import dataclass
99
from pathlib import Path
10-
from typing import List, Optional, Sequence, Union
10+
from typing import Dict, List, Optional, Sequence, Union
1111

1212
import librosa
1313
import 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

src/piper/train/vits/utils.py

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,7 @@
22

33
import logging
44
from pathlib import Path
5-
from typing import Dict, Union
5+
from typing import Dict, Optional, Union
66

77
import numpy as np
88
import torch
@@ -56,6 +56,12 @@ def load_state_dict(model, saved_state_dict):
5656
model.load_state_dict(new_state_dict)
5757

5858

59-
def get_cache_id(row_number: int, text: str, max_length: int = 50) -> str:
60-
cache_id = str(row_number) + "_" + sanitize_filename(text)
59+
def get_cache_id(
60+
row_number: int, text: str, max_length: int = 50, speaker_id: Optional[int] = None
61+
) -> str:
62+
speaker_id_str = ""
63+
if speaker_id is not None:
64+
speaker_id_str = f"_{speaker_id}"
65+
66+
cache_id = str(row_number) + speaker_id_str + "_" + sanitize_filename(text)
6167
return cache_id[:max_length]

src/piper/voice.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -303,6 +303,7 @@ def synthesize(
303303
for phoneme in itertools.chain([BOS], phonemes, [EOS]):
304304
expected_ids = self.config.phoneme_id_map.get(phoneme, [])
305305

306+
ids_to_check: Sequence[int]
306307
if phoneme != EOS:
307308
ids_to_check = list(itertools.chain(expected_ids, pad_ids))
308309
else:

0 commit comments

Comments
 (0)