From 717dad5c8e9652cb58393de8287dd6668cb8c26d Mon Sep 17 00:00:00 2001 From: Xuan-Son Nguyen Date: Wed, 5 Aug 2026 13:34:52 +0200 Subject: [PATCH] mtmd: support multi-row batching for deepseek-ocr (#26154) * mtmd: support multi-row batching for deepseek-ocr * mtmd: weave deepseek-ocr rows in one shot instead of per row (#26615) --------- Co-authored-by: Saba Fallah --- tools/mtmd/models/deepseekocr.cpp | 51 +++++++++++++++++------------- tools/mtmd/models/deepseekocr2.cpp | 16 ++++++---- tools/mtmd/models/models.h | 3 +- 3 files changed, 41 insertions(+), 29 deletions(-) diff --git a/tools/mtmd/models/deepseekocr.cpp b/tools/mtmd/models/deepseekocr.cpp index b9fea3538..0ba5a4d2a 100644 --- a/tools/mtmd/models/deepseekocr.cpp +++ b/tools/mtmd/models/deepseekocr.cpp @@ -253,6 +253,9 @@ ggml_cgraph * clip_graph_deepseekocr::build() { bool is_overview = img.add_viewsep; int n_tiles_per_row = 0; + // number of separate "row" images batched together in this graph call + // (captured now, before n_batch below gets repurposed as the SAM/ViT batch size) + const int n_rows_batch = n_batch; // note: we expect either a batch of rows or a batch of overviews, but not a mix of both @@ -272,16 +275,18 @@ ggml_cgraph * clip_graph_deepseekocr::build() { GGML_ASSERT(img.ny() % img.nx() == 0); n_tiles_per_row = img.ny() / img.nx(); - // input shape: [tile_size, tile_size * n_tiles_per_row, 3] - // we want to reshape it to [tile_size, tile_size, 3, n_tiles_per_row] - inp_raw = ggml_reshape_4d(ctx0, inp_raw, img.nx(), img.nx(), n_tiles_per_row, 3); - inp_raw = ggml_cont(ctx0, ggml_permute(ctx0, inp_raw, 0, 1, 3, 2)); + // each entry is one "row" image of shape [tile_size, tile_size * n_tiles_per_row, 3]; + // merge the tile axis into the batch axis, giving a combined SAM input of shape + // [tile_size, tile_size, 3, n_tiles_per_row * n_rows_batch] (tile fast, row slow) + inp_raw = ggml_reshape_4d(ctx0, inp_raw, img.nx() * img.nx(), n_tiles_per_row, 3, n_rows_batch); + inp_raw = ggml_cont(ctx0, ggml_permute(ctx0, inp_raw, 0, 2, 1, 3)); + inp_raw = ggml_reshape_4d(ctx0, inp_raw, img.nx(), img.nx(), 3, n_tiles_per_row * n_rows_batch); } ggml_tensor * sam_out = build_sam(inp_raw); if (!is_overview) { - n_batch = n_tiles_per_row; + n_batch = n_tiles_per_row * n_rows_batch; } const int clip_n_patches = sam_out->ne[0] * sam_out->ne[1]; @@ -354,34 +359,36 @@ ggml_cgraph * clip_graph_deepseekocr::build() { const auto w = h; const auto n_dim = cur->ne[0]; - ggml_tensor * imgnl = ggml_repeat_4d(ctx0, model.image_newline, n_dim, 1, h, 1); - cur = ggml_reshape_3d(ctx0, cur, n_dim, w, h); - cur = ggml_reshape_2d(ctx0, ggml_concat(ctx0, cur, imgnl, 1), n_dim, (w + 1) * h); - cur = ggml_concat(ctx0, cur, model.view_seperator, 1); // (n_dim, h*(w+1) + 1) + ggml_tensor * imgnl = ggml_repeat_4d(ctx0, model.image_newline, n_dim, 1, h, n_batch); + cur = ggml_reshape_4d(ctx0, cur, n_dim, w, h, n_batch); + cur = ggml_reshape_3d(ctx0, ggml_concat(ctx0, cur, imgnl, 1), n_dim, (w + 1) * h, n_batch); + ggml_tensor * vs = ggml_repeat_4d(ctx0, model.view_seperator, n_dim, 1, n_batch, 1); + cur = ggml_concat(ctx0, cur, vs, 1); // (n_dim, h*(w+1) + 1, n_batch) } else { // tile row: interleave tiles within each row, add newline per row - const int grid_x = static_cast(std::sqrt(static_cast(clip_n_patches))); - const int grid_y = grid_x; - const auto n_dim = cur->ne[0]; + const int grid_x = static_cast(std::sqrt(static_cast(clip_n_patches))); + const int grid_y = grid_x; + const auto n_dim = cur->ne[0]; - // (n_dim, clip_n_patches, n_batch) -> (n_dim, grid_x, grid_y, n_batch) - cur = ggml_reshape_4d(ctx0, cur, n_dim, grid_x, grid_y, n_batch); + // merge n_dim into the grid_x axis, freeing the 4th axis for n_rows_batch + // (n_dim, clip_n_patches, n_tiles_per_row * n_rows_batch) -> (n_dim*grid_x, grid_y, n_tiles_per_row, n_rows_batch) + cur = ggml_reshape_4d(ctx0, cur, n_dim * grid_x, grid_y, n_tiles_per_row, n_rows_batch); // tiles: re-order from A.row0 A.row1 B.row0 B.row1 ... // to A.row0 B.row0 A.row1 B.row1 ... // then add nl: A.row0 B.row0 [nl] A.row1 B.row1 [nl] ... - // interleave tiles: (n_dim, grid_x, grid_y, n_batch) -> (n_dim, grid_x, n_batch, grid_y) - cur = ggml_cont(ctx0, ggml_permute(ctx0, cur, 0, 1, 3, 2)); + // interleave tiles: -> (n_dim*grid_x, n_tiles_per_row, grid_y, n_rows_batch) + cur = ggml_cont(ctx0, ggml_permute(ctx0, cur, 0, 2, 1, 3)); - // merge: (n_dim, grid_x, n_batch, grid_y) -> (n_dim, grid_x*n_batch, grid_y, 1) - cur = ggml_reshape_4d(ctx0, cur, n_dim, grid_x * n_batch, grid_y, 1); + // merge: -> (n_dim, grid_x*n_tiles_per_row, grid_y, n_rows_batch) + cur = ggml_reshape_4d(ctx0, cur, n_dim, grid_x * n_tiles_per_row, grid_y, n_rows_batch); - // append newline per row: (n_dim, grid_x*n_batch+1, grid_y, 1) - ggml_tensor * imgnl = ggml_repeat_4d(ctx0, model.image_newline, n_dim, 1, grid_y, 1); + // append newline per row: (n_dim, grid_x*n_tiles_per_row+1, grid_y, n_rows_batch) + ggml_tensor * imgnl = ggml_repeat_4d(ctx0, model.image_newline, n_dim, 1, grid_y, n_rows_batch); cur = ggml_concat(ctx0, cur, imgnl, 1); - // flatten: (n_dim, (grid_x*n_batch+1)*grid_y) - cur = ggml_reshape_2d(ctx0, cur, n_dim, (grid_x * n_batch + 1) * grid_y); + // flatten: (n_dim, (grid_x*n_tiles_per_row+1)*grid_y, n_rows_batch) + cur = ggml_reshape_3d(ctx0, cur, n_dim, (grid_x * n_tiles_per_row + 1) * grid_y, n_rows_batch); } cb(cur, "dsocr_output", -1); diff --git a/tools/mtmd/models/deepseekocr2.cpp b/tools/mtmd/models/deepseekocr2.cpp index 056bb8180..3e8b40941 100644 --- a/tools/mtmd/models/deepseekocr2.cpp +++ b/tools/mtmd/models/deepseekocr2.cpp @@ -14,8 +14,9 @@ ggml_cgraph * clip_graph_deepseekocr2::build() { { ggml_tensor * inp; - inp = ggml_reshape_2d(ctx0, sam_out, sam_out->ne[0] * sam_out->ne[1], sam_out->ne[2]); // H*W, C - inp = ggml_cont(ctx0, ggml_permute(ctx0, inp, 1, 0, 2, 3)); + // H*W, C, B + inp = ggml_reshape_3d(ctx0, sam_out, sam_out->ne[0] * sam_out->ne[1], sam_out->ne[2], sam_out->ne[3]); + inp = ggml_cont(ctx0, ggml_permute(ctx0, inp, 1, 0, 2, 3)); // C, H*W, B auto num_image_tokens = inp->ne[1]; // H*W GGML_ASSERT(num_image_tokens == 144 || num_image_tokens == 256); @@ -32,8 +33,10 @@ ggml_cgraph * clip_graph_deepseekocr2::build() { num_queries = 144; } - // (B, num_image_tokens + num_queries, C) - inp = ggml_concat(ctx0, inp, ggml_cast(ctx0, query_embed, inp->type), 1); + // repeat the query embedding per batch item, then append: (C, num_image_tokens + num_queries, B) + query_embed = ggml_cast(ctx0, query_embed, inp->type); + query_embed = ggml_repeat_4d(ctx0, query_embed, query_embed->ne[0], num_queries, inp->ne[2], 1); + inp = ggml_concat(ctx0, inp, query_embed, 1); auto seq_len = inp->ne[1]; @@ -57,7 +60,7 @@ ggml_cgraph * clip_graph_deepseekocr2::build() { /* learned_pos_embd */ nullptr, add_rope, vit_opts); cur = ggml_cont(ctx0, - ggml_view_2d(ctx0, cur, cur->ne[0], num_queries, cur->nb[1], + ggml_view_3d(ctx0, cur, cur->ne[0], num_queries, cur->ne[2], cur->nb[1], cur->nb[2], cur->nb[1] * (cur->ne[1] - num_queries))); // only take query tokens for output ggml_build_forward_expand(gf, cur); @@ -71,7 +74,8 @@ ggml_cgraph * clip_graph_deepseekocr2::build() { // view_seperator only after the global view if (img.add_viewsep) { - cur = ggml_concat(ctx0, cur, model.view_seperator, 1); // (n_dim, 257) + ggml_tensor * vs = ggml_repeat_4d(ctx0, model.view_seperator, model.view_seperator->ne[0], 1, cur->ne[2], 1); + cur = ggml_concat(ctx0, cur, vs, 1); // (n_dim, 257, n_batch) } cb(cur, "dsocr2_output", -1); diff --git a/tools/mtmd/models/models.h b/tools/mtmd/models/models.h index eb924972b..4ee4eb374 100644 --- a/tools/mtmd/models/models.h +++ b/tools/mtmd/models/models.h @@ -138,12 +138,13 @@ struct clip_graph_deepseekocr : clip_graph { clip_graph_deepseekocr(clip_ctx * ctx, const clip_image_f32 & img) : clip_graph(ctx, img) {} ggml_cgraph * build() override; ggml_tensor * build_sam(ggml_tensor * inp); // build the SAM model - // bool support_batch() const override { return true; } // TODO: support batch for DeepSeek-OCR v1 + bool support_batch() const override { return true; } }; struct clip_graph_deepseekocr2 : clip_graph_deepseekocr { clip_graph_deepseekocr2(clip_ctx * ctx, const clip_image_f32 & img) : clip_graph_deepseekocr(ctx, img) {} ggml_cgraph * build() override; // reuses build_sam() from base + bool support_batch() const override { return true; } }; struct clip_graph_conformer : clip_graph {