crazydev919 commited on
Commit
c35d6ec
·
verified ·
1 Parent(s): 520d894

Upload app.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. app.py +122 -229
app.py CHANGED
@@ -4,188 +4,132 @@ import torch, yaml, os, sys, glob, re
4
  import librosa
5
  import soundfile as sf
6
  import torchaudio
7
- import numpy as np
8
  from huggingface_hub import snapshot_download
9
  from munch import Munch
10
  from nltk.tokenize import word_tokenize
11
  import nltk
12
  nltk.download("punkt_tab", quiet=True)
13
 
14
- # ══════════════════════════════════════════════════════════
15
- # KIKUYU load at startup (small, 145MB)
16
- # ══════════════════════════════════════════════════════════
17
- from transformers import VitsModel, AutoTokenizer
18
-
19
- print("Loading Kikuyu TTS...")
20
- kikuyu_model = VitsModel.from_pretrained("gateremark/kikuyu-tts-v1")
21
- kikuyu_tokenizer = AutoTokenizer.from_pretrained("gateremark/kikuyu-tts-v1")
22
- kikuyu_model.eval()
23
- print("✅ Kikuyu model loaded")
24
-
25
- def synthesize_kikuyu(text):
26
- inputs = kikuyu_tokenizer(text=text.strip(), return_tensors="pt")
27
- with torch.no_grad():
28
- output = kikuyu_model(**inputs)
29
- waveform = output.waveform.squeeze().cpu().numpy()
30
- sr = kikuyu_model.config.sampling_rate
31
- sf.write("/tmp/kikuyu_output.wav", waveform, sr)
32
- return "/tmp/kikuyu_output.wav"
33
-
34
- # ══════════════════════════════════════════════════════════
35
- # LUHYA — lazy load on first request (heavy, 2.18GB)
36
- # ══════════════════════════════════════════════════════════
37
- _luhya_loaded = False
38
- _luhya_model = None
39
- _luhya_sampler = None
40
- _default_style = None
41
- _model_params = None
42
- _textcleaner = None
43
- _model_dir = None
44
-
45
- def load_luhya():
46
- global _luhya_loaded, _luhya_model, _luhya_sampler
47
- global _default_style, _model_params, _textcleaner, _model_dir
48
-
49
- if _luhya_loaded:
50
- return
51
-
52
- print("Loading Luhya TTS (first request)...")
53
-
54
- # Clone + patch StyleTTS2
55
- os.system("git clone https://github.com/yl4579/StyleTTS2 /app/StyleTTS2 2>/dev/null || true")
56
- for fpath in ["/app/StyleTTS2/models.py", "/app/StyleTTS2/utils.py"]:
57
- with open(fpath) as f:
58
- code = f.read()
59
- patched = re.sub(
60
- r'torch\.load\(([^)]+)\)',
61
- lambda m: m.group(0) if 'weights_only' in m.group(1)
62
- else f'torch.load({m.group(1)}, weights_only=False)',
63
- code
64
- )
65
- with open(fpath, "w") as f:
66
- f.write(patched)
67
-
68
- if "/app/StyleTTS2" not in sys.path:
69
- sys.path.insert(0, "/app/StyleTTS2")
70
- os.chdir("/app/StyleTTS2")
71
-
72
- _model_dir = snapshot_download("crazydev919/luhya-tts")
73
-
74
- from models import build_model, load_ASR_models, load_F0_models
75
- from utils import recursive_munch
76
- from text_utils import TextCleaner
77
- from Modules.diffusion.sampler import DiffusionSampler, ADPM2Sampler, KarrasSchedule
78
-
79
- _textcleaner = TextCleaner()
80
- device = "cpu"
81
-
82
- config = yaml.safe_load(open(f"{_model_dir}/config.yml"))
83
- config["ASR_path"] = f"{_model_dir}/Utils/ASR/epoch_00080.pth"
84
- config["ASR_config"] = f"{_model_dir}/Utils/ASR/config.yml"
85
- config["F0_path"] = f"{_model_dir}/Utils/JDC/bst.t7"
86
- config["PLBERT_dir"] = f"{_model_dir}/Utils/PLBERT/"
87
-
88
- text_aligner = load_ASR_models(config["ASR_path"], config["ASR_config"])
89
- pitch_extractor = load_F0_models(config["F0_path"])
90
-
91
- from Utils.PLBERT.util import load_plbert
92
- plbert = load_plbert(config["PLBERT_dir"])
93
-
94
- _model_params = recursive_munch(config["model_params"])
95
- _luhya_model = build_model(_model_params, text_aligner, pitch_extractor, plbert)
96
- _ = [_luhya_model[k].eval() for k in _luhya_model]
97
- _ = [_luhya_model[k].to(device) for k in _luhya_model]
98
-
99
- params = torch.load(
100
- f"{_model_dir}/model.pth", map_location="cpu", weights_only=False
101
- )["net"]
102
- for key in _luhya_model:
103
- if key in params:
104
- try:
105
- _luhya_model[key].load_state_dict(params[key])
106
- except:
107
- from collections import OrderedDict
108
- sd = OrderedDict()
109
- for k, v in params[key].items():
110
- sd[k[7:] if k.startswith("module.") else k] = v
111
- _luhya_model[key].load_state_dict(sd, strict=False)
112
- _ = [_luhya_model[k].eval() for k in _luhya_model]
113
-
114
- _luhya_sampler = DiffusionSampler(
115
- _luhya_model.diffusion.diffusion,
116
- sampler=ADPM2Sampler(),
117
- sigma_schedule=KarrasSchedule(sigma_min=0.0001, sigma_max=3.0, rho=9.0),
118
- clamp=False
119
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
120
 
121
- # Precompute default style
122
- ref_candidates = (
123
- glob.glob(f"{_model_dir}/ref_wavs/*.wav") +
124
- glob.glob(f"{_model_dir}/*.wav")
125
- )
126
- DEFAULT_REF = sorted(ref_candidates)[0]
127
 
128
- to_mel = torchaudio.transforms.MelSpectrogram(
129
- n_mels=80, n_fft=2048, win_length=1200, hop_length=300)
130
- mean, std = -4, 4
 
 
131
 
132
- wave, sr = librosa.load(DEFAULT_REF, sr=24000)
133
- audio, _ = librosa.effects.trim(wave, top_db=30)
134
- wave_t = torch.from_numpy(audio).float()
135
- mel = (torch.log(1e-5 + to_mel(wave_t).unsqueeze(0)) - mean) / std
136
- mel = mel.to(device)
137
 
 
 
 
 
138
  with torch.no_grad():
139
- ref_s = _luhya_model.style_encoder(mel.unsqueeze(1))
140
- ref_p = _luhya_model.predictor_encoder(mel.unsqueeze(1))
141
- _default_style = torch.cat([ref_s, ref_p], dim=1)
142
-
143
- _luhya_loaded = True
144
- print("✅ Luhya model loaded and cached")
145
 
146
- def synthesize_luhya(text, ref_audio=None, alpha=0.3, beta=0.7, steps=5):
147
- load_luhya() # lazy load
 
 
 
 
 
148
 
 
149
  import phonemizer
150
- device = "cpu"
151
-
152
- to_mel = torchaudio.transforms.MelSpectrogram(
153
- n_mels=80, n_fft=2048, win_length=1200, hop_length=300)
154
- mean, std = -4, 4
155
-
156
- def length_to_mask(lengths):
157
- mask = torch.arange(lengths.max()).unsqueeze(0).expand(
158
- lengths.shape[0], -1).type_as(lengths)
159
- mask = torch.gt(mask + 1, lengths.unsqueeze(1))
160
- return mask
161
-
162
- def compute_style(path):
163
- wave, sr = librosa.load(path, sr=24000)
164
- audio, _ = librosa.effects.trim(wave, top_db=30)
165
- wave_t = torch.from_numpy(audio).float()
166
- mel = (torch.log(1e-5 + to_mel(wave_t).unsqueeze(0)) - mean) / std
167
- with torch.no_grad():
168
- ref_s = _luhya_model.style_encoder(mel.unsqueeze(1))
169
- ref_p = _luhya_model.predictor_encoder(mel.unsqueeze(1))
170
- return torch.cat([ref_s, ref_p], dim=1)
171
-
172
  pb = phonemizer.backend.EspeakBackend(
173
  language="sw", preserve_punctuation=True, with_stress=True)
174
- ref_s = compute_style(ref_audio) if ref_audio else _default_style
175
 
176
  ps = " ".join(word_tokenize(pb.phonemize([text.strip()])[0]))
177
- tokens = _textcleaner(ps)
178
  tokens.insert(0, 0)
179
  tokens = torch.LongTensor(tokens).to(device).unsqueeze(0)
180
 
181
  with torch.no_grad():
182
  il = torch.LongTensor([tokens.shape[-1]]).to(device)
183
  tm = length_to_mask(il).to(device)
184
- t_en = _luhya_model.text_encoder(tokens, il, tm)
185
- bd = _luhya_model.bert(tokens, attention_mask=(~tm).int())
186
- d_en = _luhya_model.bert_encoder(bd).transpose(-1, -2)
187
 
188
- sp = _luhya_sampler(
189
  noise=torch.randn((1, 256)).unsqueeze(1).to(device),
190
  embedding=bd, embedding_scale=1,
191
  features=ref_s, num_steps=int(steps)
@@ -194,9 +138,9 @@ def synthesize_luhya(text, ref_audio=None, alpha=0.3, beta=0.7, steps=5):
194
  s = beta * sp[:, 128:] + (1 - beta) * ref_s[:, 128:]
195
  ref = alpha * sp[:, :128] + (1 - alpha) * ref_s[:, :128]
196
 
197
- d = _luhya_model.predictor.text_encoder(d_en, s, il, tm)
198
- x, _ = _luhya_model.predictor.lstm(d)
199
- dur = torch.sigmoid(_luhya_model.predictor.duration_proj(x)).sum(axis=-1)
200
  pd = torch.round(dur.squeeze()).clamp(min=1)
201
 
202
  at = torch.zeros(il, int(pd.sum().data))
@@ -208,7 +152,7 @@ def synthesize_luhya(text, ref_audio=None, alpha=0.3, beta=0.7, steps=5):
208
  en = d.transpose(-1, -2) @ at.unsqueeze(0).to(device)
209
  asr = t_en @ at.unsqueeze(0).to(device)
210
 
211
- if _model_params.decoder.type == "hifigan":
212
  en_new = torch.zeros_like(en)
213
  asr_new = torch.zeros_like(asr)
214
  en_new[:, :, 0] = en[:, :, 0]
@@ -217,75 +161,24 @@ def synthesize_luhya(text, ref_audio=None, alpha=0.3, beta=0.7, steps=5):
217
  asr_new[:, :, 1:] = asr[:, :, :-1]
218
  en, asr = en_new, asr_new
219
 
220
- F0, N = _luhya_model.predictor.F0Ntrain(en, s)
221
- out = _luhya_model.decoder(asr, F0, N, ref.squeeze().unsqueeze(0))
222
 
223
  wav = out.squeeze().cpu().numpy()[..., :-50]
224
- sf.write("/tmp/luhya_output.wav", wav, 24000)
225
- return "/tmp/luhya_output.wav"
226
-
227
- # ══════════════════════════════════════════════════════════
228
- # UNIFIED ENTRY POINT
229
- # ══════════════════════════════════════════════════════════
230
- def synthesize(language, text, ref_audio=None, alpha=0.3, beta=0.7, steps=5):
231
- if language == "Kikuyu":
232
- return synthesize_kikuyu(text)
233
- else:
234
- return synthesize_luhya(text, ref_audio, alpha, beta, steps)
235
-
236
- # ══════════════════════════════════════════════════════════
237
- # GRADIO UI
238
- # ══════════════════════════════════════════════════════════
239
- with gr.Blocks(title="Kenyan Languages TTS") as demo:
240
- gr.Markdown("# 🗣️ Kenyan Languages TTS\nLuhya (Lunyore) and Kikuyu text-to-speech.")
241
-
242
- language = gr.Radio(
243
- choices=["Kikuyu", "Luhya"],
244
- value="Kikuyu",
245
- label="Language"
246
- )
247
- text = gr.Textbox(
248
- label="Text",
249
- value="Mũtũũrĩre wa ndũire nĩ kĩheo kĩa mwanya.",
250
- lines=4
251
- )
252
-
253
- with gr.Group(visible=False) as luhya_controls:
254
- gr.Markdown("**Luhya voice controls**")
255
- ref_audio = gr.Audio(label="Reference Voice (optional)", type="filepath")
256
- with gr.Row():
257
- alpha = gr.Slider(0.0, 1.0, value=0.3, step=0.1, label="Alpha")
258
- beta = gr.Slider(0.0, 1.0, value=0.7, step=0.1, label="Beta")
259
- steps = gr.Slider(1, 10, value=5, step=1, label="Steps")
260
-
261
- output_audio = gr.Audio(label="Generated Speech", type="filepath")
262
- btn = gr.Button("Generate", variant="primary")
263
-
264
- def toggle(lang):
265
- texts = {
266
- "Kikuyu": "Mũtũũrĩre wa ndũire nĩ kĩheo kĩa mwanya.",
267
- "Luhya" : "mirembe. obulani lwa bwana nyasaye.",
268
- }
269
- return gr.update(visible=(lang=="Luhya")), gr.update(value=texts[lang])
270
-
271
- language.change(fn=toggle, inputs=language, outputs=[luhya_controls, text])
272
-
273
- btn.click(
274
- fn = synthesize,
275
- inputs = [language, text, ref_audio, alpha, beta, steps],
276
- outputs = output_audio
277
- )
278
-
279
- gr.Markdown("""
280
- **API:**
281
- POST /api/predict
282
- {"data": ["Kikuyu", "your kikuyu text", null, 0.3, 0.7, 5]}
283
- {"data": ["Luhya", "your luhya text", null, 0.3, 0.7, 5]}
284
- Note: First Luhya request takes ~30s to load the model.
285
- """)
286
-
287
- demo.queue(max_size=3).launch(
288
- prevent_thread_lock = True,
289
- max_threads = 1,
290
- show_error = True,
291
  )
 
 
 
4
  import librosa
5
  import soundfile as sf
6
  import torchaudio
 
7
  from huggingface_hub import snapshot_download
8
  from munch import Munch
9
  from nltk.tokenize import word_tokenize
10
  import nltk
11
  nltk.download("punkt_tab", quiet=True)
12
 
13
+ # ── Clone + patch StyleTTS2 ────────────────────────────────
14
+ os.system("git clone https://github.com/yl4579/StyleTTS2 /app/StyleTTS2 2>/dev/null || true")
15
+
16
+ for fpath in ["/app/StyleTTS2/models.py", "/app/StyleTTS2/utils.py"]:
17
+ with open(fpath) as f:
18
+ code = f.read()
19
+ patched = re.sub(
20
+ r'torch\.load\(([^)]+)\)',
21
+ lambda m: m.group(0) if 'weights_only' in m.group(1)
22
+ else f'torch.load({m.group(1)}, weights_only=False)',
23
+ code
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
24
  )
25
+ with open(fpath, "w") as f:
26
+ f.write(patched)
27
+ print("Patched torch.load")
28
+
29
+ sys.path.insert(0, "/app/StyleTTS2")
30
+ os.chdir("/app/StyleTTS2")
31
+
32
+ model_dir = snapshot_download("crazydev919/luhya-tts")
33
+
34
+ from models import *
35
+ from utils import *
36
+ from text_utils import TextCleaner
37
+ from Modules.diffusion.sampler import DiffusionSampler, ADPM2Sampler, KarrasSchedule
38
+
39
+ device = "cpu"
40
+ textcleaner = TextCleaner()
41
+
42
+ config = yaml.safe_load(open(f"{model_dir}/config.yml"))
43
+ config["ASR_path"] = f"{model_dir}/Utils/ASR/epoch_00080.pth"
44
+ config["ASR_config"] = f"{model_dir}/Utils/ASR/config.yml"
45
+ config["F0_path"] = f"{model_dir}/Utils/JDC/bst.t7"
46
+ config["PLBERT_dir"] = f"{model_dir}/Utils/PLBERT/"
47
+
48
+ text_aligner = load_ASR_models(config["ASR_path"], config["ASR_config"])
49
+ pitch_extractor = load_F0_models(config["F0_path"])
50
+ from Utils.PLBERT.util import load_plbert
51
+ plbert = load_plbert(config["PLBERT_dir"])
52
+
53
+ model_params = recursive_munch(config["model_params"])
54
+ model = build_model(model_params, text_aligner, pitch_extractor, plbert)
55
+ _ = [model[key].eval() for key in model]
56
+ _ = [model[key].to(device) for key in model]
57
+
58
+ params = torch.load(
59
+ f"{model_dir}/model.pth", map_location="cpu", weights_only=False
60
+ )["net"]
61
+ for key in model:
62
+ if key in params:
63
+ try:
64
+ model[key].load_state_dict(params[key])
65
+ except:
66
+ from collections import OrderedDict
67
+ sd = OrderedDict()
68
+ for k, v in params[key].items():
69
+ sd[k[7:] if k.startswith("module.") else k] = v
70
+ model[key].load_state_dict(sd, strict=False)
71
+ _ = [model[key].eval() for key in model]
72
+ print("Model loaded")
73
+
74
+ sampler = DiffusionSampler(
75
+ model.diffusion.diffusion,
76
+ sampler=ADPM2Sampler(),
77
+ sigma_schedule=KarrasSchedule(sigma_min=0.0001, sigma_max=3.0, rho=9.0),
78
+ clamp=False
79
+ )
80
 
81
+ to_mel = torchaudio.transforms.MelSpectrogram(
82
+ n_mels=80, n_fft=2048, win_length=1200, hop_length=300)
83
+ mean, std = -4, 4
 
 
 
84
 
85
+ def length_to_mask(lengths):
86
+ mask = torch.arange(lengths.max()).unsqueeze(0).expand(
87
+ lengths.shape[0], -1).type_as(lengths)
88
+ mask = torch.gt(mask + 1, lengths.unsqueeze(1))
89
+ return mask
90
 
91
+ def preprocess(wave):
92
+ wave_tensor = torch.from_numpy(wave).float()
93
+ mel_tensor = to_mel(wave_tensor)
94
+ mel_tensor = (torch.log(1e-5 + mel_tensor.unsqueeze(0)) - mean) / std
95
+ return mel_tensor
96
 
97
+ def compute_style(path):
98
+ wave, sr = librosa.load(path, sr=24000)
99
+ audio, _ = librosa.effects.trim(wave, top_db=30)
100
+ mel = preprocess(audio).to(device)
101
  with torch.no_grad():
102
+ ref_s = model.style_encoder(mel.unsqueeze(1))
103
+ ref_p = model.predictor_encoder(mel.unsqueeze(1))
104
+ return torch.cat([ref_s, ref_p], dim=1)
 
 
 
105
 
106
+ ref_candidates = (
107
+ glob.glob(f"{model_dir}/ref_wavs/*.wav") +
108
+ glob.glob(f"{model_dir}/*.wav")
109
+ )
110
+ DEFAULT_REF = sorted(ref_candidates)[0]
111
+ DEFAULT_STYLE = compute_style(DEFAULT_REF)
112
+ print(f"Reference: {DEFAULT_REF}")
113
 
114
+ def synthesize(text, ref_audio=None, alpha=0.3, beta=0.7, steps=5):
115
  import phonemizer
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
116
  pb = phonemizer.backend.EspeakBackend(
117
  language="sw", preserve_punctuation=True, with_stress=True)
118
+ ref_s = compute_style(ref_audio) if ref_audio else DEFAULT_STYLE
119
 
120
  ps = " ".join(word_tokenize(pb.phonemize([text.strip()])[0]))
121
+ tokens = textcleaner(ps)
122
  tokens.insert(0, 0)
123
  tokens = torch.LongTensor(tokens).to(device).unsqueeze(0)
124
 
125
  with torch.no_grad():
126
  il = torch.LongTensor([tokens.shape[-1]]).to(device)
127
  tm = length_to_mask(il).to(device)
128
+ t_en = model.text_encoder(tokens, il, tm)
129
+ bd = model.bert(tokens, attention_mask=(~tm).int())
130
+ d_en = model.bert_encoder(bd).transpose(-1, -2)
131
 
132
+ sp = sampler(
133
  noise=torch.randn((1, 256)).unsqueeze(1).to(device),
134
  embedding=bd, embedding_scale=1,
135
  features=ref_s, num_steps=int(steps)
 
138
  s = beta * sp[:, 128:] + (1 - beta) * ref_s[:, 128:]
139
  ref = alpha * sp[:, :128] + (1 - alpha) * ref_s[:, :128]
140
 
141
+ d = model.predictor.text_encoder(d_en, s, il, tm)
142
+ x, _ = model.predictor.lstm(d)
143
+ dur = torch.sigmoid(model.predictor.duration_proj(x)).sum(axis=-1)
144
  pd = torch.round(dur.squeeze()).clamp(min=1)
145
 
146
  at = torch.zeros(il, int(pd.sum().data))
 
152
  en = d.transpose(-1, -2) @ at.unsqueeze(0).to(device)
153
  asr = t_en @ at.unsqueeze(0).to(device)
154
 
155
+ if model_params.decoder.type == "hifigan":
156
  en_new = torch.zeros_like(en)
157
  asr_new = torch.zeros_like(asr)
158
  en_new[:, :, 0] = en[:, :, 0]
 
161
  asr_new[:, :, 1:] = asr[:, :, :-1]
162
  en, asr = en_new, asr_new
163
 
164
+ F0, N = model.predictor.F0Ntrain(en, s)
165
+ out = model.decoder(asr, F0, N, ref.squeeze().unsqueeze(0))
166
 
167
  wav = out.squeeze().cpu().numpy()[..., :-50]
168
+ sf.write("/tmp/output.wav", wav, 24000)
169
+ return "/tmp/output.wav"
170
+
171
+ demo = gr.Interface(
172
+ fn = synthesize,
173
+ inputs = [
174
+ gr.Textbox(label="Luhya text", value="mirembe. obulani lwa bwana nyasaye.", lines=3),
175
+ gr.Audio(label="Reference voice (optional)", type="filepath", value=None),
176
+ gr.Slider(0.0, 1.0, value=0.3, step=0.1, label="Alpha"),
177
+ gr.Slider(0.0, 1.0, value=0.7, step=0.1, label="Beta"),
178
+ gr.Slider(1, 10, value=5, step=1, label="Diffusion steps"),
179
+ ],
180
+ outputs = gr.Audio(label="Generated speech", type="filepath"),
181
+ title = "Luhya (Lunyore) TTS",
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
182
  )
183
+
184
+ demo.launch(server_name="0.0.0.0", server_port=7860)