From 808c5ee26ba835b632dd01d6bd58be6cfd99d409 Mon Sep 17 00:00:00 2001 From: Xuan Son Nguyen Date: Thu, 1 Oct 2026 00:34:54 +0200 Subject: [PATCH] mtmd: cap max_image to ubatch for non_causal models --- tools/mtmd/clip-model.h | 5 ++++- tools/mtmd/clip.cpp | 12 ++++++++++++ tools/mtmd/clip.h | 3 +++ tools/mtmd/mtmd-cli.cpp | 11 +++++++++++ tools/mtmd/mtmd.cpp | 20 ++++++++++++++------ tools/mtmd/mtmd.h | 8 +++++++- tools/server/server-context.cpp | 27 +++++++++++++++++++++++---- 7 files changed, 74 insertions(+), 12 deletions(-) diff --git a/tools/mtmd/clip-model.h b/tools/mtmd/clip-model.h index 77248ca764..33f679fc0c 100644 --- a/tools/mtmd/clip-model.h +++ b/tools/mtmd/clip-model.h @@ -210,8 +210,11 @@ struct clip_hparams { void set_warmup_n_tokens(int n_tokens) { int n_tok_per_side = static_cast(std::sqrt(n_tokens)); GGML_ASSERT(n_tok_per_side * n_tok_per_side == n_tokens && "n_tokens must be n*n"); + // do not warmup with more tokens than the max allowed + if (custom_image_max_tokens > 0 && n_tokens > custom_image_max_tokens) { + n_tok_per_side = std::max(1, static_cast(std::sqrt(custom_image_max_tokens))); + } warmup_image_size = n_tok_per_side * patch_size * n_merge; - // TODO: support warmup size for custom token numbers } // sam vit deepseek-ocr std::vector global_attn_indices() const { diff --git a/tools/mtmd/clip.cpp b/tools/mtmd/clip.cpp index 572a7b9870..50fb6c408a 100644 --- a/tools/mtmd/clip.cpp +++ b/tools/mtmd/clip.cpp @@ -4061,6 +4061,18 @@ struct clip_cap clip_get_cap(const char * fname) { return res; } +int clip_get_image_max_tokens(const clip_ctx * ctx) { + const auto & hparams = ctx->model.hparams; + if (ctx->proj_type() == PROJECTOR_TYPE_DEEPSEEK4V) { + return hparams.dsv4_max_n_token; + } + if (hparams.image_max_pixels <= 0) { + return -1; + } + const int patch_area = hparams.patch_size * hparams.patch_size * hparams.n_merge * hparams.n_merge; + return hparams.image_max_pixels / patch_area; +} + void clip_free(clip_ctx * ctx) { if (ctx == nullptr) { return; diff --git a/tools/mtmd/clip.h b/tools/mtmd/clip.h index e07f258156..9e12702ceb 100644 --- a/tools/mtmd/clip.h +++ b/tools/mtmd/clip.h @@ -68,6 +68,9 @@ struct clip_init_result { struct clip_init_result clip_init(const char * fname, struct clip_context_params ctx_params); +// max number of output tokens per image, -1 if not dynamic size +int clip_get_image_max_tokens(const struct clip_ctx * ctx); + void clip_free(struct clip_ctx * ctx); // TODO: should be enum, not string diff --git a/tools/mtmd/mtmd-cli.cpp b/tools/mtmd/mtmd-cli.cpp index 6fe058fd8c..0c364d3fd7 100644 --- a/tools/mtmd/mtmd-cli.cpp +++ b/tools/mtmd/mtmd-cli.cpp @@ -162,6 +162,17 @@ struct mtmd_cli_context { mparams.warmup = params.warmup; mparams.image_min_tokens = params.image_min_tokens; mparams.image_max_tokens = params.image_max_tokens; + { + // non-causal models need the whole image in one ubatch + const int n_ubatch = llama_n_ubatch(lctx); + auto mem = mtmd_get_memory_usage(clip_path, mparams); + if (mem.use_non_causal && mem.image_max_tokens > n_ubatch) { + LOG_WRN("%s: cap image_max_tokens (original=%d) to n_ubatch (%d) because model needs non-causal attention on image\n", __func__, mem.image_max_tokens, n_ubatch); + LOG_WRN("%s: increase n_ubatch (-ub) to increase vision token budget\n", __func__); + mparams.image_max_tokens = n_ubatch; + mparams.image_min_tokens = std::min(mparams.image_min_tokens, n_ubatch); + } + } if (std::getenv("MTMD_DEBUG_GRAPH") != nullptr) { mparams.cb_eval_user_data = &cb_data; mparams.cb_eval = common_debug_cb_eval; diff --git a/tools/mtmd/mtmd.cpp b/tools/mtmd/mtmd.cpp index 2adf342083..9e7b1cc801 100644 --- a/tools/mtmd/mtmd.cpp +++ b/tools/mtmd/mtmd.cpp @@ -2182,8 +2182,12 @@ bool mtmd_decode_use_non_causal(const mtmd_context * ctx, const mtmd_input_chunk } switch (proj_type) { case PROJECTOR_TYPE_GEMMA4V: - // E2B (n_embd = 1536) and E4B (n_embd = 2560) always use causal - return ctx->n_embd_text != 1536 && ctx->n_embd_text != 2560; + { + // E2B (n_embd = 1536) and E4B (n_embd = 2560) always use causal + // note: use mmproj n_embd, because text model may not be provided (e.g. mtmd_get_memory_usage) + const int n_embd = clip_n_mmproj_embd(ctx->ctx_v); + return n_embd != 1536 && n_embd != 2560; + } case PROJECTOR_TYPE_GEMMA4UV: case PROJECTOR_TYPE_GEMMA3: case PROJECTOR_TYPE_DEEPSEEK4V: @@ -2708,8 +2712,8 @@ static void stub_log_callback(enum ggml_log_level, const char *, void *) { // do nothing } -std::map mtmd_get_memory_usage(const char * mmproj_fname, - struct mtmd_context_params ctx_params) { +mtmd_memory_usage mtmd_get_memory_usage(const char * mmproj_fname, + struct mtmd_context_params ctx_params) { mtmd::context_ptr ctx; auto saved_log_callback = g_logger_state.log_callback; auto saved_log_user_data = g_logger_state.log_callback_user_data; @@ -2732,10 +2736,14 @@ std::map mtmd_get_memory_usage(const char * mmproj_f if (ctx->ctx_a) { merge(ctx->ctx_a); } - return total_mem; + mtmd_memory_usage res; + res.backend_mem_usage = std::move(total_mem); + res.image_max_tokens = ctx->ctx_v ? clip_get_image_max_tokens(ctx->ctx_v) : -1; + res.use_non_causal = ctx->ctx_v ? mtmd_decode_use_non_causal(ctx.get(), nullptr) : false; + return res; } catch (const std::exception & e) { mtmd_log_set(saved_log_callback, saved_log_user_data); // restore log callback LOG_ERR("%s: error: %s\n", __func__, e.what()); - return {}; + return {{}, -1, false}; } } diff --git a/tools/mtmd/mtmd.h b/tools/mtmd/mtmd.h index c2de26eeee..bcd76a9d71 100644 --- a/tools/mtmd/mtmd.h +++ b/tools/mtmd/mtmd.h @@ -449,7 +449,13 @@ MTMD_API mtmd_input_chunks * mtmd_test_create_input_chunks(void); // Get memory usage of the current model in bytes, per backend device // Note: this is an unstable API, used internally by fit_params; it WILL be removed or changed without deprecation #ifdef __cplusplus -MTMD_API std::map mtmd_get_memory_usage( +struct mtmd_memory_usage { + std::map backend_mem_usage; + // for models that use non-causal attention, max_tokens must not exceed n_ubatch of llama_context + int image_max_tokens; + bool use_non_causal; +}; +MTMD_API struct mtmd_memory_usage mtmd_get_memory_usage( const char * mmproj_fname, struct mtmd_context_params ctx_params); #endif diff --git a/tools/server/server-context.cpp b/tools/server/server-context.cpp index fbfcbe5125..470fbd9775 100644 --- a/tools/server/server-context.cpp +++ b/tools/server/server-context.cpp @@ -1041,11 +1041,19 @@ private: mparams.progress_callback_user_data = &load_progress_mmproj; } - // optionally get the memory usage of mmproj - if (has_mmproj && params_base.fit_params) { + // get the memory usage of mmproj, also used to check image_max_tokens against n_ubatch + mtmd_memory_usage mmproj_usage = {{}, -1, false}; + int64_t mmproj_usage_t_us = 0; + if (has_mmproj) { int64_t t_start = ggml_time_us(); - auto mmproj_mem = mtmd_get_memory_usage(mmproj_path.c_str(), mparams); - int64_t t_elapsed = ggml_time_us() - t_start; + mmproj_usage = mtmd_get_memory_usage(mmproj_path.c_str(), mparams); + mmproj_usage_t_us = ggml_time_us() - t_start; + } + + // optionally fit mmproj memory usage + if (has_mmproj && params_base.fit_params) { + const auto & mmproj_mem = mmproj_usage.backend_mem_usage; + const int64_t t_elapsed = mmproj_usage_t_us; if (!mmproj_mem.empty()) { size_t total = 0; for (auto & [dev, size] : mmproj_mem) { @@ -1140,6 +1148,17 @@ private: mtmd_helper_log_set(common_log_default_callback, nullptr); } + // non-causal models need the whole image in one ubatch + { + const int n_ubatch = llama_n_ubatch(ctx_tgt); + if (mmproj_usage.use_non_causal && mmproj_usage.image_max_tokens > n_ubatch) { + SRV_WRN("cap image_max_tokens (original=%d) to n_ubatch (%d) because model needs non-causal attention on image\n", mmproj_usage.image_max_tokens, n_ubatch); + SRV_WRN("%s\n", "increase n_ubatch (-ub) to increase vision token budget"); + mparams.image_max_tokens = n_ubatch; + mparams.image_min_tokens = std::min(mparams.image_min_tokens, n_ubatch); + } + } + mctx = mtmd_init_from_file(mmproj_path.c_str(), model_tgt, mparams); if (mctx == nullptr) { SRV_ERR("failed to load multimodal model, '%s'\n", mmproj_path.c_str());