mtmd: cap max_image to ubatch for non_causal models

This commit is contained in:
Xuan Son Nguyen
2026-10-01 00:34:54 +02:00
parent ca2e2037b6
commit 808c5ee26b
7 changed files with 74 additions and 12 deletions
+4 -1
View File
@@ -210,8 +210,11 @@ struct clip_hparams {
void set_warmup_n_tokens(int n_tokens) {
int n_tok_per_side = static_cast<int>(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<int>(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<int32_t> global_attn_indices() const {
+12
View File
@@ -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;
+3
View File
@@ -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
+11
View File
@@ -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;
+14 -6
View File
@@ -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<ggml_backend_dev_t, size_t> 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<ggml_backend_dev_t, size_t> 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};
}
}
+7 -1
View File
@@ -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<ggml_backend_dev_t, size_t> mtmd_get_memory_usage(
struct mtmd_memory_usage {
std::map<ggml_backend_dev_t, size_t> 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
+23 -4
View File
@@ -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());