mtmd: add cohere2 vision support (#30062)

* cohere2 vision model

* address comments

* remove redundant mapping

* follow existing patterns

* fused linear_1
This commit is contained in:
Terrence Zhao
2026-10-07 20:07:50 +02:00
committed by GitHub
parent d6cf9acb25
commit 50a6c5cf7c
10 changed files with 146 additions and 2 deletions
+1
View File
@@ -303,6 +303,7 @@ MMPROJ_MODEL_MAP: dict[str, str] = {
"AudioFlamingo3ForConditionalGeneration": "ultravox",
"ClefModel": "clef",
"CogVLMForCausalLM": "cogvlm",
"Cohere2VisionForConditionalGeneration": "command_r",
"PplxDeciderModel": "pplx_decider",
"DeepseekOCR2ForCausalLM": "deepseek",
"DeepseekOCRForCausalLM": "deepseek",
+27 -2
View File
@@ -1,14 +1,14 @@
from __future__ import annotations
import re
from typing import Iterable, TYPE_CHECKING
from typing import Callable, Iterable, TYPE_CHECKING
import torch
if TYPE_CHECKING:
from torch import Tensor
from .base import ModelBase, TextModel, gguf, logger
from .base import MmprojModel, ModelBase, TextModel, gguf, logger
@ModelBase.register("CohereForCausalLM")
@@ -180,3 +180,28 @@ class Cohere2MoeModel(TextModel):
experts = [k for d in self._experts for k in d.keys()]
if len(experts) > 0:
raise ValueError(f"Unprocessed experts: {experts}")
@ModelBase.register("Cohere2VisionForConditionalGeneration")
# [TAG_HF_EXAMPLE_GATED] CohereLabs/command-a-vision-07-2025 is gated
@ModelBase.example("CohereLabs/command-a-plus-05-2026-bf16")
class Cohere2VisionModel(MmprojModel):
def set_gguf_parameters(self):
super().set_gguf_parameters()
self.gguf_writer.add_clip_projector_type(gguf.VisionProjectorType.COHERE2V)
self.gguf_writer.add_vision_attention_layernorm_eps(self.hparams["layer_norm_eps"])
self.gguf_writer.add_vision_projector_scale_factor(self.global_config["downsample_factor"])
self.gguf_writer.add_vision_preproc_max_tiles(self.preprocessor_config["max_patches"])
self.gguf_writer.add_vision_use_gelu(True)
def tensor_force_quant(self, name, new_name, bid, n_dims):
if ".embeddings." in name:
return gguf.GGMLQuantizationType.F32
return super().tensor_force_quant(name, new_name, bid, n_dims)
@classmethod
def filter_tensors(cls, item: tuple[str, Callable[[], Tensor]]) -> tuple[str, Callable[[], Tensor]] | None:
name, gen = item
if not name.startswith(("model.vision_tower.", "model.multi_modal_projector.")):
return None
return super().filter_tensors((name, gen))
+1
View File
@@ -6130,6 +6130,7 @@ class VisionProjectorType:
MIMO_AUDIO = "mimo_audio"
GRANITE4_VISION = "granite4_vision"
MUSE_GLIMMER = "muse-glimmer"
COHERE2V = "cohere2v"
# Items here are (block size, type size)
+1
View File
@@ -1622,6 +1622,7 @@ class TensorNameMap:
MODEL_TENSOR.V_MMPROJ: (
"aligner.w{bid}", # deepseek4v (w1 -> mm.1, w2 -> mm.2)
"multi_modal_projector.linear_{bid}",
"model.multi_modal_projector.linear_{bid}", # cohere2v
"mm_projector.proj.linear_{bid}", # Kimi-K2.5
"visual.merger.mlp.{bid}", # qwen2vl
"mlp_AR.linear_{bid}", # PaddleOCR-VL
+2
View File
@@ -507,6 +507,7 @@ enum projector_type {
PROJECTOR_TYPE_POCKETTTS_SPKENC,
PROJECTOR_TYPE_POCKETTTS_GEN,
PROJECTOR_TYPE_MUSE_GLIMMER,
PROJECTOR_TYPE_COHERE2V,
PROJECTOR_TYPE_UNKNOWN,
};
@@ -574,6 +575,7 @@ static std::map<projector_type, std::string> PROJECTOR_TYPE_NAMES = {
{ PROJECTOR_TYPE_POCKETTTS_SPKENC, "pockettts_spkenc"},
{ PROJECTOR_TYPE_POCKETTTS_GEN, "pockettts_gen"},
{ PROJECTOR_TYPE_MUSE_GLIMMER, "muse-glimmer"},
{ PROJECTOR_TYPE_COHERE2V, "cohere2v"},
};
static projector_type clip_projector_type_from_string(const std::string & str) {
+20
View File
@@ -937,6 +937,7 @@ static std::unique_ptr<clip_graph> clip_get_graph_builder(clip_ctx * ctx, const
switch (ctx->proj_type()) {
case PROJECTOR_TYPE_GEMMA3:
case PROJECTOR_TYPE_IDEFICS3:
case PROJECTOR_TYPE_COHERE2V:
case PROJECTOR_TYPE_LFM2:
case PROJECTOR_TYPE_JANUS_PRO:
case PROJECTOR_TYPE_PHI4:
@@ -1514,6 +1515,15 @@ struct clip_model_loader {
get_u32(KEY_PREPROC_IMAGE_SIZE, hparams.image_longest_edge, false);
hparams.set_limit_image_tokens();
} break;
case PROJECTOR_TYPE_COHERE2V:
{
hparams.image_pad_rf = PAD_NONE;
get_u32(KEY_PROJ_SCALE_FACTOR, hparams.n_merge);
get_u32(KEY_PREPROC_MAX_TILES, hparams.preproc_max_tiles);
if (hparams.preproc_max_tiles <= 0 || hparams.preproc_max_tiles > 256) {
throw std::runtime_error(string_format("%s: preproc_max_tiles (%d) must be in range [1, 256]\n", __func__, hparams.preproc_max_tiles));
}
} break;
case PROJECTOR_TYPE_LFM2:
{
hparams.image_resize_algo = RESIZE_ALGO_BILINEAR;
@@ -2770,6 +2780,13 @@ struct clip_model_loader {
{
model.mm_fc_w = get_tensor(string_format(TN_MM_PROJECTOR, "weight"));
} break;
case PROJECTOR_TYPE_COHERE2V:
{
model.mm_1_w = get_tensor(string_format(TN_LLAVA_PROJ, 1, "weight"));
model.mm_1_b = get_tensor(string_format(TN_LLAVA_PROJ, 1, "bias"));
model.mm_2_w = get_tensor(string_format(TN_LLAVA_PROJ, 2, "weight"));
model.mm_2_b = get_tensor(string_format(TN_LLAVA_PROJ, 2, "bias"));
} break;
case PROJECTOR_TYPE_LFM2:
{
model.mm_input_norm_w = get_tensor(TN_MM_INP_NORM, false);
@@ -4233,6 +4250,7 @@ int clip_n_output_tokens(const clip_ctx * ctx, const clip_image_f32 * img) {
case PROJECTOR_TYPE_GEMMA4V:
case PROJECTOR_TYPE_GEMMA4UV:
case PROJECTOR_TYPE_IDEFICS3:
case PROJECTOR_TYPE_COHERE2V:
case PROJECTOR_TYPE_INTERNVL:
case PROJECTOR_TYPE_NEMOTRON_V2_VL:
case PROJECTOR_TYPE_LLAMA4:
@@ -5321,6 +5339,7 @@ bool clip_encode(struct clip_ctx * ctx, struct clip_encode_params * params) {
case PROJECTOR_TYPE_GEMMA3:
case PROJECTOR_TYPE_GEMMA3NV:
case PROJECTOR_TYPE_IDEFICS3:
case PROJECTOR_TYPE_COHERE2V:
case PROJECTOR_TYPE_INTERNVL:
case PROJECTOR_TYPE_NEMOTRON_V2_VL:
case PROJECTOR_TYPE_QWEN2A:
@@ -6055,6 +6074,7 @@ int clip_n_mmproj_embd(const struct clip_ctx * ctx) {
case PROJECTOR_TYPE_KIMIK25:
case PROJECTOR_TYPE_YASA2:
case PROJECTOR_TYPE_DEEPSEEK4V:
case PROJECTOR_TYPE_COHERE2V:
return ctx->model.mm_2_w->ne[1];
case PROJECTOR_TYPE_HUNYUANVL:
return ctx->model.mm_model_proj->ne[1];
+10
View File
@@ -45,6 +45,16 @@ ggml_cgraph * clip_graph_siglip::build() {
cur = build_patch_merge_permute(cur, scale_factor);
cur = build_mm(model.mm_fc_w, cur);
} else if (proj_type == PROJECTOR_TYPE_COHERE2V) {
// tiles are square, so the pixel shuffle is the same as Idefics3
cur = build_patch_merge_permute(cur, model.hparams.n_merge);
cur = build_mm(model.mm_1_w, cur);
cur = ggml_add(ctx0, cur, model.mm_1_b);
// linear_1 output is [x, gate], HF computes silu(gate) * x
cur = ggml_swiglu_swapped(ctx0, cur);
cur = build_mm(model.mm_2_w, cur);
cur = ggml_add(ctx0, cur, model.mm_2_b);
} else if (proj_type == PROJECTOR_TYPE_LFM2) {
// pixel unshuffle block
const int scale_factor = model.hparams.n_merge;
+66
View File
@@ -1142,6 +1142,72 @@ mtmd_image_preproc_out mtmd_image_preprocessor_idefics3::preprocess(const clip_i
return output;
}
//
// mtmd_image_preprocessor_cohere2v
//
mtmd_image_preproc_out mtmd_image_preprocessor_cohere2v::preprocess(const clip_image_u8 & img) const {
const auto inst = get_slice_instructions(img.get_size());
auto sliced = slice_image(img, inst);
mtmd_image_preproc_out output;
if (sliced.slices.empty()) {
output.append_overview(hparams, sliced.overview, true);
return output;
}
// slices first, then thumbnail
output.append(hparams, sliced.slices, true);
output.append_overview(hparams, sliced.overview, true);
output.grid_x = inst.grid_size.width;
output.grid_y = inst.grid_size.height;
return output;
}
mtmd_image_preprocessor_llava_uhd::slice_instructions mtmd_image_preprocessor_cohere2v::get_slice_instructions(const clip_image_size & original_size) const {
const int tile = hparams.image_size;
// pick the grid with the least upscale; if all grids need downscale, pick the one with the least downscale
// grids are visited by tile count, then by width, same order as HF for ties
double best_down = -1.0;
double best_up = std::numeric_limits<double>::max();
clip_image_size grid_down = { 1, 1 };
clip_image_size grid_up = { 0, 0 };
for (int n = 1; n <= hparams.preproc_max_tiles; n++) {
for (int w = 1; w <= n; w++) {
if (n % w != 0) {
continue;
}
const clip_image_size g = { w, n / w };
const double scale = std::min(
(double) (g.width * tile) / original_size.width,
(double) (g.height * tile) / original_size.height);
if (scale < 1.0) {
if (scale > best_down) {
best_down = scale;
grid_down = g;
}
} else if (scale < best_up) {
best_up = scale;
grid_up = g;
}
}
}
const clip_image_size grid = grid_up.width > 0 ? grid_up : grid_down;
slice_instructions inst;
inst.overview_size = { tile, tile };
inst.refined_size = { tile * grid.width, tile * grid.height };
inst.grid_size = grid;
if (grid.width * grid.height > 1) {
for (int y = 0; y < grid.height; y++) {
for (int x = 0; x < grid.width; x++) {
inst.slices.push_back({ x * tile, y * tile, { tile, tile } });
}
}
}
return inst;
}
//
// mtmd_image_preprocessor_internvl
//
+8
View File
@@ -191,6 +191,14 @@ struct mtmd_image_preprocessor_internvl : mtmd_image_preprocessor_llava_uhd {
mtmd_image_preproc_out preprocess(const clip_image_u8 & img) const override;
};
// stretch the image to a grid of tiles, add a thumbnail if there is more than 1 tile
// ref: https://github.com/huggingface/transformers/blob/main/src/transformers/models/cohere2_vision/image_processing_cohere2_vision.py
struct mtmd_image_preprocessor_cohere2v : mtmd_image_preprocessor_llava_uhd {
mtmd_image_preprocessor_cohere2v(const clip_ctx * ctx) : mtmd_image_preprocessor_llava_uhd(ctx) {}
mtmd_image_preproc_out preprocess(const clip_image_u8 & img) const override;
slice_instructions get_slice_instructions(const clip_image_size & original_size) const override;
};
// DeepSeek-OCR (v1/v2) global view + optional local tile grid
struct mtmd_image_preprocessor_deepseekocr : mtmd_image_preprocessor {
mtmd_image_preprocessor_deepseekocr(const clip_ctx * ctx)
+10
View File
@@ -798,6 +798,16 @@ struct mtmd_context {
image_preproc = std::make_unique<mtmd_image_preprocessor_internvl>(ctx_v);
ov_img_first = false;
} break;
case PROJECTOR_TYPE_COHERE2V:
{
// <|START_OF_IMG|> (tile embeddings) <|IMG_LINE_BREAK|> ... <|END_OF_IMG|>
img_beg = "<|START_OF_IMG|>";
img_end = "<|END_OF_IMG|>";
tok_sli_img_end = {lookup_token("<|IMG_LINE_BREAK|>")};
tok_ov_img_end = tok_sli_img_end;
ov_img_first = false;
image_preproc = std::make_unique<mtmd_image_preprocessor_cohere2v>(ctx_v);
} break;
case PROJECTOR_TYPE_KIMIVL:
{
// <|media_start|> ... (image embeddings) ... <|media_end|>