CoolFace
Modelpublic

Writer/palmyra-vision-vllm

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes187downloads
preprocessing_molmo.py193 linesDownload Raw Back to root
1"""2Processor class for Molmo.3"""4 5from typing import Optional6 7import PIL8from PIL import ImageOps9from PIL.Image import Image10 11try:12    from typing import Unpack13except ImportError:14    from typing_extensions import Unpack15 16import numpy as np17import torch18 19from transformers.image_utils import ImageInput20from transformers.processing_utils import (21    TextKwargs,22    ProcessingKwargs,23    ProcessorMixin,24)25 26from transformers.tokenization_utils_base import TextInput, PreTokenizedInput27from transformers.utils import logging28 29from transformers import AutoTokenizer30from .image_preprocessing_molmo import MolmoImagesKwargs, MolmoImageProcessor31 32 33logger = logging.get_logger(__name__)34 35 36DEFAULT_IMAGE_PATCH_TOKEN = f"<im_patch>"37DEFAULT_IM_START_TOKEN = f"<im_start>"38DEFAULT_IM_END_TOKEN = f"<im_end>"39DEFAULT_IM_COL_TOKEN = f"<im_col>"40IMAGE_PROMPT = "<|image|>"41 42EXTRA_TOKENS = (DEFAULT_IM_START_TOKEN, DEFAULT_IM_END_TOKEN, DEFAULT_IMAGE_PATCH_TOKEN, DEFAULT_IM_COL_TOKEN, IMAGE_PROMPT)43 44 45def get_special_token_ids(tokenizer):46    ids = tokenizer.encode("".join(EXTRA_TOKENS), add_special_tokens=False)47    assert len(ids) == len(EXTRA_TOKENS)48    return {k: i for k, i in zip(EXTRA_TOKENS, ids)}49 50 51class MolmoTextKwargs(TextKwargs, total=False):52    style: Optional[str]53    system_prompt: Optional[str]54    message_format: Optional[str]55    always_start_with_space: Optional[bool]56    sequence_length: Optional[int]57 58 59class MolmoProcessorKwargs(ProcessingKwargs, total=False):60    text_kwargs: MolmoTextKwargs61    images_kwargs: MolmoImagesKwargs62    _defaults = {63        "images_kwargs": {64            "max_crops": 12,65            "overlap_margins": [4, 4],66            "base_image_input_size": [336, 336],67            "image_token_length_w": 12,68            "image_token_length_h": 12,69            "image_patch_size": 14,70            "image_padding_mask": True,71        },72        "text_kwargs": {73            "style": "long_caption",74            "system_prompt": "none",75            "message_format": "role",76            "always_start_with_space": True,77            "sequence_length": 1536,78            "padding": False,79        },80    }81 82 83class MolmoProcessor(ProcessorMixin):84    attributes = ["image_processor", "tokenizer"]85    image_processor_class = "AutoImageProcessor"86    tokenizer_class = ("Qwen2Tokenizer", "Qwen2TokenizerFast")87 88    def __init__(self, image_processor: MolmoImageProcessor = None, tokenizer : AutoTokenizer = None, **kwargs):89        # self.image_processor = image_processor90        # self.tokenizer = tokenizer91        super().__init__(image_processor, tokenizer)92        self._special_tokens = None93 94    @property95    def special_token_ids(self):96        if self._special_tokens is None:97            self._special_tokens = get_special_token_ids(self.tokenizer)98        return self._special_tokens99 100    def get_tokens_input(self, prompt, message_format, always_start_with_space):101        if message_format == "none" or message_format is None:102            pass103        elif message_format == "role":104            prompt = "User: " + prompt + " Assistant:"105        else:106            raise NotImplementedError(f"Message format {message_format} not implemented")107 108        if always_start_with_space:109            prompt = " " + prompt110 111        tokens = self.tokenizer.encode(prompt, add_special_tokens=False)112 113        return tokens114 115    def process(116        self,117        text: TextInput = None,118        images: ImageInput = None,119        *,120        tokens: Optional[PreTokenizedInput] = None,121        **kwargs: Unpack[MolmoProcessorKwargs],122    ):123        output_kwargs = self._merge_kwargs(124            MolmoProcessorKwargs,125            tokenizer_init_kwargs=self.tokenizer.init_kwargs,126            **kwargs,127        )128 129        if tokens is None:130            tokens = self.get_tokens_input(131                text,132                output_kwargs["text_kwargs"]["message_format"],133                output_kwargs["text_kwargs"]["always_start_with_space"],134            )135 136        image_token_id = self.special_token_ids[IMAGE_PROMPT]137 138        if images is not None:139            if not isinstance(images, (list, tuple)):140                images = [images]141            image_arrays = []142            for image in images:143                if isinstance(image, Image):144                    image = image.convert("RGB")145                    # Handle images with EXIF orientation tags, which PIL will ignore by default146                    # https://github.com/python-pillow/Pillow/issues/4703147                    img = ImageOps.exif_transpose(image)148                    image_arrays.append(np.array(image))149                else:150                    assert len(image.shape) == 3 and image.shape[-1] == 3151                    image_arrays.append(image.astype(np.uint8))152            images = image_arrays153            # For now only support inserting images at the start154            image_idx = [-1]*len(images)155        else:156            image_idx = None157 158        sequence_length = output_kwargs["text_kwargs"]["sequence_length"]159 160        image_patch_token_id = self.special_token_ids[DEFAULT_IMAGE_PATCH_TOKEN]161        image_col_token_id = self.special_token_ids[DEFAULT_IM_COL_TOKEN]162        image_start_token_id = self.special_token_ids[DEFAULT_IM_START_TOKEN]163        image_end_token_id = self.special_token_ids[DEFAULT_IM_END_TOKEN]164        out = self.image_processor.multimodal_preprocess(165            images=images,166            image_idx=image_idx,167            tokens=np.asarray(tokens).astype(np.int32),168            sequence_length=sequence_length,169            image_patch_token_id=image_patch_token_id,170            image_col_token_id=image_col_token_id,171            image_start_token_id=image_start_token_id,172            image_end_token_id=image_end_token_id,173            **output_kwargs["images_kwargs"]174        )175 176        # Prepend BOS177        # qwen2 and olmo do not have a BOS, and instead use EOS as a generic seperator token.178        bos = self.tokenizer.bos_token_id or self.tokenizer.eos_token_id179        decoder_input_tokens = np.pad(out["input_ids"], [[1, 0]], constant_values=bos)180        out["input_ids"] = decoder_input_tokens181        if "image_input_idx" in out:182            # Shift patch mapping up by one since we added BOS183            image_input_idx = out["image_input_idx"]184            out["image_input_idx"] = np.where(image_input_idx < 0, image_input_idx, image_input_idx + 1)185 186        for k, v in out.items():187            out[k] = torch.from_numpy(v)188 189        return out190 191 192MolmoProcessor.register_for_auto_class()193