MyanmarTTS 1.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- myanmar_tts/__init__.py +42 -0
- myanmar_tts/api.py +82 -0
- myanmar_tts/burmese.py +42 -0
- myanmar_tts/config.py +50 -0
- myanmar_tts/datas/__init__.py +0 -0
- myanmar_tts/datas/all_sources_dataset.py +160 -0
- myanmar_tts/datas/collate_wav.py +39 -0
- myanmar_tts/datas/dataset.py +69 -0
- myanmar_tts/datas/local_shards_dataset.py +73 -0
- myanmar_tts/models/__init__.py +0 -0
- myanmar_tts/models/diffusion_transformer.py +205 -0
- myanmar_tts/models/duration_predictor.py +40 -0
- myanmar_tts/models/estimator.py +138 -0
- myanmar_tts/models/flow_matching.py +100 -0
- myanmar_tts/models/model.py +178 -0
- myanmar_tts/models/reference_encoder.py +168 -0
- myanmar_tts/models/text_encoder.py +44 -0
- myanmar_tts/monotonic_align/__init__.py +16 -0
- myanmar_tts/monotonic_align/core.py +46 -0
- myanmar_tts/symbols.py +57 -0
- myanmar_tts/text/LICENSE +19 -0
- myanmar_tts/text/__init__.py +16 -0
- myanmar_tts/text/burmese.py +42 -0
- myanmar_tts/text/cleaners.py +10 -0
- myanmar_tts/text/symbols.py +57 -0
- myanmar_tts/utils/__init__.py +0 -0
- myanmar_tts/utils/audio.py +74 -0
- myanmar_tts/utils/load.py +43 -0
- myanmar_tts/utils/mask.py +8 -0
- myanmar_tts/utils/scheduler.py +428 -0
- myanmar_tts/vocoders/__init__.py +0 -0
- myanmar_tts/vocoders/ffgan/__init__.py +0 -0
- myanmar_tts/vocoders/ffgan/backbone.py +214 -0
- myanmar_tts/vocoders/ffgan/head.py +257 -0
- myanmar_tts/vocoders/ffgan/model.py +57 -0
- myanmar_tts/vocoders/ffgan/unify.py +60 -0
- myanmar_tts/vocoders/vocos/README.md +41 -0
- myanmar_tts/vocoders/vocos/__init__.py +0 -0
- myanmar_tts/vocoders/vocos/config.py +41 -0
- myanmar_tts/vocoders/vocos/dataset.py +57 -0
- myanmar_tts/vocoders/vocos/inference.ipynb +79 -0
- myanmar_tts/vocoders/vocos/models/__init__.py +0 -0
- myanmar_tts/vocoders/vocos/models/backbone.py +57 -0
- myanmar_tts/vocoders/vocos/models/discriminator.py +171 -0
- myanmar_tts/vocoders/vocos/models/head.py +118 -0
- myanmar_tts/vocoders/vocos/models/loss.py +66 -0
- myanmar_tts/vocoders/vocos/models/model.py +20 -0
- myanmar_tts/vocoders/vocos/models/module.py +47 -0
- myanmar_tts/vocoders/vocos/preprocess.py +45 -0
- myanmar_tts/vocoders/vocos/requirements.txt +2 -0
- myanmar_tts/vocoders/vocos/train.py +165 -0
- myanmar_tts/vocoders/vocos/utils/__init__.py +0 -0
- myanmar_tts/vocoders/vocos/utils/audio.py +74 -0
- myanmar_tts/vocoders/vocos/utils/load.py +53 -0
- myanmar_tts/vocoders/vocos/utils/scheduler.py +298 -0
- myanmartts-1.0.0.dist-info/METADATA +173 -0
- myanmartts-1.0.0.dist-info/RECORD +59 -0
- myanmartts-1.0.0.dist-info/WHEEL +4 -0
- myanmartts-1.0.0.dist-info/licenses/LICENSE +21 -0
myanmar_tts/__init__.py
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
1
|
+
"""MyanmarTTS - Burmese text-to-speech, from scratch, CC0."""
|
|
2
|
+
import os
|
|
3
|
+
from huggingface_hub import hf_hub_download
|
|
4
|
+
|
|
5
|
+
HF_REPO = "freococo/MyanmarTTS"
|
|
6
|
+
MODEL_FILE = "model_fp16.safetensors"
|
|
7
|
+
VOCODER_FILE = "vocos.pt"
|
|
8
|
+
REF_AUDIO = "samples/sample_0.wav"
|
|
9
|
+
|
|
10
|
+
__version__ = "1.0.0"
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class MyanmarTTS:
|
|
14
|
+
def __init__(self, device="cuda"):
|
|
15
|
+
self.device = device
|
|
16
|
+
self._model = None
|
|
17
|
+
self._ref = None
|
|
18
|
+
|
|
19
|
+
def _load(self):
|
|
20
|
+
if self._model is not None:
|
|
21
|
+
return
|
|
22
|
+
print(f"Downloading model files from {HF_REPO}...")
|
|
23
|
+
model_path = hf_hub_download(repo_id=HF_REPO, filename=MODEL_FILE)
|
|
24
|
+
vocos_path = hf_hub_download(repo_id=HF_REPO, filename=VOCODER_FILE)
|
|
25
|
+
self._ref = hf_hub_download(repo_id=HF_REPO, filename=REF_AUDIO)
|
|
26
|
+
|
|
27
|
+
from .api import StableTTSAPI
|
|
28
|
+
print("Loading model...")
|
|
29
|
+
self._model = StableTTSAPI(model_path, vocos_path, "vocos")
|
|
30
|
+
self._model.to(self.device)
|
|
31
|
+
print("Ready.")
|
|
32
|
+
|
|
33
|
+
def tts(self, text):
|
|
34
|
+
"""Return a numpy float32 waveform at 44100 Hz."""
|
|
35
|
+
self._load()
|
|
36
|
+
audio, _ = self._model.inference(
|
|
37
|
+
text, self._ref, "burmese", step=32, solver="dopri5", cfg=3.0
|
|
38
|
+
)
|
|
39
|
+
return audio.squeeze(0).cpu().numpy()
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
__all__ = ["MyanmarTTS"]
|
myanmar_tts/api.py
ADDED
|
@@ -0,0 +1,82 @@
|
|
|
1
|
+
import torch
|
|
2
|
+
import torch.nn as nn
|
|
3
|
+
from dataclasses import asdict
|
|
4
|
+
|
|
5
|
+
from .utils.audio import LogMelSpectrogram
|
|
6
|
+
from .config import ModelConfig, MelConfig
|
|
7
|
+
from .models.model import StableTTS
|
|
8
|
+
|
|
9
|
+
from .text import symbols
|
|
10
|
+
from .text import cleaned_text_to_sequence
|
|
11
|
+
from .text.burmese import burmese_to_ipa2
|
|
12
|
+
from .datas.dataset import intersperse
|
|
13
|
+
from .utils.audio import load_and_resample_audio
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def get_vocoder(model_path, model_name='vocos'):
|
|
17
|
+
if model_name == 'vocos':
|
|
18
|
+
from .vocoders.vocos.models.model import Vocos
|
|
19
|
+
from .config import VocosConfig, MelConfig
|
|
20
|
+
vocoder = Vocos(VocosConfig(), MelConfig())
|
|
21
|
+
vocoder.load_state_dict(torch.load(model_path, weights_only=True, map_location='cpu'))
|
|
22
|
+
vocoder.eval()
|
|
23
|
+
else:
|
|
24
|
+
raise NotImplementedError(f"Unsupported vocoder: {model_name}")
|
|
25
|
+
return vocoder
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
class StableTTSAPI(nn.Module):
|
|
29
|
+
def __init__(self, tts_model_path, vocoder_model_path, vocoder_name='vocos'):
|
|
30
|
+
super().__init__()
|
|
31
|
+
self.mel_config = MelConfig()
|
|
32
|
+
self.tts_model_config = ModelConfig()
|
|
33
|
+
|
|
34
|
+
self.mel_extractor = LogMelSpectrogram(**asdict(self.mel_config))
|
|
35
|
+
|
|
36
|
+
self.tts_model = StableTTS(len(symbols), self.mel_config.n_mels, **asdict(self.tts_model_config))
|
|
37
|
+
|
|
38
|
+
if tts_model_path.endswith(".safetensors"):
|
|
39
|
+
from safetensors.torch import load_file
|
|
40
|
+
state = load_file(tts_model_path)
|
|
41
|
+
else:
|
|
42
|
+
state = torch.load(tts_model_path, map_location='cpu', weights_only=True)
|
|
43
|
+
self.tts_model.load_state_dict(state)
|
|
44
|
+
self.tts_model.eval()
|
|
45
|
+
|
|
46
|
+
self.vocoder_model = get_vocoder(vocoder_model_path, vocoder_name)
|
|
47
|
+
self.vocoder_model.eval()
|
|
48
|
+
|
|
49
|
+
self.g2p_mapping = {
|
|
50
|
+
'burmese': burmese_to_ipa2,
|
|
51
|
+
}
|
|
52
|
+
self.supported_languages = self.g2p_mapping.keys()
|
|
53
|
+
|
|
54
|
+
@torch.inference_mode()
|
|
55
|
+
def inference(self, text, ref_audio, language, step, temperature=1.0,
|
|
56
|
+
length_scale=1.0, solver=None, cfg=3.0):
|
|
57
|
+
device = next(self.parameters()).device
|
|
58
|
+
phonemizer = self.g2p_mapping.get(language)
|
|
59
|
+
if phonemizer is None:
|
|
60
|
+
raise ValueError(f"Unsupported language: {language}")
|
|
61
|
+
|
|
62
|
+
text = phonemizer(text)
|
|
63
|
+
text = torch.tensor(
|
|
64
|
+
intersperse(cleaned_text_to_sequence(text), item=0),
|
|
65
|
+
dtype=torch.long, device=device
|
|
66
|
+
).unsqueeze(0)
|
|
67
|
+
text_length = torch.tensor([text.size(-1)], dtype=torch.long, device=device)
|
|
68
|
+
|
|
69
|
+
ref_audio = load_and_resample_audio(ref_audio, self.mel_config.sample_rate).to(device)
|
|
70
|
+
ref_audio = self.mel_extractor(ref_audio)
|
|
71
|
+
|
|
72
|
+
mel_output = self.tts_model.synthesise(
|
|
73
|
+
text, text_length, step, temperature, ref_audio,
|
|
74
|
+
length_scale, solver, cfg
|
|
75
|
+
)['decoder_outputs']
|
|
76
|
+
audio_output = self.vocoder_model(mel_output)
|
|
77
|
+
return audio_output.cpu(), mel_output.cpu()
|
|
78
|
+
|
|
79
|
+
def get_params(self):
|
|
80
|
+
tts_param = sum(p.numel() for p in self.tts_model.parameters()) / 1e6
|
|
81
|
+
vocoder_param = sum(p.numel() for p in self.vocoder_model.parameters()) / 1e6
|
|
82
|
+
return tts_param, vocoder_param
|
myanmar_tts/burmese.py
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Character-level text frontend for Burmese.
|
|
3
|
+
No real G2P: we clean the text and pass characters through directly.
|
|
4
|
+
Matches the interface of english_to_ipa2 / japanese_to_ipa2 (returns a list).
|
|
5
|
+
"""
|
|
6
|
+
import re
|
|
7
|
+
|
|
8
|
+
# Burmese Unicode ranges
|
|
9
|
+
_BURMESE_RANGES = (
|
|
10
|
+
(0x1000, 0x109F), # Myanmar
|
|
11
|
+
(0xAA60, 0xAA7F), # Myanmar Extended-A
|
|
12
|
+
(0xA9E0, 0xA9FF), # Myanmar Extended-B
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
# Allow these through the cleaner; everything else gets dropped
|
|
16
|
+
_ALLOWED_ASCII_PUNCT = set(".,!?'\"-:;()")
|
|
17
|
+
_ALLOWED_BURMESE_PUNCT = set("\u104A\u104B\u104C\u104D\u104E\u104F")
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def _is_burmese(ch):
|
|
21
|
+
o = ord(ch)
|
|
22
|
+
return any(lo <= o <= hi for lo, hi in _BURMESE_RANGES)
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def clean_burmese(text):
|
|
26
|
+
"""Keep only Burmese chars, spaces, and whitelisted punctuation."""
|
|
27
|
+
out = []
|
|
28
|
+
for ch in text:
|
|
29
|
+
if ch in (" ", "\t", "\n", "\r"):
|
|
30
|
+
out.append(" ")
|
|
31
|
+
elif _is_burmese(ch):
|
|
32
|
+
out.append(ch)
|
|
33
|
+
elif ch in _ALLOWED_ASCII_PUNCT or ch in _ALLOWED_BURMESE_PUNCT:
|
|
34
|
+
out.append(ch)
|
|
35
|
+
# else: drop silently
|
|
36
|
+
return re.sub(r"\s+", " ", "".join(out)).strip()
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def burmese_to_ipa2(text):
|
|
40
|
+
"""Return a list of character tokens (matches the IPA-G2P interface)."""
|
|
41
|
+
cleaned = clean_burmese(text)
|
|
42
|
+
return list(cleaned)
|
myanmar_tts/config.py
ADDED
|
@@ -0,0 +1,50 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
|
|
3
|
+
@dataclass
|
|
4
|
+
class MelConfig:
|
|
5
|
+
sample_rate: int = 44100
|
|
6
|
+
n_fft: int = 2048
|
|
7
|
+
win_length: int = 2048
|
|
8
|
+
hop_length: int = 512
|
|
9
|
+
f_min: float = 0.0
|
|
10
|
+
f_max: float = None
|
|
11
|
+
pad: int = 0
|
|
12
|
+
n_mels: int = 128
|
|
13
|
+
center: bool = False
|
|
14
|
+
pad_mode: str = "reflect"
|
|
15
|
+
mel_scale: str = "slaney"
|
|
16
|
+
|
|
17
|
+
def __post_init__(self):
|
|
18
|
+
if self.pad == 0:
|
|
19
|
+
self.pad = (self.n_fft - self.hop_length) // 2
|
|
20
|
+
|
|
21
|
+
@dataclass
|
|
22
|
+
class ModelConfig:
|
|
23
|
+
hidden_channels: int = 256
|
|
24
|
+
filter_channels: int = 1024
|
|
25
|
+
n_heads: int = 4
|
|
26
|
+
n_enc_layers: int = 3
|
|
27
|
+
n_dec_layers: int = 6
|
|
28
|
+
kernel_size: int = 3
|
|
29
|
+
p_dropout: int = 0.1
|
|
30
|
+
gin_channels: int = 256
|
|
31
|
+
|
|
32
|
+
@dataclass
|
|
33
|
+
class TrainConfig:
|
|
34
|
+
train_dataset_path: str = 'filelists/filelist.json'
|
|
35
|
+
test_dataset_path: str = 'filelists/filelist.json' # not used
|
|
36
|
+
batch_size: int = 32
|
|
37
|
+
learning_rate: float = 1e-4
|
|
38
|
+
num_epochs: int = 50
|
|
39
|
+
model_save_path: str = './checkpoints'
|
|
40
|
+
log_dir: str = './runs'
|
|
41
|
+
log_interval: int = 5
|
|
42
|
+
save_interval: int = 1
|
|
43
|
+
warmup_steps: int = 200
|
|
44
|
+
|
|
45
|
+
@dataclass
|
|
46
|
+
class VocosConfig:
|
|
47
|
+
input_channels: int = 128
|
|
48
|
+
dim: int = 512
|
|
49
|
+
intermediate_dim: int = 1536
|
|
50
|
+
num_layers: int = 8
|
|
File without changes
|
|
@@ -0,0 +1,160 @@
|
|
|
1
|
+
import os, random, io, glob, tarfile
|
|
2
|
+
import pyarrow.parquet as pq
|
|
3
|
+
import soundfile as sf
|
|
4
|
+
import torch
|
|
5
|
+
import torchaudio
|
|
6
|
+
from torch.utils.data import IterableDataset, get_worker_info
|
|
7
|
+
|
|
8
|
+
from ..config import MelConfig
|
|
9
|
+
from ..text import cleaned_text_to_sequence
|
|
10
|
+
from ..text.burmese import burmese_to_ipa2
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def _intersperse(lst, item):
|
|
14
|
+
r = [item] * (len(lst) * 2 + 1)
|
|
15
|
+
r[1::2] = lst
|
|
16
|
+
return r
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
SCRATCH = "/mnt/local-scratch"
|
|
20
|
+
|
|
21
|
+
SOURCES = [
|
|
22
|
+
("parquet", "/content/shards/data/run*/shard-*.parquet", "clean_text", 1.0),
|
|
23
|
+
("parquet", f"{SCRATCH}/data_all/95k/data/run*/shard-*.parquet", "original_text", 1.0),
|
|
24
|
+
("parquet", f"{SCRATCH}/data_all/bible/data/train-*.parquet", "transcript_char",1.0),
|
|
25
|
+
("parquet", f"{SCRATCH}/data_all/6k/data/shard-*.parquet", "original_text", 1.0),
|
|
26
|
+
("parquet", f"{SCRATCH}/data_all/qtext/data/train-*.parquet", "original_field", 0.5),
|
|
27
|
+
("parquet", f"{SCRATCH}/data_all/synth_mix/data/train-*.parquet", "meta_words", 0.5),
|
|
28
|
+
("parquet", f"{SCRATCH}/data_all/synth_6pr/data/train-*.parquet", "meta_words", 0.5),
|
|
29
|
+
("parquet", f"{SCRATCH}/data_all/synth_6p/data/train-*.parquet", "meta_words", 0.5),
|
|
30
|
+
("parquet", f"{SCRATCH}/data_all/synth_vox/data/train-*.parquet", "myanmar_text", 0.5),
|
|
31
|
+
("tar", f"{SCRATCH}/data_all/1m/audio_tars/*.tar", "csv", 1.0),
|
|
32
|
+
]
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
class AllSourcesDataset(IterableDataset):
|
|
36
|
+
def __init__(self, sources=SOURCES, target_sr=None, min_sec=0.5, max_sec=11.0):
|
|
37
|
+
self.target_sr = target_sr or MelConfig().sample_rate
|
|
38
|
+
self.min_sec = min_sec
|
|
39
|
+
self.max_sec = max_sec
|
|
40
|
+
|
|
41
|
+
self.entries = []
|
|
42
|
+
for kind, pattern, field, weight in sources:
|
|
43
|
+
paths = sorted(glob.glob(pattern))
|
|
44
|
+
for p in paths:
|
|
45
|
+
self.entries.append((kind, p, field, weight))
|
|
46
|
+
print(f"[AllSources] {len(self.entries)} shards across {len(sources)} sources | target_sr={self.target_sr}")
|
|
47
|
+
|
|
48
|
+
self.tar_meta = {}
|
|
49
|
+
csv_path = f"{SCRATCH}/data_all/1m/train_1m_combined_24k.csv"
|
|
50
|
+
if os.path.exists(csv_path):
|
|
51
|
+
print("[AllSources] parsing CSV...")
|
|
52
|
+
with open(csv_path, "r", encoding="utf-8") as f:
|
|
53
|
+
next(f, None)
|
|
54
|
+
for line in f:
|
|
55
|
+
parts = line.rstrip("\n").split("|", 1)
|
|
56
|
+
if len(parts) == 2:
|
|
57
|
+
fname = os.path.basename(parts[0])
|
|
58
|
+
self.tar_meta[fname] = parts[1].strip()
|
|
59
|
+
print(f"[AllSources] indexed {len(self.tar_meta)} CSV entries")
|
|
60
|
+
|
|
61
|
+
def _get_text(self, t):
|
|
62
|
+
if t is None:
|
|
63
|
+
return ""
|
|
64
|
+
if isinstance(t, list):
|
|
65
|
+
return " ".join(str(x) for x in t if x).strip()
|
|
66
|
+
return str(t).strip()
|
|
67
|
+
|
|
68
|
+
def _to_wav_phone(self, wav_bytes, text, given_sr=None):
|
|
69
|
+
try:
|
|
70
|
+
if not text:
|
|
71
|
+
return None
|
|
72
|
+
data, sr = sf.read(io.BytesIO(wav_bytes), dtype="float32")
|
|
73
|
+
if data.size == 0:
|
|
74
|
+
return None
|
|
75
|
+
wav = torch.from_numpy(data)
|
|
76
|
+
if wav.dim() > 1:
|
|
77
|
+
wav = wav.mean(dim=-1)
|
|
78
|
+
if wav.abs().max().item() < 0.001:
|
|
79
|
+
return None
|
|
80
|
+
|
|
81
|
+
actual_sr = given_sr or sr
|
|
82
|
+
if actual_sr != self.target_sr:
|
|
83
|
+
wav = torchaudio.functional.resample(wav, actual_sr, self.target_sr)
|
|
84
|
+
|
|
85
|
+
dur = wav.numel() / self.target_sr
|
|
86
|
+
if dur < self.min_sec or dur > self.max_sec:
|
|
87
|
+
return None
|
|
88
|
+
|
|
89
|
+
phones = burmese_to_ipa2(text)
|
|
90
|
+
ids = cleaned_text_to_sequence(phones)
|
|
91
|
+
if len(ids) < 3:
|
|
92
|
+
return None
|
|
93
|
+
phone = torch.tensor(_intersperse(ids, 0), dtype=torch.long)
|
|
94
|
+
return wav, phone
|
|
95
|
+
except Exception:
|
|
96
|
+
return None
|
|
97
|
+
|
|
98
|
+
def _iter_parquet(self, path, field):
|
|
99
|
+
try:
|
|
100
|
+
pf = pq.ParquetFile(path)
|
|
101
|
+
for batch in pf.iter_batches(batch_size=64):
|
|
102
|
+
d = batch.to_pydict()
|
|
103
|
+
for i in range(len(batch)):
|
|
104
|
+
text = self._get_text(d[field][i])
|
|
105
|
+
if not text:
|
|
106
|
+
continue
|
|
107
|
+
a = d["audio"][i]
|
|
108
|
+
if not isinstance(a, dict):
|
|
109
|
+
continue
|
|
110
|
+
if a.get("bytes") is not None:
|
|
111
|
+
out = self._to_wav_phone(a["bytes"], text)
|
|
112
|
+
elif a.get("array") is not None:
|
|
113
|
+
arr = a["array"]
|
|
114
|
+
sr = a["sampling_rate"]
|
|
115
|
+
bio = io.BytesIO()
|
|
116
|
+
sf.write(bio, arr, sr, format="WAV")
|
|
117
|
+
out = self._to_wav_phone(bio.getvalue(), text, given_sr=sr)
|
|
118
|
+
else:
|
|
119
|
+
continue
|
|
120
|
+
if out is not None:
|
|
121
|
+
yield out
|
|
122
|
+
except Exception as e:
|
|
123
|
+
print(f"[skip parquet] {path}: {e}")
|
|
124
|
+
return
|
|
125
|
+
|
|
126
|
+
def _iter_tar(self, path):
|
|
127
|
+
try:
|
|
128
|
+
with tarfile.open(path, "r:") as tf: # seekable mode (not stream)
|
|
129
|
+
for member in tf:
|
|
130
|
+
if not member.name.endswith(".wav"):
|
|
131
|
+
continue
|
|
132
|
+
fname = os.path.basename(member.name)
|
|
133
|
+
text = self.tar_meta.get(fname)
|
|
134
|
+
if not text:
|
|
135
|
+
continue
|
|
136
|
+
f = tf.extractfile(member)
|
|
137
|
+
if f is None:
|
|
138
|
+
continue
|
|
139
|
+
out = self._to_wav_phone(f.read(), text)
|
|
140
|
+
if out is not None:
|
|
141
|
+
yield out
|
|
142
|
+
except Exception as e:
|
|
143
|
+
print(f"[skip tar] {path}: {e}")
|
|
144
|
+
return
|
|
145
|
+
|
|
146
|
+
def __iter__(self):
|
|
147
|
+
info = get_worker_info()
|
|
148
|
+
wid = info.id if info else 0
|
|
149
|
+
nw = info.num_workers if info else 1
|
|
150
|
+
my = self.entries[wid::nw]
|
|
151
|
+
rng = random.Random(42 + wid)
|
|
152
|
+
while True:
|
|
153
|
+
rng.shuffle(my)
|
|
154
|
+
for kind, path, field, weight in my:
|
|
155
|
+
if weight < 1.0 and rng.random() > weight:
|
|
156
|
+
continue
|
|
157
|
+
if kind == "parquet":
|
|
158
|
+
yield from self._iter_parquet(path, field)
|
|
159
|
+
else:
|
|
160
|
+
yield from self._iter_tar(path)
|
|
@@ -0,0 +1,39 @@
|
|
|
1
|
+
import random
|
|
2
|
+
import torch
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
def _random_slice(x, mn=12, mx=3):
|
|
6
|
+
L = x.size(-1)
|
|
7
|
+
if L < 12:
|
|
8
|
+
return x
|
|
9
|
+
seg = random.randint(max(1, L // mn), max(1, L // mx))
|
|
10
|
+
seg = min(seg, L)
|
|
11
|
+
start = random.randint(0, L - seg)
|
|
12
|
+
return x[start:start + seg]
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def collate_fn_wav(batch):
|
|
16
|
+
wavs, phones = zip(*batch)
|
|
17
|
+
wav_lens = torch.tensor([w.size(-1) for w in wavs], dtype=torch.long)
|
|
18
|
+
phone_lens = torch.tensor([p.size(-1) for p in phones], dtype=torch.long)
|
|
19
|
+
|
|
20
|
+
sliced = [_random_slice(w) for w in wavs]
|
|
21
|
+
sliced_lens = torch.tensor([s.size(-1) for s in sliced], dtype=torch.long)
|
|
22
|
+
|
|
23
|
+
B = len(wavs)
|
|
24
|
+
max_w = int(wav_lens.max())
|
|
25
|
+
max_p = int(phone_lens.max())
|
|
26
|
+
max_s = int(sliced_lens.max())
|
|
27
|
+
|
|
28
|
+
wav_pad = torch.zeros(B, max_w)
|
|
29
|
+
phone_pad = torch.zeros(B, max_p, dtype=torch.long)
|
|
30
|
+
sliced_pad = torch.zeros(B, max_s)
|
|
31
|
+
|
|
32
|
+
for i, w in enumerate(wavs):
|
|
33
|
+
wav_pad[i, :w.size(-1)] = w
|
|
34
|
+
for i, p in enumerate(phones):
|
|
35
|
+
phone_pad[i, :p.size(-1)] = p
|
|
36
|
+
for i, s in enumerate(sliced):
|
|
37
|
+
sliced_pad[i, :s.size(-1)] = s
|
|
38
|
+
|
|
39
|
+
return (phone_pad, phone_lens, wav_pad, wav_lens, sliced_pad, sliced_lens)
|
|
@@ -0,0 +1,69 @@
|
|
|
1
|
+
import os
|
|
2
|
+
import random
|
|
3
|
+
|
|
4
|
+
import json
|
|
5
|
+
import torch
|
|
6
|
+
from torch.utils.data import Dataset
|
|
7
|
+
|
|
8
|
+
from ..text import cleaned_text_to_sequence
|
|
9
|
+
|
|
10
|
+
def intersperse(lst: list, item: int):
|
|
11
|
+
"""
|
|
12
|
+
putting a blank token between any two input tokens to improve pronunciation
|
|
13
|
+
see https://github.com/jaywalnut310/glow-tts/issues/43 for more details
|
|
14
|
+
"""
|
|
15
|
+
result = [item] * (len(lst) * 2 + 1)
|
|
16
|
+
result[1::2] = lst
|
|
17
|
+
return result
|
|
18
|
+
|
|
19
|
+
class StableDataset(Dataset):
|
|
20
|
+
def __init__(self, filelist_path, hop_length):
|
|
21
|
+
self.filelist_path = filelist_path
|
|
22
|
+
self.hop_length = hop_length
|
|
23
|
+
|
|
24
|
+
self._load_filelist(filelist_path)
|
|
25
|
+
|
|
26
|
+
def _load_filelist(self, filelist_path):
|
|
27
|
+
filelist, lengths = [], []
|
|
28
|
+
with open(filelist_path, 'r', encoding='utf-8') as f:
|
|
29
|
+
for line in f:
|
|
30
|
+
line = json.loads(line.strip())
|
|
31
|
+
filelist.append((line['mel_path'], line['phone']))
|
|
32
|
+
lengths.append(line['mel_length'])
|
|
33
|
+
|
|
34
|
+
self.filelist = filelist
|
|
35
|
+
self.lengths = lengths # length is used for DistributedBucketSampler
|
|
36
|
+
|
|
37
|
+
def __len__(self):
|
|
38
|
+
return len(self.filelist)
|
|
39
|
+
|
|
40
|
+
def __getitem__(self, idx):
|
|
41
|
+
mel_path, phone = self.filelist[idx]
|
|
42
|
+
mel = torch.load(mel_path, map_location='cpu', weights_only=True)
|
|
43
|
+
phone = torch.tensor(intersperse(cleaned_text_to_sequence(phone), 0), dtype=torch.long)
|
|
44
|
+
return mel, phone
|
|
45
|
+
|
|
46
|
+
def collate_fn(batch):
|
|
47
|
+
texts = [item[1] for item in batch]
|
|
48
|
+
mels = [item[0] for item in batch]
|
|
49
|
+
mels_sliced = [random_slice_tensor(mel) for mel in mels]
|
|
50
|
+
|
|
51
|
+
text_lengths = torch.tensor([text.size(-1) for text in texts], dtype=torch.long)
|
|
52
|
+
mel_lengths = torch.tensor([mel.size(-1) for mel in mels], dtype=torch.long)
|
|
53
|
+
mels_sliced_lengths = torch.tensor([mel_sliced.size(-1) for mel_sliced in mels_sliced], dtype=torch.long)
|
|
54
|
+
|
|
55
|
+
# pad to the same length
|
|
56
|
+
texts_padded = torch.nested.to_padded_tensor(torch.nested.nested_tensor(texts), padding=0)
|
|
57
|
+
mels_padded = torch.nested.to_padded_tensor(torch.nested.nested_tensor(mels), padding=0)
|
|
58
|
+
mels_sliced_padded = torch.nested.to_padded_tensor(torch.nested.nested_tensor(mels_sliced), padding=0)
|
|
59
|
+
|
|
60
|
+
return texts_padded, text_lengths, mels_padded, mel_lengths, mels_sliced_padded, mels_sliced_lengths
|
|
61
|
+
|
|
62
|
+
# random slice mel for reference encoder to prevent overfitting
|
|
63
|
+
def random_slice_tensor(x: torch.Tensor):
|
|
64
|
+
length = x.size(-1)
|
|
65
|
+
if length < 12:
|
|
66
|
+
return x
|
|
67
|
+
segmnt_size = random.randint(length // 12, length // 3)
|
|
68
|
+
start = random.randint(0, length - segmnt_size)
|
|
69
|
+
return x[..., start : start + segmnt_size]
|
|
@@ -0,0 +1,73 @@
|
|
|
1
|
+
import os, random, io
|
|
2
|
+
import soundfile as sf
|
|
3
|
+
import pyarrow.parquet as pq
|
|
4
|
+
import torch
|
|
5
|
+
import torchaudio
|
|
6
|
+
from torch.utils.data import IterableDataset, get_worker_info
|
|
7
|
+
from ..text import cleaned_text_to_sequence
|
|
8
|
+
from ..text.burmese import burmese_to_ipa2
|
|
9
|
+
|
|
10
|
+
def _intersperse(lst, item):
|
|
11
|
+
result = [item] * (len(lst) * 2 + 1)
|
|
12
|
+
result[1::2] = lst
|
|
13
|
+
return result
|
|
14
|
+
|
|
15
|
+
class LocalShardsDataset(IterableDataset):
|
|
16
|
+
def __init__(self, shard_paths, target_sr=44100, min_sec=0.6, max_sec=18.0):
|
|
17
|
+
self.shard_paths = list(shard_paths)
|
|
18
|
+
self.target_sr = target_sr
|
|
19
|
+
self.min_sec = min_sec
|
|
20
|
+
self.max_sec = max_sec
|
|
21
|
+
|
|
22
|
+
def _process_sample(self, row):
|
|
23
|
+
try:
|
|
24
|
+
text = (row.get("clean_text") or row.get("original_text") or "").strip()
|
|
25
|
+
if not text:
|
|
26
|
+
return None
|
|
27
|
+
a = row["audio"]
|
|
28
|
+
if isinstance(a, dict) and a.get("bytes") is not None:
|
|
29
|
+
data, sr = sf.read(io.BytesIO(a["bytes"]), dtype="float32")
|
|
30
|
+
elif isinstance(a, dict) and a.get("array") is not None:
|
|
31
|
+
import numpy as np
|
|
32
|
+
data = np.asarray(a["array"], dtype="float32")
|
|
33
|
+
sr = a["sampling_rate"]
|
|
34
|
+
else:
|
|
35
|
+
return None
|
|
36
|
+
wav = torch.from_numpy(data)
|
|
37
|
+
if wav.dim() > 1:
|
|
38
|
+
wav = wav.mean(dim=-1)
|
|
39
|
+
if sr != self.target_sr:
|
|
40
|
+
wav = torchaudio.functional.resample(wav, sr, self.target_sr)
|
|
41
|
+
dur = wav.numel() / self.target_sr
|
|
42
|
+
if dur < self.min_sec or dur > self.max_sec:
|
|
43
|
+
return None
|
|
44
|
+
phones = burmese_to_ipa2(text)
|
|
45
|
+
ids = cleaned_text_to_sequence(phones)
|
|
46
|
+
if len(ids) < 3:
|
|
47
|
+
return None
|
|
48
|
+
phone = torch.tensor(_intersperse(ids, 0), dtype=torch.long)
|
|
49
|
+
return wav, phone
|
|
50
|
+
except Exception:
|
|
51
|
+
return None
|
|
52
|
+
|
|
53
|
+
def __iter__(self):
|
|
54
|
+
info = get_worker_info()
|
|
55
|
+
wid = info.id if info else 0
|
|
56
|
+
nw = info.num_workers if info else 1
|
|
57
|
+
my = self.shard_paths[wid::nw]
|
|
58
|
+
rng = random.Random(42 + wid)
|
|
59
|
+
while True:
|
|
60
|
+
rng.shuffle(my)
|
|
61
|
+
for path in my:
|
|
62
|
+
try:
|
|
63
|
+
pf = pq.ParquetFile(path)
|
|
64
|
+
for batch in pf.iter_batches(batch_size=64):
|
|
65
|
+
d = batch.to_pydict()
|
|
66
|
+
for i in range(len(batch)):
|
|
67
|
+
row = {k: d[k][i] for k in d}
|
|
68
|
+
out = self._process_sample(row)
|
|
69
|
+
if out is not None:
|
|
70
|
+
yield out
|
|
71
|
+
except Exception as e:
|
|
72
|
+
print(f"[skip] {path}: {e}")
|
|
73
|
+
continue
|
|
File without changes
|