1# Requirements:
2# pip install onnxruntime librosa soundfile transformers numpy
3# If you want GPU inference: pip install onnxruntime-gpu (and ensure CUDA toolkit is available)
4
5import os
6import numpy as np
7import onnxruntime
8from transformers import AutoTokenizer
9import soundfile as sf
10
11TTS_ONNX_MODEL_PATH = "swahili_tts.onnx" # path to your .onnx file
12TTS_TOKENIZER_ID = "facebook/mms-tts-swh" # or whichever tokenizer you used
13OUTPUT_SAMPLE_RATE = 16000
14OUT_DIR = "tts_outputs"
15os.makedirs(OUT_DIR, exist_ok=True)
16
17def create_onnx_session(onnx_path: str):
18 """Create an ONNX Runtime session using GPU if available, otherwise CPU."""
19 providers = ["CPUExecutionProvider"]
20 try:
21 # prefer CUDA if available
22 providers = ["CUDAExecutionProvider", "CPUExecutionProvider"]
23 sess = onnxruntime.InferenceSession(onnx_path, providers=providers)
24 print("Using CUDAExecutionProvider for ONNX Runtime.")
25 except Exception:
26 sess = onnxruntime.InferenceSession(onnx_path, providers=["CPUExecutionProvider"])
27 print("CUDA not available — using CPUExecutionProvider for ONNX Runtime.")
28 return sess
29
30def generate_speech_from_onnx(text: str,
31 onnx_session: onnxruntime.InferenceSession,
32 tokenizer: AutoTokenizer,
33 out_path: str = None) -> str:
34 """
35 Synthesize speech from text using an ONNX TTS model.
36 Returns path to WAV file (16kHz, int16).
37 """
38 if not text:
39 raise ValueError("Empty text provided.")
40
41 # Tokenize to numpy inputs (match what the ONNX model expects)
42 # NOTE: many TTS tokenizers return {"input_ids": np.array(...)} — adapt if your tokenizer differs
43 inputs = tokenizer(text, return_tensors="np", padding=True)
44 # Identify ONNX input name (assume first input)
45 input_name = onnx_session.get_inputs()[0].name
46
47 # Prepare ort_inputs dict using names expected by ONNX model
48 ort_inputs = {input_name: inputs["input_ids"].astype(np.int64)}
49
50 # Run ONNX inference
51 ort_outs = onnx_session.run(None, ort_inputs)
52
53 # The model should return a raw waveform or float array convertible to waveform.
54 # In many single-file TTS ONNX exports the first output is the waveform
55 audio_array = ort_outs[0]
56
57 # Flatten in case it's multi-dim and ensure 1-D waveform
58 audio_waveform = audio_array.flatten()
59
60 # If float waveform in -1..1, convert to int16; else try to coerce to int16
61 if np.issubdtype(audio_waveform.dtype, np.floating):
62 # clip then convert
63 audio_clip = np.clip(audio_waveform, -1.0, 1.0)
64 audio_int16 = (audio_clip * 32767.0).astype(np.int16)
65 else:
66 # if it's already int16-like, cast (safeguard)
67 audio_int16 = audio_waveform.astype(np.int16)
68
69 # Compose output filename
70 if out_path is None:
71 out_path = os.path.join(OUT_DIR, f"salama_tts_{abs(hash(text)) & 0xFFFF_FFFF}.wav")
72
73 # Save with soundfile (16kHz)
74 sf.write(out_path, audio_int16, samplerate=OUTPUT_SAMPLE_RATE, subtype="PCM_16")
75 return out_path
76
77if __name__ == "__main__":
78 # Example usage
79 sess = create_onnx_session(TTS_ONNX_MODEL_PATH)
80 tokenizer = AutoTokenizer.from_pretrained(TTS_TOKENIZER_ID)
81
82 example_text = "Karibu kwenye mfumo wa SALAMA unaozalisha sauti asilia ya Kiswahili."
83 out_wav = generate_speech_from_onnx(example_text, sess, tokenizer)
84 print("Saved synthesized audio to:", out_wav)
85