GroundAnything / processing_groundinganything.py
GroundingPI's picture Skywalker0410's picture
Publish model weights, paper figures, and complete model card (#1)
37f388f
Raw History Blame
5.37 kB
"""Processor glue for GroundAnything-VLM with Kimi-K3 MoonViT preprocessing."""
from transformers.feature_extraction_utils import BatchFeature
from transformers.processing_utils import ProcessorMixin
from .media_utils import MediaInput
from .image_processing_groundinganything import GroundAnythingVLMImageProcessor
class GroundAnythingVLMProcessor(ProcessorMixin):
attributes = ["image_processor", "tokenizer"]
image_processor_class = "AutoImageProcessor"
tokenizer_class = "AutoTokenizer"
def __init__(self, image_processor=None, tokenizer=None, chat_template=None, **kwargs):
del kwargs
super().__init__(
image_processor=image_processor,
tokenizer=tokenizer,
chat_template=chat_template or getattr(tokenizer, "chat_template", None),
)
@property
def image_token(self):
return "<|image_pad|>"
@property
def image_token_id(self):
return self.tokenizer.convert_tokens_to_ids(self.image_token)
def _get_num_multimodal_tokens(self, image_sizes=None, **kwargs):
del kwargs
num_image_tokens = []
num_image_patches = []
for height, width in image_sizes or ():
image_stub = type("ImageSize", (), {"size": (width, height)})()
resize = self.image_processor.get_resize_config(
{"type": "image", "image": image_stub}
)
tokens = int(resize["num_tokens"])
num_image_tokens.append(tokens)
num_image_patches.append(tokens * self.image_processor.merge_size**2)
return {
"num_image_tokens": num_image_tokens,
"num_image_patches": num_image_patches,
}
@classmethod
def register_for_auto_class(cls, auto_class="AutoProcessor"):
cls._auto_class = auto_class
@classmethod
def from_pretrained(cls, pretrained_model_name_or_path, **kwargs):
import json
import os
from transformers import AutoTokenizer
kwargs.pop("_from_auto", None)
kwargs.pop("trust_remote_code", None)
kwargs.pop("code_revision", None)
with open(os.path.join(pretrained_model_name_or_path, "preprocessor_config.json"), encoding="utf-8") as f:
processor_config = json.load(f)
image_processor = GroundAnythingVLMImageProcessor(
media_proc_cfg=processor_config["media_proc_cfg"]
)
tokenizer = AutoTokenizer.from_pretrained(
pretrained_model_name_or_path, trust_remote_code=True, **kwargs
)
return cls(image_processor=image_processor, tokenizer=tokenizer)
def apply_chat_template(self, messages, **kwargs):
if self.chat_template and "chat_template" not in kwargs:
kwargs["chat_template"] = self.chat_template
return self.tokenizer.apply_chat_template(messages, **kwargs)
def __call__(
self,
text=None,
images=None,
return_tensors="pt",
padding=False,
**kwargs,
):
return_mm_token_type_ids = kwargs.pop("return_mm_token_type_ids", False)
if isinstance(text, str):
text = [text]
image_inputs = {}
if images is not None:
image_inputs = self.image_processor(
images=images,
return_tensors=return_tensors,
)
text = list(text)
image_index = 0
merge_length = self.image_processor.merge_size**2
for batch_index, prompt in enumerate(text):
while self.image_token in prompt:
grid = image_inputs["image_grid_thw"][image_index]
num_tokens = int(grid.prod().item()) // merge_length
prompt = prompt.replace(
self.image_token, "<|image_placeholder|>" * num_tokens, 1
)
image_index += 1
text[batch_index] = prompt.replace(
"<|image_placeholder|>", self.image_token
)
if image_index != len(image_inputs["image_grid_thw"]):
raise ValueError(
"number of image placeholders does not match image inputs"
)
text_inputs = self.tokenizer(
text,
return_tensors=return_tensors,
padding=padding,
**kwargs,
)
if return_mm_token_type_ids:
input_ids = text_inputs["input_ids"]
if hasattr(input_ids, "new_zeros"):
mm_token_type_ids = input_ids.new_zeros(input_ids.shape)
mm_token_type_ids[input_ids == self.image_token_id] = 1
else:
mm_token_type_ids = [
[int(token == self.image_token_id) for token in row]
for row in input_ids
]
text_inputs["mm_token_type_ids"] = mm_token_type_ids
return BatchFeature(data={**text_inputs, **image_inputs})
def batch_decode(self, *args, **kwargs):
return self.tokenizer.batch_decode(*args, **kwargs)
def decode(self, *args, **kwargs):
return self.tokenizer.decode(*args, **kwargs)
__all__ = ["GroundAnythingVLMProcessor"]
class GroundAnythingProcessor(GroundAnythingVLMProcessor):
"""DLM release processor identity."""