webAI-Official/webAI-ColVec1.1-8b
8191
1"""2Processing utilities for ColQwen35Bidirection retrieval.3 4Wraps the Qwen 3.5 VL processor components (image_processor, tokenizer,5video_processor) with retrieval-specific helpers for prompt construction,6MaxSim scoring, and batch handling.7 8processor kwargs:9 doc_prompt: Document prompt text appended after the image token.10 max_num_visual_tokens: Cap on visual tokens per image, controls11 resolution via max_pixels = max_num_visual_tokens × tile².12"""13 14from __future__ import annotations15 16import os17from typing import Any, List, Optional, Union18 19import numpy as np20from PIL import Image21 22from transformers import BatchFeature23from transformers.processing_utils import ProcessorMixin24from transformers.tokenization_utils_base import TextInput25from transformers.utils import logging26 27try:28 import torch29except ImportError:30 torch = None31 32logger = logging.get_logger(__name__)33 34 35def _size_value(size: Any, key: str) -> Any:36 """Read size entries from dict-like or object-like containers."""37 if size is None:38 return None39 if isinstance(size, dict):40 return size.get(key)41 return getattr(size, key, None)42 43 44class ColQwen35BidirectionProcessor(ProcessorMixin):45 """46 Processor for ColQwen35Bidirection retrieval model.47 48 Wraps Qwen 3.5's image processor, tokenizer, and video processor49 with retrieval-specific prompt construction, ``mm_token_type_ids``50 generation (required by Qwen 3.5 for 3-D position computation),51 and MaxSim scoring utilities.52 53 Visual token budget (``max_num_visual_tokens``):54 Qwen 3.5 determines visual token count from ``max_pixels`` on the55 image processor. This class converts ``max_num_visual_tokens`` into56 the equivalent ``max_pixels`` value using:57 58 tile = patch_size × merge_size # 16 × 2 = 3259 max_pixels = max_num_visual_tokens × tile² # e.g. 512 × 1024 = 524,28860 61 Lower token budgets (e.g. 512) give memory-efficient training;62 higher budgets (e.g. 2048) give finer visual granularity at inference.63 """64 65 attributes = ["image_processor", "tokenizer", "video_processor"]66 image_processor_class = "AutoImageProcessor"67 video_processor_class = "AutoVideoProcessor"68 tokenizer_class = ("Qwen2Tokenizer", "Qwen2TokenizerFast")69 70 def __init__(71 self,72 image_processor=None,73 tokenizer=None,74 video_processor=None,75 chat_template=None,76 doc_prompt: str = "Describe the image.",77 max_num_visual_tokens: Optional[int] = None,78 query_augmentation_tokens: int = 10,79 **kwargs,80 ):81 super().__init__(82 image_processor, tokenizer, video_processor,83 chat_template=chat_template, **kwargs,84 )85 86 self.doc_prompt = doc_prompt87 self.max_num_visual_tokens = max_num_visual_tokens88 self.query_augmentation_tokens = query_augmentation_tokens89 90 if max_num_visual_tokens is not None:91 self._apply_max_pixels()92 93 self.image_token = (94 tokenizer.image_token95 if getattr(tokenizer, "image_token", None)96 else "<|image_pad|>"97 )98 self.image_token_id = (99 tokenizer.image_token_id100 if getattr(tokenizer, "image_token_id", None)101 else tokenizer.convert_tokens_to_ids(self.image_token)102 )103 self.vision_start_token = (104 tokenizer.vision_start_token105 if getattr(tokenizer, "vision_start_token", None)106 else "<|vision_start|>"107 )108 self.vision_end_token = (109 tokenizer.vision_end_token110 if getattr(tokenizer, "vision_end_token", None)111 else "<|vision_end|>"112 )113 114 self.tokenizer.padding_side = "left"115 116 self._doc_prompt_template = (117 "<|im_start|>user\n"118 f"{self.vision_start_token}{self.image_token}{self.vision_end_token}"119 f"{self.doc_prompt}"120 "<|im_end|><|endoftext|>"121 )122 123 # ------------------------------------------------------------------124 # max_pixels / visual token budget125 # ------------------------------------------------------------------126 127 def _apply_max_pixels(self) -> None:128 """Sync image_processor.max_pixels with max_num_visual_tokens.129 130 Sets ``max_pixels`` on the image processor attribute AND in the131 ``size`` dict (``longest_edge``), because different versions of132 the Qwen2VLImageProcessor read from different locations. Also133 ensures ``min_pixels`` (``shortest_edge``) does not exceed134 ``max_pixels``, which would cause the resize to ignore the cap.135 """136 patch_size = getattr(self.image_processor, "patch_size", None)137 merge_size = (138 getattr(self.image_processor, "merge_size", None)139 or getattr(self.image_processor, "spatial_merge_size", None)140 )141 if patch_size is None or merge_size is None:142 logger.warning(143 "Cannot derive max_pixels: image_processor missing "144 "patch_size or merge_size/spatial_merge_size."145 )146 return147 tile = patch_size * merge_size148 max_pixels = self.max_num_visual_tokens * tile * tile149 150 self.image_processor.max_pixels = max_pixels151 size_obj = getattr(self.image_processor, "size", None)152 if size_obj is not None:153 if isinstance(size_obj, dict):154 size_obj["longest_edge"] = max_pixels155 cur_min = size_obj.get("shortest_edge")156 if cur_min is not None and cur_min > max_pixels:157 size_obj["shortest_edge"] = max_pixels158 else:159 if hasattr(size_obj, "longest_edge"):160 size_obj.longest_edge = max_pixels161 cur_min = getattr(size_obj, "shortest_edge", None)162 if cur_min is not None and cur_min > max_pixels and hasattr(size_obj, "shortest_edge"):163 size_obj.shortest_edge = max_pixels164 165 cur_min_pixels = getattr(self.image_processor, "min_pixels", 0)166 if cur_min_pixels > max_pixels:167 self.image_processor.min_pixels = max_pixels168 169 def replace_image_token(self, image_inputs: dict, image_idx: int, **kwargs) -> str:170 """Expand one ``<|image_pad|>`` placeholder into its per-image token run.171 172 ``ProcessorMixin.__call__`` delegates placeholder expansion here, and173 ``apply_chat_template`` calls ``__call__``, so without this both raise174 ``NotImplementedError`` for image inputs.175 """176 merge_length = self.image_processor.merge_size ** 2177 num_image_tokens = image_inputs["image_grid_thw"][image_idx].prod() // merge_length178 return self.image_token * num_image_tokens179 180 @classmethod181 def from_pretrained(182 cls,183 pretrained_model_name_or_path: Union[str, os.PathLike],184 *,185 max_num_visual_tokens: Optional[int] = None,186 doc_prompt: Optional[str] = None,187 query_augmentation_tokens: Optional[int] = None,188 **kwargs,189 ) -> "ColQwen35BidirectionProcessor":190 extra_kwargs: dict[str, Any] = {}191 if doc_prompt is not None:192 extra_kwargs["doc_prompt"] = doc_prompt193 if max_num_visual_tokens is not None:194 extra_kwargs["max_num_visual_tokens"] = max_num_visual_tokens195 if query_augmentation_tokens is not None:196 extra_kwargs["query_augmentation_tokens"] = query_augmentation_tokens197 198 instance = super().from_pretrained(199 pretrained_model_name_or_path, **extra_kwargs, **kwargs,200 )201 202 if max_num_visual_tokens is not None:203 instance.max_num_visual_tokens = max_num_visual_tokens204 instance._apply_max_pixels()205 206 if query_augmentation_tokens is not None:207 instance.query_augmentation_tokens = query_augmentation_tokens208 209 return instance210 211 # ------------------------------------------------------------------212 # Retrieval protocol: process_images213 # ------------------------------------------------------------------214 215 @property216 def query_augmentation_token(self) -> str:217 return self.tokenizer.pad_token218 219 def process_images(220 self,221 images: Union[Image.Image, List[Image.Image]],222 ) -> BatchFeature:223 """224 Tokenize and encode document images for retrieval.225 226 Each image is independently processed with the doc_prompt,227 and ``mm_token_type_ids`` is computed for Qwen 3.5's 3-D228 positional encoding. Multiple images are left-padded and229 concatenated into a single batch.230 """231 if not isinstance(images, list):232 images = [images]233 if len(images) == 0:234 raise ValueError("No images provided")235 236 images = [img.convert("RGB") for img in images]237 238 per_image_features: list[BatchFeature] = []239 for image in images:240 features = self._process_single_image(image)241 per_image_features.append(features)242 243 if len(per_image_features) == 1:244 return per_image_features[0]245 246 return self._left_pad_and_concat(per_image_features)247 248 def _process_single_image(self, image: Image.Image) -> BatchFeature:249 """Process one image through the full pipeline with mm_token_type_ids."""250 size_obj = getattr(self.image_processor, "size", None)251 min_pixels = _size_value(size_obj, "shortest_edge")252 if min_pixels is None:253 min_pixels = getattr(self.image_processor, "min_pixels", None)254 255 max_pixels = _size_value(size_obj, "longest_edge")256 if max_pixels is None:257 max_pixels = getattr(self.image_processor, "max_pixels", None)258 ip_kwargs: dict[str, Any] = {259 "images": [[image]],260 }261 if min_pixels is not None:262 ip_kwargs["min_pixels"] = int(min_pixels)263 if max_pixels is not None:264 ip_kwargs["max_pixels"] = int(max_pixels)265 image_inputs = self.image_processor(**ip_kwargs)266 image_grid_thw = image_inputs["image_grid_thw"]267 268 merge_size = (269 getattr(self.image_processor, "merge_size", None)270 or getattr(self.image_processor, "spatial_merge_size", None)271 )272 if merge_size is None:273 raise ValueError(274 "Image processor missing merge_size/spatial_merge_size."275 )276 merge_length = merge_size ** 2277 278 prompt = self._doc_prompt_template279 for grid in image_grid_thw:280 num_image_tokens = int(grid.prod() // merge_length) if hasattr(grid, 'prod') else int(np.prod(grid) // merge_length)281 prompt = prompt.replace(282 self.image_token,283 "<|placeholder|>" * num_image_tokens,284 1,285 )286 prompt = prompt.replace("<|placeholder|>", self.image_token)287 288 text_inputs = self.tokenizer(289 [prompt], padding="longest", return_tensors="pt",290 )291 292 input_ids = text_inputs["input_ids"]293 mm_token_type_ids = (input_ids == self.image_token_id).to(torch.int32)294 295 data = {**text_inputs, **image_inputs}296 data["mm_token_type_ids"] = mm_token_type_ids297 298 for key in ("input_ids", "attention_mask"):299 if key in data and not isinstance(data[key], torch.Tensor):300 data[key] = torch.tensor(data[key])301 302 return BatchFeature(data=data, tensor_type="pt")303 304 # ------------------------------------------------------------------305 # Retrieval protocol: process_queries306 # ------------------------------------------------------------------307 308 def process_queries(309 self,310 texts: Union[TextInput, List[TextInput]],311 ) -> BatchFeature:312 """313 Process text queries for retrieval.314 315 Each query is wrapped in a simple chat template and tokenized.316 """317 if not isinstance(texts, list):318 texts = [texts]319 if len(texts) == 0:320 raise ValueError("No texts provided")321 322 suffix = self.query_augmentation_token * self.query_augmentation_tokens323 formatted: list[str] = []324 for text in texts:325 prompt = f"<|im_start|>user\nQuery: {text}{suffix}<|im_end|><|endoftext|>"326 formatted.append(prompt)327 328 return self.tokenizer(329 formatted,330 return_tensors="pt",331 padding="longest",332 )333 334 # ------------------------------------------------------------------335 # Scoring utilities336 # ------------------------------------------------------------------337 338 def score_retrieval(339 self,340 query_embeddings: Union[torch.Tensor, List[torch.Tensor]],341 passage_embeddings: Union[torch.Tensor, List[torch.Tensor]],342 batch_size: int = 128,343 output_dtype: Optional[torch.dtype] = None,344 output_device: Union[torch.device, str] = "cpu",345 ) -> torch.Tensor:346 """347 Compute late-interaction / MaxSim retrieval scores (ColBERT-like).348 349 Args:350 query_embeddings: Per-query multi-vector embeddings.351 passage_embeddings: Per-passage multi-vector embeddings.352 batch_size: Scoring batch size.353 output_dtype: Output tensor dtype.354 output_device: Output device.355 356 Returns:357 Tensor of shape ``(n_queries, n_passages)`` with scores.358 """359 if len(query_embeddings) == 0:360 raise ValueError("No queries provided")361 if len(passage_embeddings) == 0:362 raise ValueError("No passages provided")363 364 if output_dtype is None:365 output_dtype = query_embeddings[0].dtype366 367 scores: list[torch.Tensor] = []368 for i in range(0, len(query_embeddings), batch_size):369 batch_queries = torch.nn.utils.rnn.pad_sequence(370 query_embeddings[i : i + batch_size],371 batch_first=True, padding_value=0,372 )373 batch_scores: list[torch.Tensor] = []374 for j in range(0, len(passage_embeddings), batch_size):375 batch_passages = torch.nn.utils.rnn.pad_sequence(376 passage_embeddings[j : j + batch_size],377 batch_first=True, padding_value=0,378 )379 batch_scores.append(380 torch.einsum("bnd,csd->bcns", batch_queries, batch_passages)381 .max(dim=3)[0]382 .sum(dim=2)383 )384 scores.append(385 torch.cat(batch_scores, dim=1)386 .to(output_dtype)387 .to(output_device)388 )389 return torch.cat(scores, dim=0)390 391 # ------------------------------------------------------------------392 # Internal helpers393 # ------------------------------------------------------------------394 395 @staticmethod396 def _left_pad_and_concat(397 batch_features: list[BatchFeature],398 ) -> BatchFeature:399 """400 Left-pad variable-length BatchFeature dicts and stack them.401 402 Qwen 3.5 yields a variable number of visual tokens per image403 (resolution dependent), so we align them before concatenation.404 Padding is on the left (decoder convention).405 """406 all_keys = batch_features[0].keys()407 concatenated: dict[str, Any] = {}408 409 for key in all_keys:410 tensors = [bf[key] for bf in batch_features]411 412 if not isinstance(tensors[0], torch.Tensor):413 concatenated[key] = tensors[0]414 continue415 416 if tensors[0].ndim < 2:417 concatenated[key] = torch.cat(tensors, dim=0)418 continue419 420 max_seq_len = max(t.shape[1] for t in tensors)421 padded: list[torch.Tensor] = []422 for t in tensors:423 pad_len = max_seq_len - t.shape[1]424 if pad_len > 0:425 zeros = torch.zeros(426 *t.shape[:1], pad_len, *t.shape[2:],427 dtype=t.dtype, device=t.device,428 )429 t = torch.cat([zeros, t], dim=1)430 padded.append(t)431 432 concatenated[key] = torch.cat(padded, dim=0)433 434 return BatchFeature(concatenated)435 436 437__all__ = ["ColQwen35BidirectionProcessor"]438 