mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-19 17:24:57 +02:00
mtmd: add audio out stream api
This commit is contained in:
@@ -48,29 +48,38 @@ static llama_token find_special_token(const llama_vocab * vocab, const std::stri
|
||||
return LLAMA_TOKEN_NULL;
|
||||
}
|
||||
|
||||
static void put_bytes(std::vector<char> & buf, const void * p, size_t n) {
|
||||
const char * c = (const char *) p;
|
||||
buf.insert(buf.end(), c, c + n);
|
||||
}
|
||||
|
||||
// data_sz == UINT32_MAX writes the "unknown length" sentinel (streaming), same as ffmpeg does on a pipe
|
||||
static void write_wav16_header(std::vector<char> & buf, uint32_t data_sz, int32_t rate) {
|
||||
const uint32_t riff_sz = data_sz == UINT32_MAX ? UINT32_MAX : 36 + data_sz;
|
||||
const uint32_t fmt_sz = 16, byte_rate = (uint32_t) rate * 2;
|
||||
const uint16_t fmt = 1, ch = 1, align = 2, bits = 16;
|
||||
const uint32_t rate32 = (uint32_t) rate;
|
||||
put_bytes(buf, "RIFF", 4); put_bytes(buf, &riff_sz, 4); put_bytes(buf, "WAVE", 4);
|
||||
put_bytes(buf, "fmt ", 4); put_bytes(buf, &fmt_sz, 4);
|
||||
put_bytes(buf, &fmt, 2); put_bytes(buf, &ch, 2); put_bytes(buf, &rate32, 4);
|
||||
put_bytes(buf, &byte_rate, 4); put_bytes(buf, &align, 2); put_bytes(buf, &bits, 2);
|
||||
put_bytes(buf, "data", 4); put_bytes(buf, &data_sz, 4);
|
||||
}
|
||||
|
||||
static void append_wav16_pcm(std::vector<char> & buf, const float * pcm, size_t n) {
|
||||
for (size_t i = 0; i < n; i++) {
|
||||
int16_t s = (int16_t) (std::max(-1.0f, std::min(1.0f, pcm[i])) * 32767.0f);
|
||||
put_bytes(buf, &s, 2);
|
||||
}
|
||||
}
|
||||
|
||||
static bool write_wav16(std::vector<char> & buf, const std::vector<float> & pcm, int32_t rate) {
|
||||
// RIFF chunk sizes are 32-bit; refuse to emit a file with a truncated header
|
||||
if (pcm.size() > ((size_t) UINT32_MAX - 36) / 2) {
|
||||
return false;
|
||||
}
|
||||
const uint32_t data_sz = (uint32_t) (pcm.size() * 2);
|
||||
const uint32_t riff_sz = 36 + data_sz;
|
||||
const uint32_t fmt_sz = 16, byte_rate = (uint32_t) rate * 2;
|
||||
const uint16_t fmt = 1, ch = 1, align = 2, bits = 16;
|
||||
const uint32_t rate32 = (uint32_t) rate;
|
||||
auto put = [&](const void * p, size_t n) {
|
||||
const char * c = (const char *) p;
|
||||
buf.insert(buf.end(), c, c + n);
|
||||
};
|
||||
put("RIFF", 4); put(&riff_sz, 4); put("WAVE", 4);
|
||||
put("fmt ", 4); put(&fmt_sz, 4);
|
||||
put(&fmt, 2); put(&ch, 2); put(&rate32, 4);
|
||||
put(&byte_rate, 4); put(&align, 2); put(&bits, 2);
|
||||
put("data", 4); put(&data_sz, 4);
|
||||
for (float v : pcm) {
|
||||
int16_t s = (int16_t) (std::max(-1.0f, std::min(1.0f, v)) * 32767.0f);
|
||||
put(&s, 2);
|
||||
}
|
||||
write_wav16_header(buf, (uint32_t) (pcm.size() * 2), rate);
|
||||
append_wav16_pcm(buf, pcm.data(), pcm.size());
|
||||
return true;
|
||||
}
|
||||
|
||||
@@ -119,6 +128,9 @@ public:
|
||||
prompt_batch.reset();
|
||||
n_prompt = 0;
|
||||
prompt_pos = 0;
|
||||
stream = false;
|
||||
pcm_sent = 0;
|
||||
wav_header_sent = false;
|
||||
}
|
||||
|
||||
int32_t set_input(const mtmd_helper_gen_audio_inp * inp) override {
|
||||
@@ -204,6 +216,7 @@ public:
|
||||
top_k = inp->top_k > 0 ? inp->top_k : 50;
|
||||
top_p = inp->top_p > 0 ? inp->top_p : 1.0f;
|
||||
out_type = inp->out_type;
|
||||
stream = inp->stream;
|
||||
|
||||
// the text stream keeps flowing during generation: after frame k, the input adds
|
||||
// trailing text row k on top of the codes embedding, then tts_eos, then tts_pad
|
||||
@@ -294,21 +307,46 @@ public:
|
||||
}
|
||||
|
||||
*out_sample_rate = info.sample_rate;
|
||||
|
||||
if (!stream) {
|
||||
if (out_n_samples) {
|
||||
*out_n_samples = (int64_t) audio_pcm.size();
|
||||
}
|
||||
if (out_type == MTMD_HELPER_GEN_AUDIO_OUTTYPE_PCM) {
|
||||
*out_data = (const char *) audio_pcm.data();
|
||||
*out_data_len = audio_pcm.size() * sizeof(float);
|
||||
return 0;
|
||||
}
|
||||
out_buf.clear();
|
||||
if (!write_wav16(out_buf, audio_pcm, info.sample_rate)) {
|
||||
LOG_ERR("mtmd_helper_gen_audio: output too large for WAV\n");
|
||||
return 1;
|
||||
}
|
||||
*out_data = out_buf.data();
|
||||
*out_data_len = out_buf.size();
|
||||
return 0;
|
||||
}
|
||||
|
||||
// streaming: only return audio produced since the previous call
|
||||
const size_t n_new = audio_pcm.size() - pcm_sent;
|
||||
if (out_n_samples) {
|
||||
*out_n_samples = (int64_t) audio_pcm.size();
|
||||
*out_n_samples = (int64_t) n_new;
|
||||
}
|
||||
|
||||
if (out_type == MTMD_HELPER_GEN_AUDIO_OUTTYPE_PCM) {
|
||||
*out_data = (const char *) audio_pcm.data();
|
||||
*out_data_len = audio_pcm.size() * sizeof(float);
|
||||
*out_data = (const char *) (audio_pcm.data() + pcm_sent);
|
||||
*out_data_len = n_new * sizeof(float);
|
||||
pcm_sent = audio_pcm.size();
|
||||
return 0;
|
||||
}
|
||||
|
||||
out_buf.clear();
|
||||
if (!write_wav16(out_buf, audio_pcm, info.sample_rate)) {
|
||||
LOG_ERR("mtmd_helper_gen_audio: output too large for WAV\n");
|
||||
return 1;
|
||||
if (!wav_header_sent) {
|
||||
write_wav16_header(out_buf, UINT32_MAX, info.sample_rate);
|
||||
wav_header_sent = true;
|
||||
}
|
||||
append_wav16_pcm(out_buf, audio_pcm.data() + pcm_sent, n_new);
|
||||
pcm_sent = audio_pcm.size();
|
||||
*out_data = out_buf.data();
|
||||
*out_data_len = out_buf.size();
|
||||
return 0;
|
||||
@@ -442,6 +480,9 @@ private:
|
||||
std::vector<float> h_state_buf;
|
||||
mtmd_helper_gen_audio_outtype out_type = MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV;
|
||||
std::vector<char> out_buf;
|
||||
bool stream = false;
|
||||
size_t pcm_sent = 0; // samples already returned by get_output()
|
||||
bool wav_header_sent = false;
|
||||
};
|
||||
|
||||
static std::unique_ptr<mtmd_gen_audio_pipeline> make_pipeline(llama_context * lctx, mtmd_context * mctx) {
|
||||
@@ -473,6 +514,14 @@ void mtmd_helper_gen_audio_reset(mtmd_helper_gen_audio * ctx) {
|
||||
}
|
||||
}
|
||||
|
||||
struct mtmd_helper_gen_audio_inp mtmd_helper_gen_audio_inp_default(void) {
|
||||
mtmd_helper_gen_audio_inp inp{};
|
||||
inp.top_k = 50;
|
||||
inp.top_p = 1.0f;
|
||||
inp.out_type = MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV;
|
||||
return inp;
|
||||
}
|
||||
|
||||
int32_t mtmd_helper_gen_audio_set_input(mtmd_helper_gen_audio * ctx, const mtmd_helper_gen_audio_inp * inp) {
|
||||
if (!ctx->pipeline) {
|
||||
LOG_ERR("mtmd_helper_gen_audio: unsupported or missing gen-audio pipeline\n");
|
||||
|
||||
@@ -175,6 +175,7 @@ enum mtmd_helper_gen_audio_outtype {
|
||||
MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV, // WAV PCM 16-bit LE, mono
|
||||
};
|
||||
struct mtmd_helper_gen_audio_inp {
|
||||
bool stream; // if true, output() must be called after each step_gen()
|
||||
llama_seq_id seq_id;
|
||||
|
||||
const char * prompt;
|
||||
@@ -189,6 +190,8 @@ struct mtmd_helper_gen_audio_inp {
|
||||
enum mtmd_helper_gen_audio_outtype out_type;
|
||||
};
|
||||
|
||||
MTMD_API struct mtmd_helper_gen_audio_inp mtmd_helper_gen_audio_inp_default(void);
|
||||
|
||||
MTMD_API mtmd_helper_gen_audio * mtmd_helper_gen_audio_init(
|
||||
struct llama_context * lctx,
|
||||
struct mtmd_context * mctx);
|
||||
@@ -217,6 +220,7 @@ MTMD_API int32_t mtmd_helper_gen_audio_step_gen(
|
||||
|
||||
// out_data valid until next get_output() or reset() call
|
||||
// out_n_samples (optional, can be NULL) receives the number of generated PCM samples
|
||||
// if inp->stream is true, each call returns only the audio produced since the previous call
|
||||
MTMD_API int32_t mtmd_helper_gen_audio_get_output(
|
||||
mtmd_helper_gen_audio * ctx,
|
||||
int32_t * out_sample_rate,
|
||||
|
||||
+1
-2
@@ -111,7 +111,7 @@ int main(int argc, char ** argv) {
|
||||
}
|
||||
|
||||
mtmd_helper::gen_audio gen(lctx, mctx.get());
|
||||
mtmd_helper_gen_audio_inp inp{};
|
||||
mtmd_helper_gen_audio_inp inp = mtmd_helper_gen_audio_inp_default();
|
||||
inp.seq_id = 0;
|
||||
inp.prompt = params.prompt.c_str();
|
||||
inp.prompt_len = params.prompt.size();
|
||||
@@ -119,7 +119,6 @@ int main(int argc, char ** argv) {
|
||||
inp.lang = params.tts_lang.c_str();
|
||||
inp.top_k = params.sampling.top_k;
|
||||
inp.top_p = params.sampling.top_p;
|
||||
inp.out_type = MTMD_HELPER_GEN_AUDIO_OUTTYPE_WAV;
|
||||
|
||||
//
|
||||
// stage 1: process prompt via backbone model, generate semantic representation
|
||||
|
||||
Reference in New Issue
Block a user