[model] Support MiniCPM-V 4.5 (#9610)

Signed-off-by: tc-mb <caitianchi@modelbest.cn>
Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
This commit is contained in:
tc-mb
2026-01-31 23:37:36 -08:00
committed by GitHub
co-authored by Xinyuan Tong
parent 2c036f1eb1
commit 4d28cda007
+389 -5
View File
@@ -21,7 +21,9 @@
# limitations under the License. # limitations under the License.
"""Inference-only MiniCPM-V model compatible with HuggingFace weights.""" """Inference-only MiniCPM-V model compatible with HuggingFace weights."""
import types
from functools import partial from functools import partial
from itertools import chain
from typing import ( from typing import (
Any, Any,
Callable, Callable,
@@ -56,6 +58,7 @@ from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.idefics2 import Idefics2VisionTransformer from sglang.srt.models.idefics2 import Idefics2VisionTransformer
from sglang.srt.models.llama import LlamaConfig, LlamaForCausalLM from sglang.srt.models.llama import LlamaConfig, LlamaForCausalLM
from sglang.srt.models.qwen2 import Qwen2Config, Qwen2ForCausalLM from sglang.srt.models.qwen2 import Qwen2Config, Qwen2ForCausalLM
from sglang.srt.models.qwen3 import Qwen3Config, Qwen3ForCausalLM
from sglang.srt.utils import add_prefix, flatten_nested_list from sglang.srt.utils import add_prefix, flatten_nested_list
RawImageType = Union[Image.Image, torch.Tensor] RawImageType = Union[Image.Image, torch.Tensor]
@@ -356,6 +359,218 @@ class Resampler2_5(BaseResampler):
return x return x
class Resampler4_5(BaseResampler):
def __init__(
self,
num_queries: int,
embed_dim: int,
num_heads: int,
kv_dim: Optional[int] = None,
norm_layer: Callable[[int], nn.LayerNorm] = DEFAULT_LN,
max_size: tuple[int, int] = (70, 70),
max_temporal_size=36000,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> None:
super().__init__(
num_queries,
embed_dim,
num_heads,
kv_dim,
norm_layer,
quant_config=quant_config,
prefix=prefix,
)
self.max_size = max_size
self.max_temporal_size = max_temporal_size
self._set_2d_pos_cache(self.max_size)
self._set_temporal_pos_cache(self.max_temporal_size)
self.apply(self._init_weights)
def get_1d_sincos_pos_embed_from_temporal_size(
self, embed_dim: int, pos: np.ndarray
):
"""
embed_dim: output dimension for each position
pos: a list of positions to be encoded: size (M,)
out: (M, D)
"""
assert embed_dim % 2 == 0
omega = np.arange(embed_dim // 2, dtype=np.float32)
omega /= embed_dim / 2.0
omega = 1.0 / 10000**omega # (D/2,)
pos = pos.reshape(-1) # (M,)
out = np.einsum("m,d->md", pos, omega) # (M, D/2), outer product
emb_sin = np.sin(out) # (M, D/2)
emb_cos = np.cos(out) # (M, D/2)
emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
return emb
def _set_2d_pos_cache(
self, max_size: tuple[int, int], device: torch.types.Device = "cpu"
) -> None:
pos_embed_arr = get_2d_sincos_pos_embed(
self.embed_dim, max_size, version=(2, 5)
)
pos_embed = torch.from_numpy(pos_embed_arr).float().to(device)
self.register_buffer("pos_embed", pos_embed, persistent=False)
def _adjust_pos_cache(
self, tgt_sizes: torch.Tensor, device: torch.types.Device
) -> None:
max_h = tgt_sizes[:, 0].max().item()
max_w = tgt_sizes[:, 1].max().item()
assert isinstance(max_h, int) and isinstance(max_w, int)
if max_h > self.max_size[0] or max_w > self.max_size[1]:
self.max_size = (
max(max_h, self.max_size[0]),
max(max_w, self.max_size[1]),
)
self._set_2d_pos_cache(self.max_size, device)
def _set_temporal_pos_cache(
self, max_temporal_size: int, device: torch.types.Device = "cpu"
) -> None:
temporal_size = np.arange(max_temporal_size, dtype=np.float32)
pos_embed = (
torch.from_numpy(
self.get_1d_sincos_pos_embed_from_temporal_size(
self.embed_dim, temporal_size
)
)
.float()
.to(device)
)
self.register_buffer("temporal_pos_embed", pos_embed, persistent=False)
def _adjust_temporal_pos_cache(
self, max_temporal_size: int, device: torch.types.Device = "cpu"
):
if max_temporal_size > self.max_temporal_size:
self.max_temporal_size = max_temporal_size
self._set_temporal_pos_cache(self.max_temporal_size, device)
def forward(
self, x: torch.Tensor, tgt_sizes: torch.Tensor, temporal_ids=None
) -> torch.Tensor:
assert x.shape[0] == tgt_sizes.shape[0]
bs = x.shape[0]
device = x.device
dtype = x.dtype
patch_len = tgt_sizes[:, 0] * tgt_sizes[:, 1]
self._adjust_pos_cache(tgt_sizes, device=device)
temporal_pos_emb = False
temporal_ids_flatten = None
if temporal_ids is not None:
# example: [[-1], [-1], [2, 6, 9]]
temporal_ids_flatten = list(chain.from_iterable(temporal_ids))
max_temporal_size = max(temporal_ids_flatten)
if max_temporal_size > -1:
temporal_pos_emb = True
if max_temporal_size > self.max_temporal_size:
self._adjust_temporal_pos_cache(max_temporal_size, device)
max_patch_len = patch_len.max().item()
assert isinstance(max_patch_len, int)
key_padding_mask = torch.zeros(
(bs, max_patch_len), dtype=torch.bool, device=device
)
x, _ = self.kv_proj(x) # B * L * D
x = self.ln_kv(x).permute(1, 0, 2) # L * B * D
q = self.ln_q(self.query) # Q * D
pos_embed_2d = []
pos_embed_temporal = []
for i in range(bs):
tgt_h, tgt_w = tgt_sizes[i]
if temporal_pos_emb:
if temporal_ids_flatten[i] == -1:
pos_embed_temporal.append(
torch.zeros(self.embed_dim, dtype=dtype, device=device)
)
else:
pos_embed_temporal.append(
self.temporal_pos_embed[temporal_ids_flatten[i]].to(dtype)
) # D
pos_embed_2d.append(
self.pos_embed[:tgt_h, :tgt_w, :].reshape((tgt_h * tgt_w, -1)).to(dtype)
) # patches * D
key_padding_mask[i, patch_len[i] :] = True
pos_embed_2d = torch.nn.utils.rnn.pad_sequence(
pos_embed_2d, batch_first=True, padding_value=0.0
).permute(
1, 0, 2
) # BLD => L * B * D
k = x
v = x + pos_embed_2d
if pos_embed_temporal:
k += torch.stack(pos_embed_temporal, dim=0)
bs = len(temporal_ids)
merge_k = []
merge_v = []
merge_key_padding_mask = []
start = 0
for tp in temporal_ids:
end = start + len(tp)
# # L * (end-start) * D -> (end-start) * L * D -> 1 * L*(end-start) * D
merge_k.append(
k[:, start:end, :].permute(1, 0, 2).reshape(-1, self.embed_dim)
)
merge_v.append(
v[:, start:end, :].permute(1, 0, 2).reshape(-1, self.embed_dim)
)
merge_key_padding_mask.append(
key_padding_mask[start:end, :].reshape(-1, 1)
)
start = end
k = torch.nn.utils.rnn.pad_sequence(
merge_k, batch_first=True, padding_value=0.0
).permute(
1, 0, 2
) # L*(end-start)
v = torch.nn.utils.rnn.pad_sequence(
merge_v, batch_first=True, padding_value=0.0
).permute(
1, 0, 2
) # L*(end-start)
key_padding_mask = torch.nn.utils.rnn.pad_sequence(
merge_key_padding_mask, batch_first=True, padding_value=True
).squeeze(-1)
out = self.attn(
self._repeat(q, bs), # Q * B * D
k, # L * B * D + L * B * D
v,
key_padding_mask=key_padding_mask,
)[0]
# out: Q * B * D
x = out.permute(1, 0, 2) # B * Q * D
x = self.ln_post(x)
x = x @ self.proj
return x
def get_version_by_config(config: PretrainedConfig) -> Tuple[int, ...]: def get_version_by_config(config: PretrainedConfig) -> Tuple[int, ...]:
version_float = getattr(config, "version", None) version_float = getattr(config, "version", None)
@@ -933,10 +1148,173 @@ class MiniCPMV4_0(MiniCPMBaseModel):
return pattern.pad_input_tokens(input_ids, image_inputs) return pattern.pad_input_tokens(input_ids, image_inputs)
_SUPPORT_VERSION = { class MiniCPMV4_5(MiniCPMBaseModel):
(2, 6): MiniCPMV2_6, packed_modules_mapping = {
(4, 0): MiniCPMV4_0, "qkv_proj": [
} "q_proj",
"k_proj",
"v_proj",
],
"gate_up_proj": [
"gate_proj",
"up_proj",
],
}
# LoRA specific attributes
supported_lora_modules = [
# vision encoder
"fc1",
"fc2",
"out_proj",
# language model
"qkv_proj", # same name with vision encoder
"o_proj",
"gate_up_proj",
"down_proj",
# resampler
"kv_proj",
]
# BitandBytes specific attributes
bitsandbytes_stacked_params_mapping = {
# shard_name, weight_name, index
"q_proj": ("qkv_proj", 0),
"k_proj": ("qkv_proj", 1),
"v_proj": ("qkv_proj", 2),
"gate_proj": ("gate_up_proj", 0),
"up_proj": ("gate_up_proj", 1),
}
embedding_modules = {}
embedding_padding_modules = []
def __init__(
self,
config: PretrainedConfig,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__(config=config, quant_config=quant_config, prefix=prefix)
assert self.version == (4, 5)
def init_llm(
self,
config: Qwen3Config,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> nn.Module:
llm = Qwen3ForCausalLM(config=config, quant_config=quant_config, prefix=prefix)
llm.get_input_embeddings = types.MethodType(
lambda self: self.model.get_input_embeddings(), llm
)
return llm
def init_vision_module(
self,
config: PretrainedConfig,
quant_config: Optional[QuantizationConfig],
prefix: str = "",
) -> nn.Module:
model = Idefics2VisionTransformer(
config=config.vision_config, quant_config=quant_config, prefix=prefix
)
if self.config.drop_vision_last_layer:
model.encoder.layers = model.encoder.layers[:-1]
setattr(model, "embed_dim", model.embeddings.embed_dim)
setattr(model, "patch_size", model.embeddings.patch_size)
return model
def init_resampler(
self,
embed_dim: int,
vision_dim: int,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
) -> nn.Module:
with set_default_torch_dtype(torch.float16):
# The resampler in 2.6 remains consistent with the one in 2.5.
resampler = Resampler4_5(
num_queries=self.config.query_num,
embed_dim=embed_dim,
num_heads=embed_dim // 128,
kv_dim=vision_dim,
quant_config=quant_config,
prefix=prefix,
)
return resampler.to(device="cuda", dtype=torch.get_default_dtype())
def get_vision_embedding(
self,
pixel_values: List[torch.Tensor],
patch_attn_mask: Optional[torch.Tensor] = None,
tgt_sizes: Optional[torch.Tensor] = None,
) -> torch.Tensor:
vision_embedding = self.vpm(
pixel_values,
patch_attention_mask=patch_attn_mask,
tgt_sizes=tgt_sizes,
)
return vision_embedding
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
# list of tensors
pixel_values = flatten_nested_list([item.feature for item in items])
tgt_sizes = torch.stack(
flatten_nested_list([item.tgt_size for item in items]), dim=0
)
assert len(pixel_values) == tgt_sizes.shape[0]
device = self.vpm.embeddings.position_embedding.weight.device
dtype = self.vpm.embeddings.position_embedding.weight.dtype
all_pixel_values_lst = [
i.flatten(end_dim=1).permute(1, 0) for i in pixel_values
]
max_patches = (tgt_sizes[:, 0] * tgt_sizes[:, 1]).max().item()
assert isinstance(max_patches, int)
all_pixel_values = torch.nn.utils.rnn.pad_sequence(
all_pixel_values_lst, batch_first=True, padding_value=0.0
)
B, L, _ = all_pixel_values.shape
all_pixel_values = all_pixel_values.permute(0, 2, 1).reshape(B, 3, -1, L)
patch_attn_mask = torch.zeros(
(B, 1, max_patches), dtype=torch.bool, device=device
)
tgt_sizes_tensor = tgt_sizes.clone().to(device=patch_attn_mask.device)
mask_shapes = tgt_sizes_tensor[:, 0] * tgt_sizes_tensor[:, 1]
patch_attn_mask[:, 0, :] = torch.arange(
patch_attn_mask.size(2), device=patch_attn_mask.device
).unsqueeze(0) < mask_shapes.unsqueeze(1)
vision_embedding = self.vpm(
all_pixel_values.type(dtype),
patch_attention_mask=patch_attn_mask,
tgt_sizes=tgt_sizes,
)
return self.resampler(vision_embedding, tgt_sizes)
def pad_input_ids(self, input_ids: List[int], image_inputs: MultimodalInputs):
# Get all special token IDs
im_start_id: int = image_inputs.im_start_id
im_end_id: int = image_inputs.im_end_id
slice_start_id: int = image_inputs.slice_start_id
slice_end_id: int = image_inputs.slice_end_id
media_token_pairs = [(im_start_id, im_end_id), (slice_start_id, slice_end_id)]
pattern = MultiModalityDataPaddingPatternTokenPairs(media_token_pairs)
return pattern.pad_input_tokens(input_ids, image_inputs)
def eval(self):
super().eval()
return self
_SUPPORT_VERSION = {(2, 6): MiniCPMV2_6, (4, 0): MiniCPMV4_0, (4, 5): MiniCPMV4_5}
class MiniCPMV: class MiniCPMV:
@@ -971,7 +1349,13 @@ class MiniCPMV:
# Dispatch class based on version # Dispatch class based on version
instance_class = _SUPPORT_VERSION.get(version) instance_class = _SUPPORT_VERSION.get(version)
if instance_class is None: if instance_class is None:
raise ValueError("Currently, MiniCPMV only supports versions 2.6 and 4.0") supported_versions = ", ".join(
[f"{v[0]}.{v[1]}" for v in sorted(_SUPPORT_VERSION.keys())]
)
raise ValueError(
f"Currently, MiniCPMV only supports versions "
f"{supported_versions}. Got version: {version}"
)
try: try:
minicpmv = instance_class( minicpmv = instance_class(