File size: 5,367 Bytes
37f388f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
"""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."""