File size: 10,539 Bytes
a381a62
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
"""Strict persistent-visual CVRR wrapper for released InternVL3 chat models."""

from __future__ import annotations

import copy
import io
import pathlib

import torch

from .source_gemma import (
    GemmaCVRR,
    LayerCall,
    _StopAfterCell,
    _capture_call,
    _gather_ids,
    _gather_rows,
    _replace_rows,
    inject_cell_lora,
)
from .source_helpers import _layer_hidden
from .source_helpers import _normalize_tiles, dynamic_tiles


class InternVLCVRR(GemmaCVRR):
    def __init__(
        self,
        model_path: str,
        *,
        ell_star: int,
        steps: int = 4,
        beta: float = 0.33,
        rank: int = 32,
        alpha: float = 12.0,
        dropout: float = 0.01,
        device: str | torch.device = "cuda:0",
        offline: bool = True,
    ):
        # Bypass GemmaCVRR.__init__, retaining its audited recurrence, adapter
        # toggling, loss, and serialization methods.
        torch.nn.Module.__init__(self)
        from transformers import AutoModel, AutoTokenizer

        import torch.distributed as dist

        self._owns_process_group = False
        if dist.is_available() and not dist.is_initialized():
            import os
            import tempfile

            rendezvous = pathlib.Path(tempfile.gettempdir()) / f"cvrr_iv_train_{os.getpid()}"
            dist.init_process_group(
                "gloo", init_method=f"file://{rendezvous}", rank=0, world_size=1
            )
            self._owns_process_group = True

        self.model_path = str(model_path)
        self.device_ref = torch.device(device)
        self.tokenizer = AutoTokenizer.from_pretrained(
            model_path,
            trust_remote_code=True,
            use_fast=False,
            local_files_only=offline,
        )
        self.base_model = AutoModel.from_pretrained(
            model_path,
            trust_remote_code=True,
            local_files_only=offline,
            low_cpu_mem_usage=True,
            use_flash_attn=False,
            dtype=torch.bfloat16,
            device_map=str(self.device_ref),
        )
        self.model_type = "internvl_chat"
        self.layers = self.base_model.language_model.model.layers
        self.ell_star = int(ell_star)
        self.cell_index = self.ell_star + 1
        self.upper_start = self.cell_index + 1
        if not 0 <= self.ell_star <= len(self.layers) - 2:
            raise ValueError("ell_star must leave a cell and upper decoder")
        if steps < 2 or not 0.0 <= beta <= 1.0:
            raise ValueError("invalid recurrence depth or beta")
        self.steps = int(steps)
        self.beta = float(beta)
        for parameter in self.base_model.parameters():
            parameter.requires_grad_(False)
        self.lora = inject_cell_lora(
            self.layers[self.cell_index],
            rank=rank,
            alpha=alpha,
            dropout=dropout,
            suffixes={"wqkv", "wo", "w1", "w2", "w3"},
        )
        for module in self.lora.values():
            module.to(self.device_ref)
        self.rank = int(rank)
        self.alpha = float(alpha)
        self.adapter_dropout = float(dropout)
        self.image_token_id = int(
            self.tokenizer.convert_tokens_to_ids("<IMG_CONTEXT>")
        )
        self.base_model.img_context_token_id = self.image_token_id
        self.base_model.eval()

    def _pad_token_id(self) -> int:
        return int(self.base_model.config.llm_config.pad_token_id)

    def _modality(self, mm_inputs):
        return mm_inputs["input_ids"].eq(self.image_token_id).long()

    def _initial_multimodal(self, mm_inputs):
        calls: dict[int, LayerCall] = {}
        captured = {}
        cell = self.layers[self.cell_index]

        def stop(_module, _args, output):
            captured["hidden"] = _layer_hidden(output).detach()
            raise _StopAfterCell

        pre = cell.register_forward_pre_hook(
            _capture_call(calls, self.cell_index), with_kwargs=True
        )
        post = cell.register_forward_hook(stop)
        try:
            with torch.no_grad(), self.adapters(False):
                try:
                    self.base_model(
                        **mm_inputs,
                        use_cache=False,
                        output_hidden_states=False,
                        return_dict=True,
                    )
                except _StopAfterCell:
                    pass
        finally:
            pre.remove()
            post.remove()
        if "hidden" not in captured or self.cell_index not in calls:
            raise RuntimeError("failed to capture InternVL recurrent cell")
        return captured["hidden"], calls[self.cell_index]

    def _text_context(self, question_ids, question_mask):
        calls: dict[int, LayerCall] = {}
        captured = {}
        handles = []
        for index in range(self.cell_index, len(self.layers)):
            handles.append(
                self.layers[index].register_forward_pre_hook(
                    _capture_call(calls, index), with_kwargs=True
                )
            )

        def capture(_module, _args, output):
            captured["anchor"] = _layer_hidden(output).detach()

        handles.append(self.layers[self.cell_index].register_forward_hook(capture))
        try:
            with torch.no_grad(), self.adapters(False):
                self.base_model.language_model(
                    input_ids=question_ids,
                    attention_mask=question_mask,
                    use_cache=False,
                    output_hidden_states=False,
                    return_dict=True,
                )
        finally:
            for handle in handles:
                handle.remove()
        missing = [
            index
            for index in range(self.cell_index, len(self.layers))
            if index not in calls
        ]
        if missing or "anchor" not in captured:
            raise RuntimeError(f"failed to capture InternVL text path: {missing}")
        return captured["anchor"], calls

    def _upper(self, state, text_calls):
        hidden = state
        with self.adapters(False):
            for index in range(self.upper_start, len(self.layers)):
                hidden = self._call_layer(
                    self.layers[index], hidden, text_calls[index]
                )
        hidden = self.base_model.language_model.model.norm(hidden)
        return self.base_model.language_model.output(hidden).float()


