Team Ai
Modelpublic

webAI-Official/webAI-ColVec1.1-8b

sourceHugging Faceotherupdated 2mo agoView on Hugging Face
8likes191downloads
processing_colqwen35_bidirection.py438 linesDownload Raw Back to root
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