mtmd: cap max_image to ubatch for non_causal models (#29773)

This commit is contained in:
Xuan-Son Nguyen
2026-10-01 11:55:11 +02:00
committed by GitHub
parent 5503b04b05
commit 552f18f912
7 changed files with 74 additions and 12 deletions
+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());