class InternVLArrowCollator:
    def __init__(self, model: InternVLCVRR, *, max_tiles: int = 12):
        self.tokenizer = model.tokenizer
        self.template = copy.deepcopy(model.base_model.conv_template)
        self.system_message = model.base_model.system_message
        self.num_image_token = int(model.base_model.num_image_token)
        self.image_size = int(
            model.base_model.config.force_image_size
            or model.base_model.config.vision_config.image_size
        )
        self.use_thumbnail = bool(model.base_model.config.use_thumbnail)
        self.max_tiles = int(max_tiles)

    @staticmethod
    def _pad(values, pad):
        width = max(value.shape[0] for value in values)
        output = values[0].new_full((len(values), width), pad)
        for index, value in enumerate(values):
            output[index, : value.shape[0]] = value
        return output

    def _query(self, question, hint, num_tiles):
        template = copy.deepcopy(self.template)
        template.system_message = self.system_message
        template.append_message(
            template.roles[0],
            "<image>\n" + str(question).strip() + str(hint),
        )
        template.append_message(template.roles[1], None)
        query = template.get_prompt()
        visual = (
            "<img>"
            + "<IMG_CONTEXT>" * self.num_image_token * num_tiles
            + "</img>"
        )
        return query.replace("<image>", visual, 1)

    def __call__(self, features):
        from PIL import Image

        rows = []
        for feature in features:
            raw = feature["image_bytes"]
            if isinstance(raw, memoryview):
                raw = raw.tobytes()
            with Image.open(io.BytesIO(raw)) as opened:
                tiles = dynamic_tiles(
                    opened.convert("RGB"),
                    image_size=self.image_size,
                    max_tiles=self.max_tiles,
                    thumbnail=self.use_thumbnail,
                )
            # The released InternViT does not cast inputs internally; its model
            # card explicitly converts pixel_values to bfloat16 before forward.
            pixels = _normalize_tiles(tiles).to(torch.bfloat16)
            query = self._query(
                feature["fixed_question"], feature["fixed_hint"], len(tiles)
            )
            tokenized = self.tokenizer(query, return_tensors="pt")
            prompt = tokenized.input_ids[0]
            full = self.tokenizer(
                query + str(feature["fixed_answer"]).strip(),
                return_tensors="pt",
            ).input_ids[0]
            if not torch.equal(full[: prompt.numel()], prompt):
                raise RuntimeError("InternVL answer serialization changed the prompt prefix")
            full = torch.cat(
                (full, full.new_tensor([self.tokenizer.eos_token_id]))
            )
            answer = full[prompt.numel() :]
            rows.append(
                {
                    "input_ids": full,
                    "attention_mask": torch.cat(
                        (tokenized.attention_mask[0], torch.ones_like(answer))
                    ),
                    "labels": torch.cat((torch.full_like(prompt, -100), answer)),
                    "pixel_values": pixels,
                    "image_flags": torch.ones(len(tiles), 1, dtype=torch.long),
                }
            )
        return {
            "mm_inputs": {
                "input_ids": self._pad(
                    [row["input_ids"] for row in rows], self.tokenizer.pad_token_id
                ),
                "attention_mask": self._pad(
                    [row["attention_mask"] for row in rows], 0
                ),
                "pixel_values": torch.cat(
                    [row["pixel_values"] for row in rows], dim=0
                ),
                "image_flags": torch.cat(
                    [row["image_flags"] for row in rows], dim=0
                ),
            },
            "mm_labels": self._pad([row["labels"] for row in rows], -100),
        }