Writer/palmyra-vision-vllm
0187
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 