From 8e3cad1aa2a5e5f2101e5906efbda612983898df Mon Sep 17 00:00:00 2001 From: Concedo <39025047+LostRuins@users.noreply.github.com> Date: Thu, 16 Jan 2025 12:04:58 +0800 Subject: [PATCH] added audio caching, as a hacky fix for ST TTS bug --- expose.h | 1 + koboldcpp.py | 4 +++- otherarch/tts_adapter.cpp | 36 +++++++++++++++++++++++++++++++----- 3 files changed, 35 insertions(+), 6 deletions(-) diff --git a/expose.h b/expose.h index 12ad17a99..91eeca404 100644 --- a/expose.h +++ b/expose.h @@ -220,6 +220,7 @@ struct tts_generation_inputs const int speaker_seed = 0; const int audio_seed = 0; const bool quiet = false; + const bool nocache = false; }; struct tts_generation_outputs { diff --git a/koboldcpp.py b/koboldcpp.py index 785807eba..071580a68 100644 --- a/koboldcpp.py +++ b/koboldcpp.py @@ -296,7 +296,8 @@ class tts_generation_inputs(ctypes.Structure): _fields_ = [("prompt", ctypes.c_char_p), ("speaker_seed", ctypes.c_int), ("audio_seed", ctypes.c_int), - ("quiet", ctypes.c_bool)] + ("quiet", ctypes.c_bool), + ("nocache", ctypes.c_bool)] class tts_generation_outputs(ctypes.Structure): _fields_ = [("status", ctypes.c_int), @@ -1372,6 +1373,7 @@ def tts_generate(genparams): aseed = -1 inputs.audio_seed = aseed inputs.quiet = is_quiet + inputs.nocache = genparams.get("nocache", False) ret = handle.tts_generate(inputs) outstr = "" if ret.status==1: diff --git a/otherarch/tts_adapter.cpp b/otherarch/tts_adapter.cpp index 65da52077..ccbe891a5 100644 --- a/otherarch/tts_adapter.cpp +++ b/otherarch/tts_adapter.cpp @@ -421,6 +421,9 @@ static llama_context * cts_ctx = nullptr; //codes to speech static int ttsdebugmode = 0; static std::string ttsplatformenv, ttsdeviceenv, ttsvulkandeviceenv; static std::string last_generated_audio = ""; +static std::string last_generation_settings_prompt = ""; //for caching purposes to fix ST bug +static int last_generation_settings_speaker_seed; +static int last_generation_settings_audio_seed; static std::vector last_speaker_codes; //will store cached speaker static int last_speaker_seed = -999; @@ -555,6 +558,23 @@ tts_generation_outputs ttstype_generate(const tts_generation_inputs inputs) std::mt19937 tts_rng(audio_seed); std::mt19937 speaker_rng(speaker_seed); + //if we can reuse an old generation, do so + if(!inputs.nocache + && last_generation_settings_audio_seed == inputs.audio_seed + && last_generation_settings_speaker_seed == inputs.speaker_seed + && last_generated_audio!="" + && last_generation_settings_prompt == std::string(inputs.prompt)) + { + if(ttsdebugmode==1 || !inputs.quiet) + { + printf("\nReusing Cached Audio."); + output.data = last_generated_audio.c_str(); + output.status = 1; + return output; + } + } + + int n_decode = 0; int n_predict = 2048; //will be updated later bool next_token_uses_guide_token = true; @@ -845,18 +865,19 @@ tts_generation_outputs ttstype_generate(const tts_generation_inputs inputs) const int n_sr = 24000; // original sampling rate const int t_sr = 16000; //final target sampling rate - // zero out first 0.1 seconds or 0.05 depending on whether its seeded - const int cutout = (speaker_seed>0?(n_sr/10):(n_sr/20)); + // zero out first x seconds depending on whether its seeded + const int cutout = t_sr/4; + + audio = resample_wav(audio,n_sr,t_sr); //resample to 16k + for (int i = 0; i < cutout; ++i) { audio[i] = 0.0f; } //add some silence at the end - for (int i = 0; i < n_sr/10; ++i) { + for (int i = 0; i < t_sr/10; ++i) { audio.push_back(0.0f); } - audio = resample_wav(audio,n_sr,t_sr); //resample to 16k - last_generated_audio = save_wav16_base64(audio, t_sr); ttstime = timer_check(); @@ -867,6 +888,11 @@ tts_generation_outputs ttstype_generate(const tts_generation_inputs inputs) output.data = last_generated_audio.c_str(); output.status = 1; + + last_generation_settings_audio_seed = inputs.audio_seed; + last_generation_settings_speaker_seed = inputs.speaker_seed; + last_generation_settings_prompt = std::string(inputs.prompt); + return output; } }