From 7942e61b3a1d1c382074bceff06f88895c438612 Mon Sep 17 00:00:00 2001 From: enne2 Date: Tue, 18 Aug 2026 15:10:02 +0200 Subject: [PATCH] feat: anti-loop guards + KV cache configurabile Prevenzione allucinazioni/loop infiniti (Qwen3-TTS autoregressivo): - block_repeated_ngrams: maschera i token che ricreerebbero un n-gram gia visto (n=4) - has_repeating_cycle: ferma la generazione su cicli periodici (periodo 1-16, 4 ripetizioni) - stuck detector: token dominante nella finestra recente (4 occorrenze in 8 token) - fallback EOS quando tutti i logits sono mascherati (evita NaN) - KV cache talker configurabile (--kv-cache, default 8192): 4096 overflowava con reference lunghe + testi lunghi ("decode would overflow cache") - parametri esposti via API (no_repeat_ngram_size, loop_max_period, loop_repeats, loop_window) e CLI (--no-repeat-ngram, --loop-period, --loop-repeats, --loop-window) - Docker: TTS_KV_CACHE env (default 8192) Validato: testo 2788 char che prima degenerava in loop ora sintetizza pulito (141s via API, nessuna ripetizione). --- docker/tts-server/docker-compose.yml | 1 + docker/tts-server/entrypoint.sh | 1 + src/gguf-weights.h | 9 ++++ src/pipeline-tts.cpp | 73 ++++++++++++++++++++++++++-- src/pipeline-tts.h | 8 ++- src/qwen.cpp | 7 ++- src/qwen.h | 19 ++++++++ src/sampling.h | 67 +++++++++++++++++++++++++ src/tts-server.h | 28 +++++++++-- tools/qwen-tts.cpp | 30 ++++++++++++ tools/tts-server.cpp | 17 +++++++ 11 files changed, 250 insertions(+), 10 deletions(-) diff --git a/docker/tts-server/docker-compose.yml b/docker/tts-server/docker-compose.yml index 22af18e..69eb831 100644 --- a/docker/tts-server/docker-compose.yml +++ b/docker/tts-server/docker-compose.yml @@ -33,6 +33,7 @@ services: HOST: 0.0.0.0 PORT: "8881" VOICE_DIR: /voices + TTS_KV_CACHE: "${TTS_KV_CACHE:-8192}" volumes: - "${MODELS_DIR:-../../models}:/models:ro" - "${VOICES_DIR:-../../voices}:/voices:ro" diff --git a/docker/tts-server/entrypoint.sh b/docker/tts-server/entrypoint.sh index b591039..18281d9 100755 --- a/docker/tts-server/entrypoint.sh +++ b/docker/tts-server/entrypoint.sh @@ -19,6 +19,7 @@ extra_args=() [ -n "$CODEC_CHUNK_DUR" ] && extra_args+=(--codec-chunk-dur "$CODEC_CHUNK_DUR") [ -n "$CODEC_LEFT_DUR" ] && extra_args+=(--codec-left-dur "$CODEC_LEFT_DUR") [ -n "$MAX_BATCH" ] && extra_args+=(--max-batch "$MAX_BATCH") +[ -n "$TTS_KV_CACHE" ] && extra_args+=(--kv-cache "$TTS_KV_CACHE") [ -n "$MAX_PREFILL_TOKENS" ] && extra_args+=(--max-prefill-tokens "$MAX_PREFILL_TOKENS") [ "$NO_FA" = "1" ] && extra_args+=(--no-fa) [ "$CLAMP_FP16" = "1" ] && extra_args+=(--clamp-fp16) diff --git a/src/gguf-weights.h b/src/gguf-weights.h index af820f4..7888aec 100644 --- a/src/gguf-weights.h +++ b/src/gguf-weights.h @@ -368,6 +368,15 @@ static uint32_t gf_get_u32(const GGUFModel & gf, const char * key) { return gguf_get_val_u32(gf.gguf, idx); } +// Read a uint32 KV value with a fallback default when the key is absent. +static uint32_t gf_get_u32_opt(const GGUFModel & gf, const char * key, uint32_t def) { + int64_t idx = gguf_find_key(gf.gguf, key); + if (idx < 0) { + return def; + } + return gguf_get_val_u32(gf.gguf, idx); +} + // Read a float32 KV value (returns 0 if not found) static float gf_get_f32(const GGUFModel & gf, const char * key) { int64_t idx = gguf_find_key(gf.gguf, key); diff --git a/src/pipeline-tts.cpp b/src/pipeline-tts.cpp index 211a16c..2eba3a1 100644 --- a/src/pipeline-tts.cpp +++ b/src/pipeline-tts.cpp @@ -103,6 +103,11 @@ static void parse_generation_defaults(const GGUFModel & gf, GenerationDefaults & g.subtalker_top_p = gf_get_f32(gf, "generation.subtalker_top_p"); g.subtalker_temperature = gf_get_f32(gf, "generation.subtalker_temperature"); g.max_new_tokens = (int) gf_get_u32(gf, "generation.max_new_tokens"); + // Anti-loop guards: optional GGUF keys, 0 = engine default. + g.no_repeat_ngram_size = (int) gf_get_u32_opt(gf, "generation.no_repeat_ngram_size", 0); + g.loop_max_period = (int) gf_get_u32_opt(gf, "generation.loop_max_period", 0); + g.loop_repeats = (int) gf_get_u32_opt(gf, "generation.loop_repeats", 0); + g.loop_window = (int) gf_get_u32_opt(gf, "generation.loop_window", 0); } // Ensure the static predictor graph set for batch width N exists: the @@ -152,7 +157,8 @@ bool pipeline_tts_load(PipelineTTS * pt, bool use_fa, bool clamp_fp16, int max_batch, - float codec_chunk_sec) { + float codec_chunk_sec, + int talker_kv_size) { pt->bp = bp; pt->backend = bp.backend; pt->sched = NULL; @@ -260,11 +266,11 @@ bool pipeline_tts_load(PipelineTTS * pt, } // KV caches, one set per slot: the talker holds the LM context up - // to 4096 positions (the longest ICL prompt observed is ~250 + - // max_new_tokens ~ 1500, so 4096 has 60% headroom). Predictor holds - // one frame of 16 sub-steps per slot. + // to talker_kv_size positions (default 8192; 4096 was the historical + // value and overflowed with long references + long texts). Predictor + // holds one frame of 16 sub-steps per slot. if (!kv_cache_init(&pt->talker_kv, pt->talker.num_hidden_layers, pt->talker.num_key_value_heads, - pt->talker.head_dim, 4096, pt->max_batch, pt->backend)) { + pt->talker.head_dim, talker_kv_size > 0 ? talker_kv_size : 8192, pt->max_batch, pt->backend)) { ggml_backend_sched_free(pt->sched); pt->sched = NULL; pipeline_codec_free(&pt->codec); @@ -1171,6 +1177,23 @@ void tts_engine_step(TtsEngine * e, std::vector * retired) { // codec_eos. Then run the upstream sampling chain. Timer t_host; apply_suppress(s.logits.data(), vocab, vocab - 1024, vocab, codec_eos_id); + // Anti-loop: block candidate tokens that would recreate an already + // seen n-gram (deterministic phrase-loop prevention). + if (p->no_repeat_ngram_size > 1) { + block_repeated_ngrams(s.logits.data(), vocab, s.talker_history.data(), + (int) s.talker_history.size(), p->no_repeat_ngram_size); + // If everything got masked (all -inf), fall back to EOS instead + // of feeding an all -inf distribution to softmax (NaN). + float max_l = -INFINITY; + for (int v = 0; v < vocab; v++) { + if (s.logits[v] > max_l) { + max_l = s.logits[v]; + } + } + if (max_l == -INFINITY) { + s.logits[codec_eos_id] = 0.0f; + } + } float u_c0 = 0.0f; int c0 = sample_top_k_p(s.logits.data(), vocab, s.talker_T, p->top_k, p->top_p, p->repetition_penalty, s.talker_history.data(), (int) s.talker_history.size(), s.job->resolved_seed, @@ -1275,6 +1298,46 @@ void tts_engine_step(TtsEngine * e, std::vector * retired) { s.all_codes.push_back(codes); s.talker_history.push_back(s.pending_c0); + // Anti-loop: stop as soon as the generated c0 sequence + // shows a periodic repetition (cycle detector) or a + // "stuck" state (one token dominating the recent window). + // The current frame is still decoded; the slot is skipped + // from the next step on. + bool looped = false; + if (p->loop_max_period > 0 && p->loop_repeats > 1 && + has_repeating_cycle(s.talker_history.data(), (int) s.talker_history.size(), + p->loop_max_period, p->loop_repeats, p->loop_window)) { + looped = true; + } + // Stuck detector: a single c0 token dominating a tight + // recent window (>= loop_repeats occurrences in the last + // 2*loop_repeats tokens) is a strong loop signature even + // when the tokens are not exactly periodic. The tight + // window avoids false positives on natural speech. + if (!looped && p->loop_max_period > 0 && p->loop_repeats > 1) { + const int n_hist = (int) s.talker_history.size(); + const int win = std::min(p->loop_window, p->loop_repeats * 2); + const int begin = std::max(0, n_hist - win); + const int window = n_hist - begin; + if (window >= p->loop_repeats) { + int32_t last = s.talker_history.back(); + int cnt = 0; + for (int h = begin; h < n_hist; h++) { + if (s.talker_history[(size_t) h] == last) { + cnt++; + } + } + if (cnt >= p->loop_repeats) { + looped = true; + } + } + } + if (looped) { + qt_log(QT_LOG_INFO, "[Pipeline] repetition loop detected at step %d (slot %d), stopping", + s.step, i); + s.finished = true; + } + // Streaming slots stage this frame through // all_codes.back() and has_frame; the shared codec // flush after this loop decodes every lane in one diff --git a/src/pipeline-tts.h b/src/pipeline-tts.h index 4055e57..3fdb842 100644 --- a/src/pipeline-tts.h +++ b/src/pipeline-tts.h @@ -72,6 +72,11 @@ struct GenerationDefaults { float subtalker_top_p; float subtalker_temperature; int max_new_tokens; + // Anti-loop guards (0 = use engine default). + int no_repeat_ngram_size; + int loop_max_period; + int loop_repeats; + int loop_window; }; struct PromptPrefixCacheEntry { @@ -210,7 +215,8 @@ bool pipeline_tts_load(PipelineTTS * pt, bool use_fa, bool clamp_fp16, int max_batch, - float codec_chunk_sec); + float codec_chunk_sec, + int talker_kv_size); void pipeline_tts_free(PipelineTTS * pt); diff --git a/src/qwen.cpp b/src/qwen.cpp index 9f96e61..df0f625 100644 --- a/src/qwen.cpp +++ b/src/qwen.cpp @@ -223,6 +223,7 @@ void qt_init_default_params(struct qt_init_params * p) { p->use_fa = true; p->clamp_fp16 = false; p->max_batch = 1; + p->talker_kv_size = 8192; p->codec_chunk_sec = QT_CODEC_CHUNK_SEC_DEFAULT; } @@ -247,6 +248,10 @@ void qt_tts_default_params(struct qt_tts_params * p) { p->subtalker_temperature = 0.9f; p->subtalker_top_k = 50; p->subtalker_top_p = 1.0f; + p->no_repeat_ngram_size = 4; + p->loop_max_period = 16; + p->loop_repeats = 4; + p->loop_window = 64; p->dump_dir = nullptr; p->cancel = nullptr; p->cancel_user_data = nullptr; @@ -362,7 +367,7 @@ struct qt_context * qt_init(const struct qt_init_params * params) { } if (!pipeline_tts_load(&q->pt, params->talker_path, params->codec_path, q->bp, params->use_fa, - params->clamp_fp16, max_batch, chunk_sec)) { + params->clamp_fp16, max_batch, chunk_sec, params->talker_kv_size)) { qt_throw("qt_init: pipeline_tts_load failed for '%s' / '%s'", params->talker_path, params->codec_path); } diff --git a/src/qwen.h b/src/qwen.h index bbbc47a..f121fb2 100644 --- a/src/qwen.h +++ b/src/qwen.h @@ -142,6 +142,14 @@ struct qt_init_params { // and thread safe in both modes. int max_batch; + // Talker KV cache size in positions (default: 8192). The autoregressive + // talker holds prompt + generated c0 tokens in this cache; a value too + // small for the prompt length + max_new_tokens causes a hard overflow + // ("decode would overflow cache") and degraded/looping output as the + // model approaches the limit. 4096 was the historical default; 8192 + // doubles the headroom for long references and long texts. + int talker_kv_size; + // Chunk width of the buffered codec decode, in seconds of // audio, resolved to an integer frame count at the codec frame rate // by qt_init and applied to every synthesis on the handle. A chunk @@ -313,6 +321,17 @@ struct qt_tts_params { int subtalker_top_k; float subtalker_top_p; + // Anti-loop / anti-hallucination guards (defaults: ngram 4, period 8, + // repeats 4, window 64). no_repeat_ngram_size masks candidate tokens + // that would recreate an already-seen n-gram (0 disables). The cycle + // detector stops generation when the last `loop_window` c0 tokens + // contain a periodic repetition of period <= loop_max_period repeated + // at least loop_repeats times (loop_max_period <= 0 disables). + int no_repeat_ngram_size; + int loop_max_period; + int loop_repeats; + int loop_window; + // Intermediate tensor dump directory. NULL disables dumps. Debug // only, slows the run. const char * dump_dir; diff --git a/src/sampling.h b/src/sampling.h index e6c13fd..1febc2d 100644 --- a/src/sampling.h +++ b/src/sampling.h @@ -68,6 +68,73 @@ static inline void apply_repetition_penalty(float * logits, } } +// Block candidate tokens that would recreate an already-seen n-gram +// (no_repeat_ngram, HF-style). For every previous occurrence of the +// current (n-1)-gram suffix in `history`, the token that followed it is +// masked to -inf. This deterministically prevents verbatim phrase loops. +// ngram_size <= 1 disables the filter. The caller must ensure that at +// least one logit stays finite (e.g. EOS) or fall back to EOS. +static inline void block_repeated_ngrams(float * logits, + int V, + const int32_t * history, + int n_history, + int ngram_size) { + if (ngram_size <= 1 || n_history < ngram_size - 1) { + return; + } + const int prefix_len = ngram_size - 1; + const int32_t * suffix = history + (n_history - prefix_len); + for (int start = 0; start + ngram_size <= n_history; start++) { + bool match = true; + for (int j = 0; j < prefix_len; j++) { + if (history[start + j] != suffix[j]) { + match = false; + break; + } + } + if (!match) { + continue; + } + int32_t cand = history[start + prefix_len]; + if (cand >= 0 && cand < V) { + logits[cand] = -INFINITY; + } + } +} + +// Detect periodic repetition in the generated token sequence. Returns +// true when the last `window` tokens contain a cycle of period in +// [1, max_period] repeated at least `repeats` times. Catches "A B A B" +// style loops before they can run away. max_period <= 0 disables. +static inline bool has_repeating_cycle(const int32_t * history, + int n_history, + int max_period, + int repeats, + int window) { + if (max_period <= 0 || repeats < 2 || n_history < 2) { + return false; + } + const int begin = std::max(0, n_history - window); + for (int period = 1; period <= max_period; period++) { + const int required = period * repeats; + if (n_history - begin < required) { + continue; + } + const int cycle_start = n_history - required; + bool repeating = true; + for (int i = cycle_start + period; i < n_history; i++) { + if (history[i] != history[i - period]) { + repeating = false; + break; + } + } + if (repeating) { + return true; + } + } + return false; +} + // Stochastic sampler. Pipeline mirrors HF generate() in F32 : // 1. repetition_penalty(history) // 2. temperature divide diff --git a/src/tts-server.h b/src/tts-server.h index e68b941..aaefb20 100644 --- a/src/tts-server.h +++ b/src/tts-server.h @@ -56,6 +56,12 @@ struct tts_request { float temperature; // 0 selects greedy decoding float top_p; // in (0, 1] float repetition_penalty; // strictly positive + + // Optional anti-loop overrides. -1 marks unset (engine defaults). + int no_repeat_ngram_size; // 0 disables n-gram blocking + int loop_max_period; // 0 disables the cycle detector + int loop_repeats; // repeats required to declare a cycle + int loop_window; // lookback window for cycle detection }; // One voice registration parsed from the POST /v1/audio/voices JSON body. @@ -188,6 +194,10 @@ static bool tts_parse_request(const std::string & body, tts_request & req, std:: req.temperature = NAN; req.top_p = NAN; req.repetition_penalty = NAN; + req.no_repeat_ngram_size = -1; + req.loop_max_period = -1; + req.loop_repeats = -1; + req.loop_window = -1; auto opt_int = [&](const char * key, int64_t lo, int64_t hi, int64_t & out) -> bool { yyjson_val * v = yyjson_obj_get(root, key); @@ -216,12 +226,24 @@ static bool tts_parse_request(const std::string & body, tts_request & req, std:: int64_t max_new = -1; int64_t top_k = -1; + int64_t ngram = -1; + int64_t lperiod = -1; + int64_t lreps = -1; + int64_t lwin = -1; bool ok = opt_int("seed", INT64_MIN, INT64_MAX, req.seed) && opt_int("max_new_tokens", 1, INT32_MAX, max_new) && opt_int("top_k", 0, INT32_MAX, top_k) && opt_num("temperature", 0.0, FLT_MAX, req.temperature) && opt_num("top_p", DBL_MIN, 1.0, req.top_p) && - opt_num("repetition_penalty", DBL_MIN, FLT_MAX, req.repetition_penalty); - req.max_new_tokens = (int) max_new; - req.top_k = (int) top_k; + opt_num("repetition_penalty", DBL_MIN, FLT_MAX, req.repetition_penalty) && + opt_int("no_repeat_ngram_size", 0, INT32_MAX, ngram) && + opt_int("loop_max_period", 0, INT32_MAX, lperiod) && + opt_int("loop_repeats", 0, INT32_MAX, lreps) && + opt_int("loop_window", 0, INT32_MAX, lwin); + req.max_new_tokens = (int) max_new; + req.top_k = (int) top_k; + req.no_repeat_ngram_size = (int) ngram; + req.loop_max_period = (int) lperiod; + req.loop_repeats = (int) lreps; + req.loop_window = (int) lwin; yyjson_doc_free(doc); diff --git a/tools/qwen-tts.cpp b/tools/qwen-tts.cpp index f5ea6e6..0a17776 100644 --- a/tools/qwen-tts.cpp +++ b/tools/qwen-tts.cpp @@ -53,6 +53,7 @@ static void print_usage(const char * prog) { " --ref-spk and --ref-text, enables ICL clone mode)\n" " --ref-text Transcript file for the reference (enables ICL clone mode)\n" " --max-new Max new audio frames (default: 2048)\n" + " --kv-cache Talker KV cache size in positions (default: 8192)\n" " --codec-chunk-dur Codec decode chunk duration in seconds (default: 24.0)\n" " --stream-by-line Flush synthesis at each newline, one WAV header per line (-o '-')\n\n" "Sampling:\n" @@ -62,6 +63,10 @@ static void print_usage(const char * prog) { " --top-k Talker top-k (default: 50, 0 disables)\n" " --top-p Talker top-p (default: 1.0)\n" " --rep-pen Talker repetition penalty (default: 1.05)\n" + " --no-repeat-ngram Block repeated n-grams (default: 4, 0 disables)\n" + " --loop-period Cycle detector max period (default: 8, 0 disables)\n" + " --loop-repeats Cycle detector repeats required (default: 4)\n" + " --loop-window Cycle detector lookback window (default: 64)\n" " --sub-temp Sub-talker temperature (default: 0.9)\n" " --sub-top-k Sub-talker top-k (default: 50)\n" " --sub-top-p Sub-talker top-p (default: 1.0)\n\n" @@ -92,6 +97,10 @@ struct Args { int top_k; float top_p; float repetition_penalty; + int no_repeat_ngram_size; + int loop_max_period; + int loop_repeats; + int loop_window; int subtalker_top_k; float subtalker_top_p; float subtalker_temperature; @@ -100,6 +109,7 @@ struct Args { bool clamp_fp16; bool stream_by_line; float codec_chunk_sec; + int talker_kv_size; }; // Read all of stdin into a string. Binary mode on Windows so UTF-16 input @@ -188,6 +198,10 @@ static bool parse_args(int argc, char ** argv, Args & a) { a.top_k = 50; a.top_p = 1.0f; a.repetition_penalty = 1.05f; + a.no_repeat_ngram_size = 4; + a.loop_max_period = 16; + a.loop_repeats = 4; + a.loop_window = 64; a.subtalker_do_sample = true; a.subtalker_top_k = 50; a.subtalker_top_p = 1.0f; @@ -198,6 +212,7 @@ static bool parse_args(int argc, char ** argv, Args & a) { // Chunk sentinel : qt_init resolves a non positive value to the // library default. a.codec_chunk_sec = 0.0f; + a.talker_kv_size = 8192; for (int i = 1; i < argc; i++) { const char * arg = argv[i]; if (std::strcmp(arg, "-h") == 0 || std::strcmp(arg, "--help") == 0) { @@ -245,6 +260,14 @@ static bool parse_args(int argc, char ** argv, Args & a) { a.top_p = (float) std::atof(argv[++i]); } else if (std::strcmp(arg, "--rep-pen") == 0 && i + 1 < argc) { a.repetition_penalty = (float) std::atof(argv[++i]); + } else if (std::strcmp(arg, "--no-repeat-ngram") == 0 && i + 1 < argc) { + a.no_repeat_ngram_size = std::atoi(argv[++i]); + } else if (std::strcmp(arg, "--loop-period") == 0 && i + 1 < argc) { + a.loop_max_period = std::atoi(argv[++i]); + } else if (std::strcmp(arg, "--loop-repeats") == 0 && i + 1 < argc) { + a.loop_repeats = std::atoi(argv[++i]); + } else if (std::strcmp(arg, "--loop-window") == 0 && i + 1 < argc) { + a.loop_window = std::atoi(argv[++i]); } else if (std::strcmp(arg, "--sub-temp") == 0 && i + 1 < argc) { a.subtalker_temperature = (float) std::atof(argv[++i]); } else if (std::strcmp(arg, "--sub-top-k") == 0 && i + 1 < argc) { @@ -259,6 +282,8 @@ static bool parse_args(int argc, char ** argv, Args & a) { a.stream_by_line = true; } else if (std::strcmp(arg, "--codec-chunk-dur") == 0 && i + 1 < argc) { a.codec_chunk_sec = (float) std::atof(argv[++i]); + } else if (std::strcmp(arg, "--kv-cache") == 0 && i + 1 < argc) { + a.talker_kv_size = std::atoi(argv[++i]); } else if (std::strcmp(arg, "-o") == 0 && i + 1 < argc) { a.out_wav = argv[++i]; } else { @@ -281,6 +306,7 @@ static int run(const Args & a) { iparams.use_fa = a.use_fa; iparams.clamp_fp16 = a.clamp_fp16; iparams.codec_chunk_sec = a.codec_chunk_sec; + iparams.talker_kv_size = a.talker_kv_size; qt_context * q = qt_init(&iparams); if (!q) { @@ -407,6 +433,10 @@ static int run(const Args & a) { params.top_k = a.top_k; params.top_p = a.top_p; params.repetition_penalty = a.repetition_penalty; + params.no_repeat_ngram_size = a.no_repeat_ngram_size; + params.loop_max_period = a.loop_max_period; + params.loop_repeats = a.loop_repeats; + params.loop_window = a.loop_window; params.subtalker_do_sample = a.subtalker_do_sample; params.subtalker_temperature = a.subtalker_temperature; params.subtalker_top_k = a.subtalker_top_k; diff --git a/tools/tts-server.cpp b/tools/tts-server.cpp index 4bdb1f4..453b15d 100644 --- a/tools/tts-server.cpp +++ b/tools/tts-server.cpp @@ -51,6 +51,7 @@ static void print_usage(const char * prog) { " --port Listen port (default: 8080)\n" " --lang Language label (default: auto)\n" " --max-batch Concurrent requests batched on the GPU (default: 1)\n" + " --kv-cache Talker KV cache size in positions (default: 8192)\n" " --no-fa Disable flash attention\n" " --clamp-fp16 Clamp hidden states to FP16 range\n" " --codec-chunk-dur Codec decode chunk duration in seconds, wav responses (default: 24.0)\n", @@ -73,6 +74,7 @@ int main(int argc, char ** argv) { bool use_fa = true; bool clamp_fp16 = false; int max_batch = 1; + int talker_kv_size = 8192; // Chunk sentinel : qt_init resolves a non positive value to the // library default. float codec_chunk_dur = 0.0f; @@ -97,6 +99,8 @@ int main(int argc, char ** argv) { clamp_fp16 = true; } else if (!std::strcmp(arg, "--max-batch") && i + 1 < argc) { max_batch = std::atoi(argv[++i]); + } else if (!std::strcmp(arg, "--kv-cache") && i + 1 < argc) { + talker_kv_size = std::atoi(argv[++i]); } else if (!std::strcmp(arg, "--codec-chunk-dur") && i + 1 < argc) { codec_chunk_dur = (float) std::atof(argv[++i]); } else if (!std::strcmp(arg, "--help") || !std::strcmp(arg, "-h")) { @@ -121,6 +125,7 @@ int main(int argc, char ** argv) { iparams.use_fa = use_fa; iparams.clamp_fp16 = clamp_fp16; iparams.max_batch = max_batch; + iparams.talker_kv_size = talker_kv_size; iparams.codec_chunk_sec = codec_chunk_dur; struct qt_context * q = qt_init(&iparams); @@ -297,6 +302,18 @@ int main(int argc, char ** argv) { if (!std::isnan(req.repetition_penalty)) { p.repetition_penalty = req.repetition_penalty; } + if (req.no_repeat_ngram_size != -1) { + p.no_repeat_ngram_size = req.no_repeat_ngram_size; + } + if (req.loop_max_period != -1) { + p.loop_max_period = req.loop_max_period; + } + if (req.loop_repeats != -1) { + p.loop_repeats = req.loop_repeats; + } + if (req.loop_window != -1) { + p.loop_window = req.loop_window; + } // Trampoline : the C ABI on_chunk forwards to the C++ sink. const tts_sink * sink_ptr = &sink;