from __future__ import annotations from abc import ABC, abstractmethod from typing import Any, Optional, Tuple import torch import torch.nn as nn from transformers.configuration_utils import PretrainedConfig from transformers.modeling_utils import PreTrainedModel from .configuration_intern_vit import InternVisionConfig from .modeling_intern_vit import InternVisionModel from .qts_plus_tokenizer import QTSplusTokenizer, QTSplusTokenizerConfig def qts_integrate_embeddings( vision_features: torch.Tensor, input_ids: torch.Tensor, attention_mask: torch.Tensor, labels: Optional[torch.Tensor] = None, image_token_id: Optional[int] = None, video_token_id: Optional[int] = None, text_model_embed_layer: Optional[nn.Embedding] = None, kept_indices: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor, Optional[torch.Tensor]]: """Replace multimodal placeholder token embeddings with vision features. Supports two prompt formats: - multiple placeholders (e.g. InternVL `` repeated per vision token) - a single placeholder token (expanded into N vision tokens) """ if text_model_embed_layer is None: raise ValueError("text_model_embed_layer is required") if input_ids.dtype is not torch.long: input_ids = input_ids.long() placeholder_token_id = video_token_id if video_token_id is not None else image_token_id if placeholder_token_id is None: raise ValueError("Either `image_token_id` or `video_token_id` must be provided") inputs_embeds = text_model_embed_layer(input_ids) if vision_features.ndim != 2: raise ValueError(f"vision_features must be [N, D], got {tuple(vision_features.shape)}") if input_ids.ndim != 2 or input_ids.shape[0] != 1: raise ValueError("Only batch_size==1 is currently supported") pos = (input_ids[0] == int(placeholder_token_id)).nonzero(as_tuple=False).flatten() if pos.numel() == 0: raise ValueError("No multimodal placeholder tokens found in input_ids") n_feats = int(vision_features.shape[0]) if n_feats <= 0: raise ValueError("vision_features must contain at least one vector") # Single placeholder: expand into N tokens. if pos.numel() == 1 and n_feats >= 1: insert_at = int(pos.item()) vision_features = vision_features.to(inputs_embeds.device, inputs_embeds.dtype) pre = inputs_embeds[:, :insert_at, :] post = inputs_embeds[:, insert_at + 1 :, :] inputs_embeds = torch.cat([pre, vision_features.unsqueeze(0), post], dim=1) pre_mask = attention_mask[:, :insert_at] post_mask = attention_mask[:, insert_at + 1 :] feats_mask = torch.ones((1, n_feats), device=attention_mask.device, dtype=attention_mask.dtype) attention_mask = torch.cat([pre_mask, feats_mask, post_mask], dim=1) if labels is not None: pre_lab = labels[:, :insert_at] post_lab = labels[:, insert_at + 1 :] feats_lab = torch.full((1, n_feats), -100, device=labels.device, dtype=labels.dtype) labels = torch.cat([pre_lab, feats_lab, post_lab], dim=1) return inputs_embeds, attention_mask.to(inputs_embeds.device), labels # Multi-placeholder: drop unselected placeholders, then replace remaining. vision_features = vision_features.to(inputs_embeds.device, inputs_embeds.dtype) m_placeholders = int(pos.numel()) if n_feats > m_placeholders: raise ValueError( f"Number of vision features ({n_feats}) exceeds placeholder tokens ({m_placeholders}). " "Ensure the prompt inserts enough tokens." ) if n_feats < m_placeholders: if kept_indices is not None: keep_idx = kept_indices.flatten().to(device=pos.device, dtype=torch.long) keep_idx = keep_idx[(keep_idx >= 0) & (keep_idx < m_placeholders)] if keep_idx.numel() != n_feats: keep_idx = torch.arange(n_feats, device=pos.device, dtype=torch.long) order = torch.argsort(keep_idx) keep_idx = keep_idx[order] vision_features = vision_features[order.to(device=vision_features.device)] keep_mask = torch.zeros((m_placeholders,), device=pos.device, dtype=torch.bool) keep_mask[keep_idx] = True drop_pos = pos[~keep_mask] else: drop_pos = pos[n_feats:] if drop_pos.numel() > 0: keep_seq = torch.ones((input_ids.shape[1],), device=input_ids.device, dtype=torch.bool) keep_seq[drop_pos] = False input_ids = input_ids[:, keep_seq] attention_mask = attention_mask[:, keep_seq] inputs_embeds = inputs_embeds[:, keep_seq, :] if labels is not None: labels = labels[:, keep_seq] pos = (input_ids[0] == int(placeholder_token_id)).nonzero(as_tuple=False).flatten() # Replace placeholder embeddings. if int(pos.numel()) != n_feats: raise ValueError(f"Placeholder tokens ({int(pos.numel())}) != vision features ({n_feats}) after trimming") for i in range(n_feats): inputs_embeds[0, int(pos[i].item()), :] = vision_features[i, :] if labels is not None and n_feats > 0: labels = labels.clone() labels[0, pos[:n_feats]] = -100 return inputs_embeds, attention_mask.to(inputs_embeds.device), labels class InternVL2_5VisionConfig(PretrainedConfig): model_type = "internvl2_5_vision" is_composition = True def __init__( self, vision_config: Optional[dict[str, Any]] = None, llm_hidden_size: Optional[int] = None, select_layer: int = -1, force_image_size: Optional[int] = None, downsample_ratio: float = 0.5, ps_version: str = "v2", **kwargs: Any, ) -> None: super().__init__(**kwargs) if vision_config is None: vision_config = {"architectures": ["InternVisionModel"]} self.vision_config = InternVisionConfig(**vision_config) self.select_layer = int(select_layer) self.force_image_size = int(force_image_size) if force_image_size is not None else None self.downsample_ratio = float(downsample_ratio) self.ps_version = str(ps_version) self.hidden_size = int(self.vision_config.hidden_size) self.out_hidden_size = int(llm_hidden_size) if llm_hidden_size is not None else int(self.hidden_size) self.llm_hidden_size = int(self.out_hidden_size) self.architectures = ["InternVL2_5VisionTower"] def to_dict(self) -> dict[str, Any]: out = dict(self.__dict__) out["vision_config"] = self.vision_config.to_dict() out["model_type"] = self.__class__.model_type return out class InternVL2_5VisionTower(PreTrainedModel): config_class = InternVL2_5VisionConfig main_input_name = "pixel_values" def __init__(self, config: InternVL2_5VisionConfig): super().__init__(config) vision_cfg = config.vision_config if config.force_image_size is not None: vision_cfg = InternVisionConfig(**vision_cfg.to_dict()) vision_cfg.image_size = int(config.force_image_size) self.vision_model = InternVisionModel(vision_cfg) self.select_layer = int(config.select_layer) self.downsample_ratio = float(config.downsample_ratio) self.ps_version = str(config.ps_version) vit_hidden_size = int(vision_cfg.hidden_size) llm_hidden_size = int(config.out_hidden_size) mlp_in = vit_hidden_size * int(1 / self.downsample_ratio) ** 2 self.mlp1 = nn.Sequential( nn.LayerNorm(mlp_in), nn.Linear(mlp_in, llm_hidden_size), nn.GELU(), nn.Linear(llm_hidden_size, llm_hidden_size), ) self.post_init() def pixel_shuffle(self, x: torch.Tensor, scale_factor: float = 0.5) -> torch.Tensor: n, w, h, c = x.size() x = x.view(n, w, int(h * scale_factor), int(c / scale_factor)) x = x.permute(0, 2, 1, 3).contiguous() x = x.view( n, int(h * scale_factor), int(w * scale_factor), int(c / (scale_factor * scale_factor)), ) if self.ps_version != "v1": x = x.permute(0, 2, 1, 3).contiguous() return x def extract_feature(self, pixel_values: torch.Tensor) -> torch.Tensor: if self.select_layer == -1: vit_out = self.vision_model( pixel_values=pixel_values, output_hidden_states=False, return_dict=True, ).last_hidden_state else: vit_out = self.vision_model( pixel_values=pixel_values, output_hidden_states=True, return_dict=True, ).hidden_states[self.select_layer] vit_out = vit_out[:, 1:, :] # drop CLS h = w = int(vit_out.shape[1] ** 0.5) vit_out = vit_out.reshape(vit_out.shape[0], h, w, -1) vit_out = self.pixel_shuffle(vit_out, scale_factor=self.downsample_ratio) vit_out = vit_out.reshape(vit_out.shape[0], -1, vit_out.shape[-1]) vit_out = self.mlp1(vit_out) return vit_out def get_image_features(self, pixel_values: torch.Tensor) -> torch.Tensor: return self.extract_feature(pixel_values) def forward(self, pixel_values: torch.Tensor, **_: Any) -> torch.Tensor: return self.get_image_features(pixel_values) def build_vision_tower(config: PretrainedConfig) -> InternVL2_5VisionTower: vision_cfg = getattr(config, "vision_config", None) if not isinstance(vision_cfg, dict): raise ValueError("Missing `vision_config` in model config for InternVL2.5 vision tower") llm_hidden = getattr(config, "hidden_size", None) if not isinstance(llm_hidden, int) or llm_hidden <= 0: llm_hidden = getattr(config, "llm_hidden_size", None) if not isinstance(llm_hidden, int) or llm_hidden <= 0: raise ValueError("Missing `hidden_size` / `llm_hidden_size` in config") vt_cfg = InternVL2_5VisionConfig( vision_config=vision_cfg, llm_hidden_size=int(llm_hidden), select_layer=int(getattr(config, "select_layer", -1)), force_image_size=getattr(config, "force_image_size", None), downsample_ratio=float(getattr(config, "downsample_ratio", 0.5)), ps_version=str(getattr(config, "ps_version", "v2")), ) return InternVL2_5VisionTower(vt_cfg) def build_qts_plus_tower(config: PretrainedConfig) -> QTSplusTokenizer: vision_dim = getattr(config, "vision_embed_size", None) if not isinstance(vision_dim, int) or vision_dim <= 0: vision_dim = getattr(config, "hidden_size", None) if not isinstance(vision_dim, int) or vision_dim <= 0: raise ValueError("Missing `vision_embed_size` / `hidden_size` in config") lm_heads = getattr(config, "num_attention_heads", None) if not isinstance(lm_heads, int) or lm_heads <= 0: raise ValueError("Missing `num_attention_heads` in config") if vision_dim % lm_heads != 0: raise ValueError(f"vision_embed_size ({vision_dim}) must be divisible by num_attention_heads ({lm_heads})") kv_heads = getattr(config, "num_key_value_heads", None) cfg = QTSplusTokenizerConfig( embedding_dim=int(vision_dim), n_heads=int(lm_heads), num_kv_heads=int(kv_heads) if isinstance(kv_heads, int) and kv_heads > 0 else None, tau_s=float(getattr(config, "qts_plus_tau_s", 0.1)), nmax=int(getattr(config, "qts_plus_nmax", 2560)), rho_min=float(getattr(config, "qts_plus_rho_min", 0.05)), rho_max=float(getattr(config, "qts_plus_rho_max", 0.5)), block_dropout=float(getattr(config, "qts_plus_block_dropout", 0.0)), reencode=bool(getattr(config, "qts_plus_reencode", False)), scoring_layers=int(getattr(config, "qts_plus_scoring_layers", 1)), reencode_layers=int(getattr(config, "qts_plus_reencode_layers", 0)), lambda_t=float(getattr(config, "lambda_t", 1.0)), lambda_m=float(getattr(config, "lambda_m", 1.7)), lambda_s=float(getattr(config, "lambda_s", 0.05)), project_text_if_needed=bool(getattr(config, "project_text_if_needed", False)), ) return QTSplusTokenizer(cfg) class QTSplusMetaModel: def __init__(self, config: PretrainedConfig): super().__init__(config) self.config = config self.vision_tower = None if getattr(config, "vision_tower", None) in {"internvl2_5_vision", "internvl_vision"}: self.vision_tower = build_vision_tower(config) self.qts_plus = None if getattr(config, "enable_qts_plus", False): self.qts_plus = build_qts_plus_tower(config) def get_qts_plus_tower(self): return getattr(self, "qts_plus", None) def get_vision_tower(self): return getattr(self, "vision_tower", None) class QTSplusMetaForCausalLM(ABC): @abstractmethod def get_model(self): # pragma: no cover raise NotImplementedError def get_qts_plus_tower(self): return self.get_model().get_qts_plus_tower() def get_vision_tower(self): return self.get_model().get_vision_tower() def prepare_inputs_for_multimodal( self, vision_input: Optional[torch.FloatTensor] = None, input_ids: Optional[torch.LongTensor] = None, position_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, past_key_values: Optional[list[torch.FloatTensor]] = None, labels: Optional[torch.LongTensor] = None, question_input_ids: Optional[torch.LongTensor] = None, image_token_id: Optional[int] = None, video_token_id: Optional[int] = None, mode: str = "train", ): if attention_mask is None and input_ids is not None: attention_mask = torch.ones_like(input_ids, dtype=torch.long, device=input_ids.device) # Default: no multimodal inputs -> no-op. if vision_input is None: z = torch.tensor(0.0, device=input_ids.device if input_ids is not None else None) return vision_input, position_ids, attention_mask, past_key_values, None, labels, z, z, z if question_input_ids is None: raise ValueError("`question_input_ids` is required for QTSplus InternVL2.5 inference/training.") if question_input_ids.dtype is not torch.long: question_input_ids = question_input_ids.long() if question_input_ids.ndim == 1: question_input_ids = question_input_ids.unsqueeze(0) vision_tower = self.get_vision_tower() qts_plus_tower = self.get_qts_plus_tower() text_embed_layer = self.get_model().get_input_embeddings() if vision_tower is None or qts_plus_tower is None: raise ValueError("Both `vision_tower` and `qts_plus` must be initialized for multimodal inference.") # Normalize `vision_input` into a pixel_values tensor. if isinstance(vision_input, list): if len(vision_input) == 0: z = torch.tensor(0.0, device=input_ids.device) return None, position_ids, attention_mask, past_key_values, None, labels, z, z, z vision_input = vision_input[0] pixel_values = vision_input.get("pixel_values") if isinstance(vision_input, dict) else vision_input if not isinstance(pixel_values, torch.Tensor): raise ValueError(f"vision_input must be a torch.Tensor or dict with pixel_values, got {type(vision_input)}") if pixel_values.ndim == 3: # [3, H, W] pixel_values = pixel_values.unsqueeze(0).unsqueeze(0) # [1, 1, 3, H, W] elif pixel_values.ndim == 4: # [B, 3, H, W] or [T, 3, H, W] b_txt = int(question_input_ids.shape[0]) if pixel_values.shape[0] == b_txt: pixel_values = pixel_values.unsqueeze(1) # [B, 1, 3, H, W] else: pixel_values = pixel_values.unsqueeze(0) # [1, T, 3, H, W] elif pixel_values.ndim != 5: raise ValueError(f"Unsupported InternVL pixel_values shape: {tuple(pixel_values.shape)}") b, t, c, h, w = pixel_values.shape pixel_values_flat = pixel_values.reshape(b * t, c, h, w) try: vt_param = next(vision_tower.parameters()) vt_device = vt_param.device vt_dtype = vt_param.dtype except StopIteration: vt_device = pixel_values_flat.device vt_dtype = pixel_values_flat.dtype vision_features = vision_tower.get_image_features(pixel_values_flat.to(device=vt_device, dtype=vt_dtype)) if not (isinstance(vision_features, torch.Tensor) and vision_features.ndim == 3): raise ValueError(f"vision_tower must return [B, N, D], got {type(vision_features)} {vision_features.shape}") vision_features = vision_features.reshape(b, t * vision_features.shape[1], vision_features.shape[2]) text_embeddings = text_embed_layer(question_input_ids.to(text_embed_layer.weight.device)) vision_features = vision_features.to(device=text_embeddings.device, dtype=text_embeddings.dtype) try: qts_plus_tower.to(device=text_embeddings.device, dtype=text_embeddings.dtype) except Exception: qts_plus_tower.to(device=text_embeddings.device) qts_plus_out = qts_plus_tower(vision_features, text_embeddings, mode=mode) z_list = qts_plus_out["Z"] if not (isinstance(z_list, list) and len(z_list) == 1 and isinstance(z_list[0], torch.Tensor)): raise ValueError("Expected QTSplusTokenizer to return a list of 1 tensor for batch_size==1") kept = None try: kept_list = qts_plus_out.get("indices") kept = kept_list[0] if isinstance(kept_list, list) and len(kept_list) == 1 else None except Exception: kept = None if image_token_id is None: image_token_id = getattr(self.config, "image_token_id", 92546) inputs_embeds, attention_mask, labels = qts_integrate_embeddings( vision_features=z_list[0], input_ids=input_ids, attention_mask=attention_mask, labels=labels, image_token_id=image_token_id, video_token_id=video_token_id, text_model_embed_layer=text_embed_layer, kept_indices=kept, ) add_loss = qts_plus_out.get("add_loss") or {} flops_loss = add_loss.get("flops", 0.0) kv_loss = add_loss.get("kv", 0.0) smooth_loss = add_loss.get("smooth", 0.0) # Return `inputs_embeds` so the LM consumes the integrated embeddings. return ( vision_input, position_ids, attention_mask, past_key_values, inputs_embeds, labels, flops_loss, kv_loss, smooth_loss, )