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}