diff --git a/tools/mtmd/clip.cpp b/tools/mtmd/clip.cpp index b9dd5e8452..5969bf7599 100644 --- a/tools/mtmd/clip.cpp +++ b/tools/mtmd/clip.cpp @@ -3898,8 +3898,9 @@ struct clip_init_result clip_init(const char * fname, struct clip_context_params struct clip_cap clip_get_cap(const char * fname) { clip_cap res; clip_model_loader loader(fname, /* skip_tensors= */ true); - res.has_vision = loader.has_vision; - res.has_audio = loader.has_audio; + res.has_vision = loader.has_vision; + res.has_audio = loader.has_audio; + res.has_gen_audio = loader.has_gen_audio; return res; } diff --git a/tools/mtmd/clip.h b/tools/mtmd/clip.h index a5b7137752..56c200482c 100644 --- a/tools/mtmd/clip.h +++ b/tools/mtmd/clip.h @@ -134,5 +134,6 @@ std::map clip_get_mem_usage(const struct clip_ctx * struct clip_cap { bool has_vision; bool has_audio; + bool has_gen_audio; }; struct clip_cap clip_get_cap(const char * fname); diff --git a/tools/mtmd/mtmd.cpp b/tools/mtmd/mtmd.cpp index 4063d28e07..fdef2079da 100644 --- a/tools/mtmd/mtmd.cpp +++ b/tools/mtmd/mtmd.cpp @@ -2510,10 +2510,11 @@ struct mtmd_caps mtmd_get_cap_from_file(const char * fname) { mtmd_caps cap; cap.inp_audio = tmp.has_audio; cap.inp_vision = tmp.has_vision; + cap.gen_audio = tmp.has_gen_audio; return cap; } catch (const std::exception & e) { LOG_ERR("%s: failed to get capabilities from file '%s': %s\n", __func__, fname, e.what()); - return mtmd_caps{ false, false }; + return mtmd_caps{ false, false, false }; } } diff --git a/tools/mtmd/mtmd.h b/tools/mtmd/mtmd.h index 78587f3fed..084507bec6 100644 --- a/tools/mtmd/mtmd.h +++ b/tools/mtmd/mtmd.h @@ -337,6 +337,7 @@ MTMD_API void mtmd_log_set(ggml_log_callback log_callback, void * user_data); struct mtmd_caps { bool inp_vision; bool inp_audio; + bool gen_audio; }; MTMD_API struct mtmd_caps mtmd_get_cap_from_file(const char * mmproj_fname); diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index 9f6490e08e..a5d1c2732a 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -42,8 +42,11 @@ constexpr int HTTP_POLLING_SECONDS = 1; static common_speculative_output_limits server_output_limits(const common_params & params) { if (!params.mmproj.path.empty()) { - // gen-audio (TTS) capability isn't known until the mmproj loads, size generously - return { params.n_batch, params.n_batch }; + const auto mcaps = mtmd_get_cap_from_file(params.mmproj.path.c_str()); + if (mcaps.gen_audio) { + // some TTS models decode with embeddings output for the whole batch + return { params.n_batch, params.n_batch }; + } } if (params.embedding || diff --git a/tools/server/server-models.cpp b/tools/server/server-models.cpp index 93e940951e..a292766f12 100644 --- a/tools/server/server-models.cpp +++ b/tools/server/server-models.cpp @@ -624,7 +624,7 @@ void server_models::load_models() { /* progress */ {}, /* exit_code */ 0, /* stop_timeout */ DEFAULT_STOP_TIMEOUT, - /* multimodal */ mtmd_caps{false, false}, + /* multimodal */ mtmd_caps{false, false, false}, // /* need_download */ false, }; add_model(std::move(meta)); @@ -797,7 +797,7 @@ void server_models::load_models() { /* progress */ {}, /* exit_code */ 0, /* stop_timeout */ DEFAULT_STOP_TIMEOUT, - /* multimodal */ mtmd_caps{false, false}, + /* multimodal */ mtmd_caps{false, false, false}, // /* need_download */ false, }; add_model(std::move(meta));