# -*- coding: utf-8 -*- """ MiniCPM5-Vision-2B: Self-Contained Omni-Modal Vision-Language Architecture. Combines MiniCPM5-2B LLaMA backbone, SigLIP-so400m-patch14-384 vision encoder, and 2x2 Spatial Unshuffle Projector. """ import copy import math from typing import List, Optional, Tuple, Union import numpy as np from PIL import Image import torch import torch.nn as nn from torchvision import transforms from transformers.configuration_utils import PretrainedConfig from transformers.modeling_utils import PreTrainedModel from transformers.models.llama.configuration_llama import LlamaConfig from transformers.models.llama.modeling_llama import LlamaForCausalLM, LlamaModel from transformers.models.siglip.configuration_siglip import SiglipVisionConfig from transformers.models.siglip.modeling_siglip import SiglipVisionModel from transformers.modeling_outputs import CausalLMOutputWithPast class MiniCPM5VSliceConfig(PretrainedConfig): r"""Configuration for dynamic image slicing in MiniCPM5-V.""" def __init__( self, max_slice_nums=9, scale_resolution=448, patch_size=14, slice_mode=True, **kwargs ): super().__init__(**kwargs) self.max_slice_nums = max_slice_nums self.scale_resolution = scale_resolution self.patch_size = patch_size self.slice_mode = slice_mode class MiniCPM5VConfig(PretrainedConfig): model_type = "minicpm5_v" is_composition = True def __init__( self, text_config=None, vision_config=None, slice_config=None, projector_hidden_act="gelu", spatial_downsample_factor=2, drop_vision_last_layer=False, image_token_id=130074, slice_start_token_id=130075, slice_end_token_id=130076, box_start_token_id=130077, box_end_token_id=130078, ref_start_token_id=130079, ref_end_token_id=130080, **kwargs ): super().__init__(**kwargs) if text_config is None: text_config = { "vocab_size": 130561, "hidden_size": 2048, "intermediate_size": 6144, "num_hidden_layers": 42, "num_attention_heads": 16, "num_key_value_heads": 2, "head_dim": 128, "max_position_embeddings": 131072, "rms_norm_eps": 1e-6, "rope_theta": 5000000.0, "hidden_act": "silu", "tie_word_embeddings": False, "torch_dtype": "bfloat16", } if vision_config is None: vision_config = { "hidden_size": 1152, "image_size": 384, "intermediate_size": 4304, "num_attention_heads": 16, "num_hidden_layers": 27, "patch_size": 14, "hidden_act": "gelu_pytorch_tanh", "layer_norm_eps": 1e-6, } if slice_config is None: slice_config = { "max_slice_nums": 9, "scale_resolution": 448, "patch_size": 14, "slice_mode": True, } if isinstance(text_config, dict): self.text_config = LlamaConfig(**text_config) else: self.text_config = text_config if isinstance(vision_config, dict): self.vision_config = SiglipVisionConfig(**vision_config) else: self.vision_config = vision_config if isinstance(slice_config, dict): self.slice_config = MiniCPM5VSliceConfig(**slice_config) else: self.slice_config = slice_config self.projector_hidden_act = projector_hidden_act self.spatial_downsample_factor = spatial_downsample_factor self.drop_vision_last_layer = drop_vision_last_layer self.image_token_id = image_token_id self.slice_start_token_id = slice_start_token_id self.slice_end_token_id = slice_end_token_id self.box_start_token_id = box_start_token_id self.box_end_token_id = box_end_token_id self.ref_start_token_id = ref_start_token_id self.ref_end_token_id = ref_end_token_id def to_dict(self): output = copy.deepcopy(self.__dict__) output["text_config"] = self.text_config.to_dict() output["vision_config"] = self.vision_config.to_dict() output["slice_config"] = self.slice_config.to_dict() output["model_type"] = self.__class__.model_type return output class MiniCPM5VSpatialUnshuffleProjector(nn.Module): def __init__(self, in_dim=4608, out_dim=2048, downsample_factor=2): super().__init__() self.downsample_factor = downsample_factor self.mlp = nn.Sequential( nn.Linear(in_dim, out_dim, bias=True), nn.GELU(), nn.Linear(out_dim, out_dim, bias=True) ) def forward(self, x: torch.Tensor, grid_h: int = 32, grid_w: int = 32) -> torch.Tensor: b, n, c = x.shape x = x.view(b, grid_h, grid_w, c) h_down = grid_h // self.downsample_factor w_down = grid_w // self.downsample_factor x = x.view(b, h_down, self.downsample_factor, w_down, self.downsample_factor, c) x = x.permute(0, 1, 3, 2, 4, 5).contiguous() x = x.view(b, h_down * w_down, c * (self.downsample_factor ** 2)) return self.mlp(x) class MiniCPM5VPreTrainedModel(PreTrainedModel): config_class = MiniCPM5VConfig base_model_prefix = "model" supports_gradient_checkpointing = True _no_split_modules = ["LlamaDecoderLayer", "SiglipVisionEmbeddings", "SiglipEncoderLayer"] class MiniCPM5VForConditionalGeneration(MiniCPM5VPreTrainedModel): def __init__(self, config: MiniCPM5VConfig): super().__init__(config) self.config = config self.vpm = SiglipVisionModel(config.vision_config) in_dim = config.vision_config.hidden_size * (config.spatial_downsample_factor ** 2) out_dim = config.text_config.hidden_size self.resampler = MiniCPM5VSpatialUnshuffleProjector(in_dim=in_dim, out_dim=out_dim, downsample_factor=config.spatial_downsample_factor) self.llm = LlamaForCausalLM(config.text_config) self.post_init() def get_input_embeddings(self): return self.llm.get_input_embeddings() def set_input_embeddings(self, value): self.llm.set_input_embeddings(value) def get_output_embeddings(self): return self.llm.get_output_embeddings() def set_output_embeddings(self, new_embeddings): self.llm.set_output_embeddings(new_embeddings) def encode_vision(self, pixel_values: torch.Tensor) -> torch.Tensor: vision_outputs = self.vpm(pixel_values=pixel_values, interpolate_pos_encoding=True) hidden_states = vision_outputs.last_hidden_state img_size = pixel_values.shape[-1] patch_size = self.config.vision_config.patch_size grid_dim = img_size // patch_size return self.resampler(hidden_states, grid_h=grid_dim, grid_w=grid_dim) def forward( self, input_ids: Optional[torch.LongTensor] = None, pixel_values: Optional[torch.FloatTensor] = None, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_values: Optional[List[torch.FloatTensor]] = None, inputs_embeds: Optional[torch.FloatTensor] = None, labels: Optional[torch.LongTensor] = None, use_cache: Optional[bool] = None, output_attentions: Optional[bool] = None, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, ) -> Union[Tuple, CausalLMOutputWithPast]: return_dict = return_dict if return_dict is not None else self.config.use_return_dict if inputs_embeds is None: safe_input_ids = input_ids.clone() image_mask = (safe_input_ids == self.config.image_token_id) safe_input_ids[image_mask] = 0 inputs_embeds = self.get_input_embeddings()(safe_input_ids) if pixel_values is not None: vision_embeds = self.encode_vision(pixel_values) flat_vision_embeds = vision_embeds.view(-1, vision_embeds.shape[-1]) image_token_mask = (input_ids == self.config.image_token_id) num_image_tokens = image_token_mask.sum().item() if num_image_tokens > 0: b, s, d = inputs_embeds.shape flat_embeds = inputs_embeds.view(-1, d).clone() flat_mask = image_token_mask.view(-1) n = min(num_image_tokens, flat_vision_embeds.shape[0]) match_indices = torch.nonzero(flat_mask, as_tuple=False).squeeze(-1)[:n] flat_embeds[match_indices] = flat_vision_embeds[:n].to(dtype=flat_embeds.dtype) inputs_embeds = flat_embeds.view(b, s, d) return self.llm( input_ids=None, inputs_embeds=inputs_embeds, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, labels=labels, use_cache=use_cache, output_attentions=output_attentions, output_hidden_states=output_hidden_states, return_dict=return_dict, ) def preprocess_inputs( self, image: Image.Image, prompt: str, tokenizer, device=None ): r"""Convenience helper to format visual inputs and prompt tokens into model tensor feeds.""" dev = device if device is not None else next(self.parameters()).device img_proc = MiniCPM5VImageProcessor() slices_tensor = img_proc.preprocess_image(image).to(device=dev, dtype=self.dtype) num_slices = slices_tensor.shape[0] total_vis_tokens = num_slices * 256 image_tokens = [self.config.image_token_id] * total_vis_tokens text_tokens = tokenizer.encode(prompt, add_special_tokens=False) input_ids = torch.tensor([image_tokens + text_tokens], dtype=torch.long, device=dev) attention_mask = torch.ones_like(input_ids) return { "input_ids": input_ids, "pixel_values": slices_tensor, "attention_mask": attention_mask } @torch.no_grad() def generate( self, input_ids: torch.LongTensor, pixel_values: Optional[torch.FloatTensor] = None, **kwargs ): """Generate new tokens from image+text inputs. NOTE: Returns only the newly generated tokens (shape: [batch, new_tokens]). The input prefix is NOT included in the output. Decode outputs[0] directly: response = tokenizer.decode(outputs[0], skip_special_tokens=True) """ safe_input_ids = input_ids.clone() image_mask = (safe_input_ids == self.config.image_token_id) safe_input_ids[image_mask] = 0 inputs_embeds = self.get_input_embeddings()(safe_input_ids) if pixel_values is not None: vision_embeds = self.encode_vision(pixel_values) flat_vision_embeds = vision_embeds.view(-1, vision_embeds.shape[-1]) image_token_mask = (input_ids == self.config.image_token_id) num_image_tokens = image_token_mask.sum().item() if num_image_tokens > 0: b, s, d = inputs_embeds.shape flat_embeds = inputs_embeds.view(-1, d).clone() flat_mask = image_mask.view(-1) n = min(num_image_tokens, flat_vision_embeds.shape[0]) match_indices = torch.nonzero(flat_mask, as_tuple=False).squeeze(-1)[:n] flat_embeds[match_indices] = flat_vision_embeds[:n].to(dtype=flat_embeds.dtype) inputs_embeds = flat_embeds.view(b, s, d) return self.llm.generate( inputs_embeds=inputs_embeds, **kwargs ) class MiniCPM5VImageProcessor: def __init__( self, image_size: int = 448, max_slice_nums: int = 9, scale_resolution: int = 448, patch_size: int = 14, image_mean: Tuple[float, float, float] = (0.5, 0.5, 0.5), image_std: Tuple[float, float, float] = (0.5, 0.5, 0.5), ): self.image_size = image_size self.max_slice_nums = max_slice_nums self.scale_resolution = scale_resolution self.patch_size = patch_size self.image_mean = image_mean self.image_std = image_std self.transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize(mean=self.image_mean, std=self.image_std) ]) def get_slice_grid(self, orig_w: int, orig_h: int) -> Tuple[int, int]: orig_aspect = orig_w / max(1, orig_h) best_grid = (1, 1) min_error = float("inf") for total_slices in range(1, self.max_slice_nums + 1): for gw in range(1, total_slices + 1): if total_slices % gw == 0: gh = total_slices // gw grid_aspect = gw / gh error = abs(math.log(orig_aspect / grid_aspect)) if error < min_error: min_error = error best_grid = (gw, gh) return best_grid def slice_image(self, image: Image.Image) -> List[Image.Image]: image = image.convert("RGB") w, h = image.size overview = image.resize((self.scale_resolution, self.scale_resolution), Image.Resampling.BICUBIC) gw, gh = self.get_slice_grid(w, h) if gw == 1 and gh == 1: return [overview] target_w = gw * self.scale_resolution target_h = gh * self.scale_resolution resized_full = image.resize((target_w, target_h), Image.Resampling.BICUBIC) slices = [] for j in range(gh): for i in range(gw): box = ( i * self.scale_resolution, j * self.scale_resolution, (i + 1) * self.scale_resolution, (j + 1) * self.scale_resolution, ) slices.append(resized_full.crop(box)) return slices + [overview] def preprocess_image(self, image: Image.Image) -> torch.Tensor: slices = self.slice_image(image) tensors = [self.transform(s) for s in slices] return torch.stack(tensors, dim=0)