whiskers.git / android / app / src / main / kotlin / com / k2fsa / sherpa / onnx / Tts.kt
Tts.kt416 lines · 11.8 KB · raw
1// Copyright (c)  2023  Xiaomi Corporation
2package com.k2fsa.sherpa.onnx
3
4import android.content.res.AssetManager
5
6data class OfflineTtsVitsModelConfig(
7    var model: String = "",
8    var lexicon: String = "",
9    var tokens: String = "",
10    var dataDir: String = "",
11    var dictDir: String = "", // unused
12    var noiseScale: Float = 0.667f,
13    var noiseScaleW: Float = 0.8f,
14    var lengthScale: Float = 1.0f,
15)
16
17data class OfflineTtsMatchaModelConfig(
18    var acousticModel: String = "",
19    var vocoder: String = "",
20    var lexicon: String = "",
21    var tokens: String = "",
22    var dataDir: String = "",
23    var dictDir: String = "", // unused
24    var noiseScale: Float = 1.0f,
25    var lengthScale: Float = 1.0f,
26)
27
28data class OfflineTtsKokoroModelConfig(
29    var model: String = "",
30    var voices: String = "",
31    var tokens: String = "",
32    var dataDir: String = "",
33    var lexicon: String = "",
34    var lang: String = "",
35    var dictDir: String = "", // unused
36    var lengthScale: Float = 1.0f,
37)
38
39data class OfflineTtsZipVoiceModelConfig(
40    var tokens: String = "",
41    var encoder: String = "",
42    var decoder: String = "",
43    var vocoder: String = "",
44    var dataDir: String = "",
45    var lexicon: String = "",
46    var featScale: Float = 0.1f,
47    var tShift: Float = 0.5f,
48    var targetRms: Float = 0.1f,
49    var guidanceScale: Float = 1.0f,
50)
51
52data class OfflineTtsKittenModelConfig(
53    var model: String = "",
54    var voices: String = "",
55    var tokens: String = "",
56    var dataDir: String = "",
57    var lengthScale: Float = 1.0f,
58)
59
60/**
61 * Configuration for Pocket TTS models.
62 *
63 * See https://k2-fsa.github.io/sherpa/onnx/tts/pocket/index.html for details.
64 *
65 * @property lmFlow Path to the LM flow model (.onnx)
66 * @property lmMain Path to the LM main model (.onnx)
67 * @property encoder Path to the encoder model (.onnx)
68 * @property decoder Path to the decoder model (.onnx)
69 * @property textConditioner Path to the text conditioner model (.onnx)
70 * @property vocabJson Path to vocabulary JSON file
71 * @property tokenScoresJson Path to token scores JSON file
72 */
73data class OfflineTtsPocketModelConfig(
74  var lmFlow: String = "",
75  var lmMain: String = "",
76  var encoder: String = "",
77  var decoder: String = "",
78  var textConditioner: String = "",
79  var vocabJson: String = "",
80  var tokenScoresJson: String = "",
81  var voiceEmbeddingCacheCapacity: Int = 50,
82)
83
84data class OfflineTtsSupertonicModelConfig(
85  var durationPredictor: String = "",
86  var textEncoder: String = "",
87  var vectorEstimator: String = "",
88  var vocoder: String = "",
89  var ttsJson: String = "",
90  var unicodeIndexer: String = "",
91  var voiceStyle: String = "",
92)
93
94data class OfflineTtsModelConfig(
95    var vits: OfflineTtsVitsModelConfig = OfflineTtsVitsModelConfig(),
96    var matcha: OfflineTtsMatchaModelConfig = OfflineTtsMatchaModelConfig(),
97    var kokoro: OfflineTtsKokoroModelConfig = OfflineTtsKokoroModelConfig(),
98    var zipvoice: OfflineTtsZipVoiceModelConfig = OfflineTtsZipVoiceModelConfig(),
99    var kitten: OfflineTtsKittenModelConfig = OfflineTtsKittenModelConfig(),
100    var pocket: OfflineTtsPocketModelConfig = OfflineTtsPocketModelConfig(),
101    var supertonic: OfflineTtsSupertonicModelConfig = OfflineTtsSupertonicModelConfig(),
102
103    var numThreads: Int = 1,
104    var debug: Boolean = false,
105    var provider: String = "cpu",
106)
107
108data class OfflineTtsConfig(
109    var model: OfflineTtsModelConfig = OfflineTtsModelConfig(),
110    var ruleFsts: String = "",
111    var ruleFars: String = "",
112    var maxNumSentences: Int = 1,
113    var silenceScale: Float = 0.2f,
114)
115
116class GeneratedAudio(
117    val samples: FloatArray,
118    val sampleRate: Int,
119) {
120    fun save(filename: String) =
121        saveImpl(filename = filename, samples = samples, sampleRate = sampleRate)
122
123    private external fun saveImpl(
124        filename: String,
125        samples: FloatArray,
126        sampleRate: Int
127    ): Boolean
128}
129
130data class GenerationConfig(
131    var silenceScale: Float = 0.2f,
132    var speed: Float = 1.0f,
133    var sid: Int = 0,
134    var referenceAudio: FloatArray? = null,
135    var referenceSampleRate: Int = 0,
136    var referenceText: String? = null,
137    var numSteps: Int = 5,
138    var extra: Map<String, String>? = null
139)
140
141class OfflineTts(
142    assetManager: AssetManager? = null,
143    var config: OfflineTtsConfig,
144) {
145    private var ptr: Long
146
147    init {
148        ptr = if (assetManager != null) {
149            newFromAsset(assetManager, config)
150        } else {
151            newFromFile(config)
152        }
153        require(ptr != 0L) {
154            "Invalid OfflineTtsConfig: failed to create native OfflineTts"
155        }
156    }
157
158    fun sampleRate() = getSampleRate(ptr)
159
160    fun numSpeakers() = getNumSpeakers(ptr)
161
162    fun generate(
163        text: String,
164        sid: Int = 0,
165        speed: Float = 1.0f
166    ): GeneratedAudio {
167        return generateImpl(ptr, text = text, sid = sid, speed = speed)
168    }
169
170    fun generateWithCallback(
171        text: String,
172        sid: Int = 0,
173        speed: Float = 1.0f,
174        callback: (samples: FloatArray) -> Int
175    ): GeneratedAudio {
176        return generateWithCallbackImpl(
177            ptr,
178            text = text,
179            sid = sid,
180            speed = speed,
181            callback = callback
182        )
183    }
184
185    fun generateWithConfig(
186      text: String,
187      config: GenerationConfig
188    ): GeneratedAudio {
189        return generateWithConfigImpl(ptr, text, config, null)
190    }
191
192    fun generateWithConfigAndCallback(
193        text: String,
194        config: GenerationConfig,
195        callback: (samples: FloatArray) -> Int
196    ): GeneratedAudio {
197        return generateWithConfigImpl(ptr, text, config, callback)
198    }
199
200    fun allocate(assetManager: AssetManager? = null) {
201        if (ptr == 0L) {
202            ptr = if (assetManager != null) {
203                newFromAsset(assetManager, config)
204            } else {
205                newFromFile(config)
206            }
207            require(ptr != 0L) {
208                "Invalid OfflineTtsConfig: failed to create native OfflineTts"
209            }
210        }
211    }
212
213    fun free() {
214        if (ptr != 0L) {
215            delete(ptr)
216            ptr = 0
217        }
218    }
219
220    protected fun finalize() {
221        if (ptr != 0L) {
222            delete(ptr)
223            ptr = 0
224        }
225    }
226
227    fun release() = finalize()
228
229    private external fun newFromAsset(
230        assetManager: AssetManager,
231        config: OfflineTtsConfig,
232    ): Long
233
234    private external fun newFromFile(
235        config: OfflineTtsConfig,
236    ): Long
237
238    private external fun delete(ptr: Long)
239    private external fun getSampleRate(ptr: Long): Int
240    private external fun getNumSpeakers(ptr: Long): Int
241
242    // The returned array has two entries:
243    //  - the first entry is an 1-D float array containing audio samples.
244    //    Each sample is normalized to the range [-1, 1]
245    //  - the second entry is the sample rate
246    private external fun generateImpl(
247        ptr: Long,
248        text: String,
249        sid: Int = 0,
250        speed: Float = 1.0f
251    ): GeneratedAudio
252
253    private external fun generateWithCallbackImpl(
254        ptr: Long,
255        text: String,
256        sid: Int = 0,
257        speed: Float = 1.0f,
258        callback: (samples: FloatArray) -> Int
259    ): GeneratedAudio
260
261
262    private external fun generateWithConfigImpl(
263        ptr: Long,
264        text: String,
265        config: GenerationConfig,
266        callback: ((samples: FloatArray) -> Int)?
267    ): GeneratedAudio
268
269    companion object {
270        init {
271            System.loadLibrary("sherpa-onnx-jni")
272        }
273    }
274}
275
276// please refer to
277// https://k2-fsa.github.io/sherpa/onnx/tts/pretrained_models/index.html
278// to download models
279fun getOfflineTtsConfig(
280    modelDir: String,
281    modelName: String, // for VITS
282    acousticModelName: String, // for Matcha
283    vocoder: String, // for Matcha
284    voices: String, // for Kokoro or kitten
285    lexicon: String,
286    dataDir: String,
287    dictDir: String, // unused
288    ruleFsts: String,
289    ruleFars: String,
290    numThreads: Int? = null,
291    isKitten: Boolean = false,
292    isSupertonic: Boolean = false,
293    durationPredictor: String = "", // for Supertonic
294    textEncoder: String = "", // for Supertonic
295    vectorEstimator: String = "", // for Supertonic
296    supertonicVocoder: String = "", // for Supertonic
297    ttsJson: String = "", // for Supertonic
298    unicodeIndexer: String = "", // for Supertonic
299    voiceStyle: String = "", // for Supertonic
300): OfflineTtsConfig {
301    // For Matcha TTS, please set
302    // acousticModelName, vocoder
303
304    // For Kokoro TTS, please set
305    // modelName, voices
306
307    // For Kitten TTS, please set
308    // modelName, voices, isKitten
309
310    // For VITS, please set
311    // modelName
312
313    // For Supertonic TTS, please set
314    // isSupertonic, durationPredictor, textEncoder, vectorEstimator,
315    // supertonicVocoder, ttsJson, unicodeIndexer, voiceStyle
316
317    val numberOfThreads = if (numThreads != null) {
318        numThreads
319    } else if (voices.isNotEmpty()) {
320        // for Kokoro and Kitten TTS models, we use more threads
321        4
322    } else {
323        2
324    }
325
326    if (!isSupertonic && modelName.isEmpty() && acousticModelName.isEmpty()) {
327        throw IllegalArgumentException("Please specify a TTS model")
328    }
329
330    if (modelName.isNotEmpty() && acousticModelName.isNotEmpty()) {
331        throw IllegalArgumentException("Please specify either a VITS or a Matcha model, but not both")
332    }
333
334    if (acousticModelName.isNotEmpty() && vocoder.isEmpty()) {
335        throw IllegalArgumentException("Please provide vocoder for Matcha TTS")
336    }
337
338    val vits = if (modelName.isNotEmpty() && voices.isEmpty() && !isSupertonic) {
339        OfflineTtsVitsModelConfig(
340            model = "$modelDir/$modelName",
341            lexicon = "$modelDir/$lexicon",
342            tokens = "$modelDir/tokens.txt",
343            dataDir = dataDir,
344        )
345    } else {
346        OfflineTtsVitsModelConfig()
347    }
348
349    val matcha = if (acousticModelName.isNotEmpty()) {
350        OfflineTtsMatchaModelConfig(
351            acousticModel = "$modelDir/$acousticModelName",
352            vocoder = vocoder,
353            lexicon = "$modelDir/$lexicon",
354            tokens = "$modelDir/tokens.txt",
355            dataDir = dataDir,
356        )
357    } else {
358        OfflineTtsMatchaModelConfig()
359    }
360
361    val kokoro = if (voices.isNotEmpty() && !isKitten && !isSupertonic) {
362        OfflineTtsKokoroModelConfig(
363            model = "$modelDir/$modelName",
364            voices = "$modelDir/$voices",
365            tokens = "$modelDir/tokens.txt",
366            dataDir = dataDir,
367            lexicon = when {
368                lexicon == "" -> lexicon
369                "," in lexicon -> lexicon
370                else -> "$modelDir/$lexicon"
371            },
372        )
373    } else {
374        OfflineTtsKokoroModelConfig()
375    }
376
377    val kitten = if (isKitten) {
378        OfflineTtsKittenModelConfig(
379            model = "$modelDir/$modelName",
380            voices = "$modelDir/$voices",
381            tokens = "$modelDir/tokens.txt",
382            dataDir = dataDir,
383        )
384    } else {
385        OfflineTtsKittenModelConfig()
386    }
387
388    val supertonic = if (isSupertonic) {
389        OfflineTtsSupertonicModelConfig(
390            durationPredictor = "$modelDir/$durationPredictor",
391            textEncoder = "$modelDir/$textEncoder",
392            vectorEstimator = "$modelDir/$vectorEstimator",
393            vocoder = "$modelDir/$supertonicVocoder",
394            ttsJson = "$modelDir/$ttsJson",
395            unicodeIndexer = "$modelDir/$unicodeIndexer",
396            voiceStyle = "$modelDir/$voiceStyle",
397        )
398    } else {
399        OfflineTtsSupertonicModelConfig()
400    }
401
402    return OfflineTtsConfig(
403        model = OfflineTtsModelConfig(
404            vits = vits,
405            matcha = matcha,
406            kokoro = kokoro,
407            kitten = kitten,
408            supertonic = supertonic,
409            numThreads = numberOfThreads,
410            debug = true,
411            provider = "cpu",
412        ),
413        ruleFsts = ruleFsts,
414        ruleFars = ruleFars,
415    )
416}