diff --git a/tools/mtmd/mtmd-helper-gen.cpp b/tools/mtmd/mtmd-helper-gen.cpp index dcd802c6ad..e98e9dfeeb 100644 --- a/tools/mtmd/mtmd-helper-gen.cpp +++ b/tools/mtmd/mtmd-helper-gen.cpp @@ -83,7 +83,7 @@ public: // (e.g. continuous/diffusion models); such pipelines read whatever they need // directly off h_state_in instead virtual int32_t step(llama_token sampled, const float * h_state_in, const float ** h_state_out) = 0; - virtual int32_t get_output(int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len) = 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; @@ -255,12 +255,15 @@ public: return 0; } - int32_t get_output(int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len) override { + int32_t get_output(int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len, int64_t * out_n_samples) override { if (!flush_c2w()) { return 1; } *out_sample_rate = info.sample_rate; + 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(); @@ -445,9 +448,9 @@ int32_t mtmd_helper_gen_audio_step(mtmd_helper_gen_audio * ctx, llama_token samp } int32_t mtmd_helper_gen_audio_get_output(mtmd_helper_gen_audio * ctx, int32_t * out_sample_rate, - const char ** out_data, size_t * out_data_len) { + const char ** out_data, size_t * out_data_len, int64_t * out_n_samples) { if (!ctx->pipeline) { return 1; } - return ctx->pipeline->get_output(out_sample_rate, out_data, out_data_len); + return ctx->pipeline->get_output(out_sample_rate, out_data, out_data_len, out_n_samples); } diff --git a/tools/mtmd/mtmd-helper.h b/tools/mtmd/mtmd-helper.h index f7fd722335..c065039f8f 100644 --- a/tools/mtmd/mtmd-helper.h +++ b/tools/mtmd/mtmd-helper.h @@ -204,11 +204,14 @@ MTMD_API int32_t mtmd_helper_gen_audio_step( const float ** h_state_out); // out_data valid until next get_output() or reset() call +// out_n_samples (optional, can be NULL) receives the number of generated PCM samples, +// which combined with out_sample_rate gives the output audio duration MTMD_API int32_t mtmd_helper_gen_audio_get_output( mtmd_helper_gen_audio * ctx, int32_t * out_sample_rate, const char ** out_data, - size_t * out_data_len); + size_t * out_data_len, + int64_t * out_n_samples); #ifdef __cplusplus } // extern "C" @@ -247,8 +250,8 @@ struct gen_audio { int32_t step(llama_token sampled, const float * h_state, const float ** h_state_out) { return mtmd_helper_gen_audio_step(ctx.get(), sampled, h_state, h_state_out); } - int32_t get_output(int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len) { - return mtmd_helper_gen_audio_get_output(ctx.get(), out_sample_rate, out_data, out_data_len); + int32_t get_output(int32_t * out_sample_rate, const char ** out_data, size_t * out_data_len, int64_t * out_n_samples = nullptr) { + return mtmd_helper_gen_audio_get_output(ctx.get(), out_sample_rate, out_data, out_data_len, out_n_samples); } }; diff --git a/tools/tts/tts.cpp b/tools/tts/tts.cpp index cc2bca6a40..23e797dc12 100644 --- a/tools/tts/tts.cpp +++ b/tools/tts/tts.cpp @@ -113,7 +113,6 @@ int main(int argc, char ** argv) { LOG_ERR("set_input failed\n"); return 1; } - LOG_INF("prompt eval took %.2f seconds\n", (ggml_time_us() - t_prompt_start_us) / 1e6); // codec_0 (backbone) EOS token: ordinary LLM sampling concern, kept out of the // model-agnostic audio-generation helper @@ -139,6 +138,7 @@ int main(int argc, char ** argv) { const float * h_state = llama_get_embeddings_ith(lctx, -1); tts_timings timings; + const int64_t t_gen_start_us = ggml_time_us(); for (; n_frames < max_new && sampled != codec_eos_tok; n_frames++) { const float * h_next = nullptr; @@ -150,16 +150,24 @@ int main(int argc, char ** argv) { sampled = sample_codec0(); timings.report(n_frames + 1); } + const double t_gen_s = (ggml_time_us() - t_gen_start_us) / 1e6; int32_t sample_rate = 0; const char * data = nullptr; size_t data_len = 0; - if (gen.get_output(&sample_rate, &data, &data_len) != 0) { + int64_t n_samples = 0; + if (gen.get_output(&sample_rate, &data, &data_len, &n_samples) != 0) { LOG_ERR("get_output failed\n"); return 1; } LOG_INF("generated %d frames, %zu bytes of WAV audio (%d Hz)\n", n_frames, data_len, sample_rate); + + const double t_prompt_s = (t_gen_start_us - t_prompt_start_us) / 1e6; + const double t_total_s = t_prompt_s + t_gen_s; + const double audio_s = sample_rate > 0 ? (double) n_samples / sample_rate : 0.0; + LOG_INF("timings: prompt eval %.2fs + generation %.2fs = total %.2fs\n", t_prompt_s, t_gen_s, t_total_s); + LOG_INF(" output audio = %.2fs (audio time = %.2fx process time)\n", audio_s, t_total_s > 0 ? audio_s / t_total_s : 0.0); FILE * f = fopen(params.out_file.c_str(), "wb"); if (!f) { LOG_ERR("failed to open %s\n", params.out_file.c_str());