diff --git a/tools/mtmd/mtmd-helper-gen.cpp b/tools/mtmd/mtmd-helper-gen.cpp index f1ecb4f317..884b11680e 100644 --- a/tools/mtmd/mtmd-helper-gen.cpp +++ b/tools/mtmd/mtmd-helper-gen.cpp @@ -6,6 +6,7 @@ #include #include +#include #include #include #include @@ -92,6 +93,7 @@ public: virtual int32_t step_gen(llama_token sampled, const float * h_state_in, const float ** h_state_out, bool * out_stop) = 0; virtual int32_t get_output(int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len, int64_t * out_n_samples) = 0; + protected: llama_context * lctx; mtmd_context * mctx; @@ -481,6 +483,10 @@ public: prompt_pos = 0; step_idx = 0; eos_step = -1; + chunks.clear(); + chunk_idx = 0; + n_voice_pos = 0; + chunk_budget = 0; } int32_t set_input(const mtmd_helper_gen_audio_inp * inp) override { @@ -503,8 +509,6 @@ public: LOG_ERR("mtmd_helper_gen_audio: empty prompt\n"); return 1; } - // the model may pin the tail length, otherwise guess it from the text like the reference - frames_after_eos = count_words(text) <= 4 ? 5 : 3; std::vector ids(text.size() + 16); int n_ids = llama_tokenize(vocab, text.c_str(), (int32_t) text.size(), ids.data(), @@ -515,6 +519,14 @@ public: } ids.resize((size_t) n_ids); + // long inputs degrade badly, the reference splits them and restarts each piece from + // the voice conditioning, see split_into_best_sentences() + chunks = split_chunks(ids); + chunk_idx = 0; + if (chunks.size() > 1) { + LOG_INF("mtmd_helper_gen_audio: %d tokens split into %zu chunks\n", n_ids, chunks.size()); + } + const int n_e = n_embd; auto push_row = [&](llama_token t) { prompt_embd_buf.insert(prompt_embd_buf.end(), @@ -529,10 +541,14 @@ public: } prompt_embd_buf.insert(prompt_embd_buf.end(), voice.begin(), voice.end()); } - for (llama_token t : ids) { + // every later chunk rewinds to here and re-prompts, so the voice stays primed + n_voice_pos = (int) (prompt_embd_buf.size() / (size_t) n_e); + + for (llama_token t : chunks[0]) { push_row(t); } push_row(audio_bos); + arm_chunk_budget(0); n_prompt = (int) (prompt_embd_buf.size() / (size_t) n_e); prompt_batch.reset(new decode_embd_batch(prompt_embd_buf.data(), n_prompt, 1, n_e)); @@ -591,11 +607,15 @@ public: if (out.is_eos && eos_step < 0) { eos_step = step_idx; } - // the frame of the stopping step is discarded, matching _autoregressive_generation() - if (eos_step >= 0 && step_idx >= eos_step + frames_after_eos) { - *out_stop = true; - *h_state_out = nullptr; - return 0; + // the frame of the stopping step is discarded, matching _autoregressive_generation(). + // the budget is the reference's fallback for a chunk whose eos head never fires + const bool chunk_done = (eos_step >= 0 && step_idx >= eos_step + frames_after_eos) || + step_idx >= chunk_budget; + if (chunk_done) { + if (eos_step < 0) { + LOG_WRN("mtmd_helper_gen_audio: chunk %zu hit its budget without end-of-speech\n", chunk_idx); + } + return finish_chunk(h_state_out, out_stop); } feats_buf.insert(feats_buf.end(), out.feats, out.feats + out.n_feats); @@ -675,6 +695,142 @@ private: return true; } + // token ids of the pieces the reference splits on, see split_into_best_sentences(). + // the leading token is dropped, it is the tokenizer's dummy prefix + std::vector punct_ids(const char * s) const { + std::vector ids(16); + const int n = llama_tokenize(vocab, s, (int32_t) strlen(s), ids.data(), (int32_t) ids.size(), false, false); + if (n <= 1) { + return {}; + } + return std::vector(ids.begin() + 1, ids.begin() + n); + } + + // cut after runs of boundary tokens, so punctuation stays with the sentence it ends + static std::vector> split_on(const std::vector & ids, + const std::vector & boundary) { + std::vector> out; + size_t start = 0; + bool prev_was_boundary = false; + for (size_t i = 0; i < ids.size(); i++) { + const bool is_boundary = std::find(boundary.begin(), boundary.end(), ids[i]) != boundary.end(); + if (!is_boundary && prev_was_boundary) { + out.emplace_back(ids.begin() + start, ids.begin() + i); + start = i; + } + prev_was_boundary = is_boundary; + } + out.emplace_back(ids.begin() + start, ids.end()); + return out; + } + + std::vector> split_chunks(const std::vector & ids) const { + if ((int) ids.size() <= max_chunk_tokens) { + return { ids }; + } + const std::vector eos_punct = punct_ids(".!...?"); + const std::vector mid_punct = punct_ids(",;:"); + + // oversized sentences are split again on weaker punctuation, else words get skipped + std::vector> segments; + for (auto & seg : split_on(ids, eos_punct)) { + if ((int) seg.size() <= max_chunk_tokens) { + segments.push_back(std::move(seg)); + continue; + } + auto sub = split_on(seg, mid_punct); + if (sub.size() > 1) { + for (auto & s : sub) { + segments.push_back(std::move(s)); + } + } else { + segments.push_back(std::move(seg)); + } + } + + std::vector> out; + for (auto & seg : segments) { + if (seg.empty()) { + continue; + } + if (!out.empty() && (int) (out.back().size() + seg.size()) <= max_chunk_tokens) { + out.back().insert(out.back().end(), seg.begin(), seg.end()); + } else { + out.push_back(std::move(seg)); + } + } + if (out.empty()) { + out.push_back(ids); + } + for (const auto & c : out) { + if ((int) c.size() > max_chunk_tokens) { + LOG_WRN("mtmd_helper_gen_audio: chunk of %zu tokens exceeds the %d token budget, " + "generation may skip words\n", c.size(), max_chunk_tokens); + } + } + return out; + } + + // _estimate_max_gen_len() plus the per-chunk tail guess, both in frames + void arm_chunk_budget(size_t idx) { + const int n_tok = (int) chunks[idx].size(); + chunk_budget = (int) std::ceil((n_tok / 3.0 + 2.0) * frame_rate); + // the reference guesses the tail from the word count, approximated here by tokens + frames_after_eos = n_tok <= 6 ? 5 : 3; + step_idx = 0; + eos_step = -1; + } + + // ends the current chunk and, if there is another, re-prompts it on top of the voice + int32_t finish_chunk(const float ** h_state_out, bool * out_stop) { + if (!flush_gen_wav()) { + return 1; + } + // the decoder restarts too, the next chunk's audio is not continuous with this one + dec_state.clear(); + + if (chunk_idx + 1 >= chunks.size()) { + *out_stop = true; + *h_state_out = nullptr; + return 0; + } + chunk_idx++; + + // drop this chunk's text and audio, keep the voice conditioning + llama_memory_seq_rm(llama_get_memory(lctx), seq_id, n_voice_pos, -1); + pos = n_voice_pos; + + const int n_e = n_embd; + prompt_embd_buf.clear(); + for (llama_token t : chunks[chunk_idx]) { + prompt_embd_buf.insert(prompt_embd_buf.end(), + tok_embd.begin() + (size_t) t * n_e, + tok_embd.begin() + (size_t) (t + 1) * n_e); + } + prompt_embd_buf.insert(prompt_embd_buf.end(), + tok_embd.begin() + (size_t) audio_bos * n_e, + tok_embd.begin() + (size_t) (audio_bos + 1) * n_e); + arm_chunk_budget(chunk_idx); + + // bounded by max_chunk_tokens + 1, so one decode is enough + const int n_rows = (int) (prompt_embd_buf.size() / (size_t) n_e); + decode_embd_batch batch(prompt_embd_buf.data(), n_rows, 1, n_e); + batch.set_position_normal(pos, seq_id); + batch.batch.logits[n_rows - 1] = 1; + if (llama_decode(lctx, batch.batch) != 0) { + LOG_ERR("mtmd_helper_gen_audio: chunk prompt decode failed\n"); + return 1; + } + pos += n_rows; + prompt_embd_buf.clear(); + + const float * he = llama_get_embeddings_ith(lctx, -1); + h_state_buf.assign(he, he + n_embd); + *h_state_out = h_state_buf.data(); + *out_stop = false; + return 0; + } + // same normalization as prepare_text_prompt() in the reference, it affects quality static std::string prepare_text(const std::string & in) { std::string s; @@ -789,6 +945,14 @@ private: int step_idx = 0; int eos_step = -1; int frames_after_eos = 3; + // long inputs are split, each chunk restarts from the voice conditioning + static constexpr int max_chunk_tokens = 50; // MAX_TOKEN_PER_CHUNK in the reference + static constexpr double frame_rate = 12.5; + std::vector> chunks; + size_t chunk_idx = 0; + int n_voice_pos = 0; // KV positions held by the voice conditioning + int chunk_budget = 0; + // latents are decoded a window at a time, the decoder state bridges the windows size_t window_frames = 8; std::vector feats_buf;