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.
Files changed (59) hide show
  1. myanmar_tts/__init__.py +42 -0
  2. myanmar_tts/api.py +82 -0
  3. myanmar_tts/burmese.py +42 -0
  4. myanmar_tts/config.py +50 -0
  5. myanmar_tts/datas/__init__.py +0 -0
  6. myanmar_tts/datas/all_sources_dataset.py +160 -0
  7. myanmar_tts/datas/collate_wav.py +39 -0
  8. myanmar_tts/datas/dataset.py +69 -0
  9. myanmar_tts/datas/local_shards_dataset.py +73 -0
  10. myanmar_tts/models/__init__.py +0 -0
  11. myanmar_tts/models/diffusion_transformer.py +205 -0
  12. myanmar_tts/models/duration_predictor.py +40 -0
  13. myanmar_tts/models/estimator.py +138 -0
  14. myanmar_tts/models/flow_matching.py +100 -0
  15. myanmar_tts/models/model.py +178 -0
  16. myanmar_tts/models/reference_encoder.py +168 -0
  17. myanmar_tts/models/text_encoder.py +44 -0
  18. myanmar_tts/monotonic_align/__init__.py +16 -0
  19. myanmar_tts/monotonic_align/core.py +46 -0
  20. myanmar_tts/symbols.py +57 -0
  21. myanmar_tts/text/LICENSE +19 -0
  22. myanmar_tts/text/__init__.py +16 -0
  23. myanmar_tts/text/burmese.py +42 -0
  24. myanmar_tts/text/cleaners.py +10 -0
  25. myanmar_tts/text/symbols.py +57 -0
  26. myanmar_tts/utils/__init__.py +0 -0
  27. myanmar_tts/utils/audio.py +74 -0
  28. myanmar_tts/utils/load.py +43 -0
  29. myanmar_tts/utils/mask.py +8 -0
  30. myanmar_tts/utils/scheduler.py +428 -0
  31. myanmar_tts/vocoders/__init__.py +0 -0
  32. myanmar_tts/vocoders/ffgan/__init__.py +0 -0
  33. myanmar_tts/vocoders/ffgan/backbone.py +214 -0
  34. myanmar_tts/vocoders/ffgan/head.py +257 -0
  35. myanmar_tts/vocoders/ffgan/model.py +57 -0
  36. myanmar_tts/vocoders/ffgan/unify.py +60 -0
  37. myanmar_tts/vocoders/vocos/README.md +41 -0
  38. myanmar_tts/vocoders/vocos/__init__.py +0 -0
  39. myanmar_tts/vocoders/vocos/config.py +41 -0
  40. myanmar_tts/vocoders/vocos/dataset.py +57 -0
  41. myanmar_tts/vocoders/vocos/inference.ipynb +79 -0
  42. myanmar_tts/vocoders/vocos/models/__init__.py +0 -0
  43. myanmar_tts/vocoders/vocos/models/backbone.py +57 -0
  44. myanmar_tts/vocoders/vocos/models/discriminator.py +171 -0
  45. myanmar_tts/vocoders/vocos/models/head.py +118 -0
  46. myanmar_tts/vocoders/vocos/models/loss.py +66 -0
  47. myanmar_tts/vocoders/vocos/models/model.py +20 -0
  48. myanmar_tts/vocoders/vocos/models/module.py +47 -0
  49. myanmar_tts/vocoders/vocos/preprocess.py +45 -0
  50. myanmar_tts/vocoders/vocos/requirements.txt +2 -0
  51. myanmar_tts/vocoders/vocos/train.py +165 -0
  52. myanmar_tts/vocoders/vocos/utils/__init__.py +0 -0
  53. myanmar_tts/vocoders/vocos/utils/audio.py +74 -0
  54. myanmar_tts/vocoders/vocos/utils/load.py +53 -0
  55. myanmar_tts/vocoders/vocos/utils/scheduler.py +298 -0
  56. myanmartts-1.0.0.dist-info/METADATA +173 -0
  57. myanmartts-1.0.0.dist-info/RECORD +59 -0
  58. myanmartts-1.0.0.dist-info/WHEEL +4 -0
  59. myanmartts-1.0.0.dist-info/licenses/LICENSE +21 -0
@@ -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