Image-Text-to-Text
Transformers
Safetensors
English
qwen2_5vl_ca
feature-extraction
conversational
custom_code
ameroyer nielsr HF Staff commited on
Commit
5f7de59
·
0 Parent(s):

Super-squash branch 'main' using huggingface_hub

Browse files

Co-authored-by: nielsr <nielsr@users.noreply.huggingface.co>

.gitattributes ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ *.7z filter=lfs diff=lfs merge=lfs -text
2
+ *.arrow filter=lfs diff=lfs merge=lfs -text
3
+ *.bin filter=lfs diff=lfs merge=lfs -text
4
+ *.bz2 filter=lfs diff=lfs merge=lfs -text
5
+ *.ckpt filter=lfs diff=lfs merge=lfs -text
6
+ *.ftz filter=lfs diff=lfs merge=lfs -text
7
+ *.gz filter=lfs diff=lfs merge=lfs -text
8
+ *.h5 filter=lfs diff=lfs merge=lfs -text
9
+ *.joblib filter=lfs diff=lfs merge=lfs -text
10
+ *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
+ *.mlmodel filter=lfs diff=lfs merge=lfs -text
12
+ *.model filter=lfs diff=lfs merge=lfs -text
13
+ *.msgpack filter=lfs diff=lfs merge=lfs -text
14
+ *.npy filter=lfs diff=lfs merge=lfs -text
15
+ *.npz filter=lfs diff=lfs merge=lfs -text
16
+ *.onnx filter=lfs diff=lfs merge=lfs -text
17
+ *.ot filter=lfs diff=lfs merge=lfs -text
18
+ *.parquet filter=lfs diff=lfs merge=lfs -text
19
+ *.pb filter=lfs diff=lfs merge=lfs -text
20
+ *.pickle filter=lfs diff=lfs merge=lfs -text
21
+ *.pkl filter=lfs diff=lfs merge=lfs -text
22
+ *.pt filter=lfs diff=lfs merge=lfs -text
23
+ *.pth filter=lfs diff=lfs merge=lfs -text
24
+ *.rar filter=lfs diff=lfs merge=lfs -text
25
+ *.safetensors filter=lfs diff=lfs merge=lfs -text
26
+ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
27
+ *.tar.* filter=lfs diff=lfs merge=lfs -text
28
+ *.tar filter=lfs diff=lfs merge=lfs -text
29
+ *.tflite filter=lfs diff=lfs merge=lfs -text
30
+ *.tgz filter=lfs diff=lfs merge=lfs -text
31
+ *.wasm filter=lfs diff=lfs merge=lfs -text
32
+ *.xz filter=lfs diff=lfs merge=lfs -text
33
+ *.zip filter=lfs diff=lfs merge=lfs -text
34
+ *.zst filter=lfs diff=lfs merge=lfs -text
35
+ *tfevents* filter=lfs diff=lfs merge=lfs -text
Notice ADDED
@@ -0,0 +1,2 @@
 
 
 
1
+ CASA-Qwen2_5-VL-3B is finetuned from Qwen2.5-VL-3B with additional CASA layers.
2
+ Qwen is licensed under the Qwen LICENSE AGREEMENT, Copyright (c) Alibaba Cloud. All Rights Reserved.
README.md ADDED
@@ -0,0 +1,92 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ base_model:
3
+ - Qwen/Qwen2.5-VL-3B-Instruct
4
+ datasets:
5
+ - HuggingFaceM4/FineVision
6
+ - mvp-lab/LLaVA-OneVision-1.5-Instruct-Data
7
+ language:
8
+ - en
9
+ license: cc-by-nc-sa-4.0
10
+ pipeline_tag: image-text-to-text
11
+ library_name: transformers
12
+ ---
13
+
14
+ # CASA-Qwen2_5-VL-3B
15
+
16
+ This repository contains the model weights for **CASA-Qwen2_5-VL-3B**, introduced in the paper [CASA: Cross-Attention over Self-Attention for Efficient Vision-Language Fusion](https://huggingface.co/papers/2512.19535).
17
+
18
+ This model is a [Qwen-2.5VL-3B-Instruct](https://huggingface.co/Qwen/Qwen2.5-VL-3B-Instruct) model adapted from token insertion to a cross-attention-based architecture.
19
+
20
+ - **Paper:** [CASA: Cross-Attention over Self-Attention for Efficient Vision-Language Fusion](https://arxiv.org/abs/2512.19535)
21
+ - **Project Page:** [kyutai.org/casa](https://kyutai.org/casa)
22
+ - **Code:** [github.com/kyutai-labs/casa](https://github.com/kyutai-labs/casa)
23
+
24
+ ## Sample Usage
25
+
26
+ This model requires `trust_remote_code=True` to load the custom architecture. Below is a snippet to run inference using `transformers`.
27
+
28
+ ```python
29
+ import torch
30
+ from transformers.models.auto.modeling_auto import AutoModel
31
+ from transformers.models.auto.processing_auto import AutoProcessor
32
+
33
+ model_id = "kyutai/CASA-Qwen2_5-VL-3B"
34
+ model = AutoModel.from_pretrained(
35
+ model_id,
36
+ torch_dtype=torch.bfloat16,
37
+ attn_implementation="flash_attention_2",
38
+ trust_remote_code=True,
39
+ ).cuda()
40
+
41
+ processor = AutoProcessor.from_pretrained(
42
+ model_id,
43
+ trust_remote_code=True,
44
+ )
45
+
46
+ conversation = [
47
+ {
48
+ "role": "user",
49
+ "content": [
50
+ {
51
+ "type": "image",
52
+ "image": "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/transformers/tasks/ai2d-demo.png",
53
+ },
54
+ {
55
+ "type": "text",
56
+ "text": "Describe this image.",
57
+ },
58
+ ],
59
+ },
60
+ ]
61
+
62
+ inputs = processor.tokenize_messages(messages=conversation)
63
+ inputs = inputs.to(model.device)
64
+ input_len = inputs["input_ids"].shape[1]
65
+
66
+ output_ids = model.generate_from_image(
67
+ **inputs,
68
+ max_new_tokens=512,
69
+ pre_image_tokens=processor.pre_image_tokens,
70
+ post_image_tokens=processor.post_image_tokens,
71
+ eos_token_id=model.generation_config.eos_token_id,
72
+ )[0, input_len:]
73
+
74
+ response = processor.tokenizer.decode(output_ids, skip_special_tokens=True)
75
+ print(response)
76
+ ```
77
+
78
+ ## Citation
79
+
80
+ ```bibtex
81
+ @article{kyutai2025casa,
82
+ author = {Moritz B\"ohle and Am\'elie Royer and Juliette Marrie and Edouard Grave and Patrick P\'erez},
83
+ year = {2025},
84
+ title = {CASA: Cross-Attention over Self-Attention for Efficient Vision-Language Fusion},
85
+ journal = {ArXiv},
86
+ url = {https://arxiv.org/abs/2512.19535}
87
+ }
88
+ ```
89
+
90
+ ## License
91
+
92
+ The code in the official repository is provided under the **MIT license**. The weights for this model are released under the **CC-BY-NC-SA 4.0 license**. Additionally, as this model includes weights from Qwen2.5-VL-3B, it is subject to the [Qwen RESEARCH LICENSE AGREEMENT](https://huggingface.co/Qwen/Qwen2.5-VL-3B-Instruct/blob/main/LICENSE).
config.json ADDED
@@ -0,0 +1,80 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "V2Qwen2_5VL"
4
+ ],
5
+ "auto_map": {
6
+ "AutoConfig": "configuration_qwen2_5vl_ca.Qwen2_5_VLCAConfig",
7
+ "AutoModel": "modeling_qwen2_5vl_ca.V2Qwen2_5VL"
8
+ },
9
+ "attention_dropout": 0.0,
10
+ "bos_token_id": 151643,
11
+ "cross_attention": true,
12
+ "head_dim": 128,
13
+ "hidden_act": "silu",
14
+ "hidden_size": 2048,
15
+ "image_token_id": 151655,
16
+ "initializer_range": 0.02,
17
+ "intermediate_size": 11008,
18
+ "max_position_embeddings": 128000,
19
+ "max_window_layers": 70,
20
+ "model_type": "qwen2_5vl_ca",
21
+ "num_attention_heads": 16,
22
+ "num_hidden_layers": 36,
23
+ "num_key_value_heads": 2,
24
+ "rms_norm_eps": 1e-06,
25
+ "rope_scaling": {
26
+ "mrope_section": [
27
+ 16,
28
+ 24,
29
+ 24
30
+ ],
31
+ "rope_type": "default",
32
+ "type": "default"
33
+ },
34
+ "rope_theta": 1000000.0,
35
+ "sliding_window": 32768,
36
+ "tie_word_embeddings": true,
37
+ "torch_dtype": "bfloat16",
38
+ "transformers_version": "4.51.3",
39
+ "use_cache": true,
40
+ "use_sliding_window": false,
41
+ "video_token_id": 151656,
42
+ "vision_config": {
43
+ "depth": 32,
44
+ "fullatt_block_indexes": [
45
+ 7,
46
+ 15,
47
+ 23,
48
+ 31
49
+ ],
50
+ "hidden_act": "silu",
51
+ "hidden_size": 1280,
52
+ "image_mean": [
53
+ 0.48145466,
54
+ 0.4578275,
55
+ 0.40821073
56
+ ],
57
+ "image_std": [
58
+ 0.26862954,
59
+ 0.26130258,
60
+ 0.27577711
61
+ ],
62
+ "in_channels": 3,
63
+ "in_chans": 3,
64
+ "intermediate_size": 3420,
65
+ "model_type": "qwen2_5_vl",
66
+ "num_heads": 16,
67
+ "out_hidden_size": 2048,
68
+ "patch_size": 14,
69
+ "spatial_merge_size": 2,
70
+ "spatial_patch_size": 14,
71
+ "temporal_patch_size": 1,
72
+ "tokens_per_second": 2,
73
+ "window_size": 112
74
+ },
75
+ "vision_end_token_id": 151653,
76
+ "vision_start_token_id": 151652,
77
+ "vision_token_id": 151654,
78
+ "vocab_size": 151936,
79
+ "xa_layers": []
80
+ }
configuration_qwen2_5vl_ca.py ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Any
2
+
3
+ from transformers.models.qwen2_5_vl.configuration_qwen2_5_vl import Qwen2_5_VLConfig
4
+
5
+
6
+ class Qwen2_5_VLCAConfig(Qwen2_5_VLConfig):
7
+ """Qwen2.5-VL config augmented with cross-attention options"""
8
+
9
+ model_type = "qwen2_5vl_ca"
10
+
11
+ def __init__(
12
+ self,
13
+ *args: Any,
14
+ # Our fusion mechanism
15
+ xa_layers: None | tuple = None,
16
+ cross_attention: bool = False,
17
+ xa_share_qkvo: bool = False,
18
+ **kwargs: Any,
19
+ ):
20
+ super().__init__(*args, **kwargs)
21
+ self.head_dim = self.hidden_size // self.num_attention_heads
22
+ # Our fusion mechanisms
23
+ self.xa_layers = xa_layers
24
+ self.cross_attention = cross_attention
25
+ self.xa_share_qkvo = xa_share_qkvo
26
+ # Derived from vision_start/end_token_id set by parent
27
+ self.pre_image_tokens = (
28
+ [self.vision_start_token_id]
29
+ if getattr(self, "vision_start_token_id", None) is not None
30
+ else []
31
+ )
32
+ self.post_image_tokens = (
33
+ [self.vision_end_token_id]
34
+ if getattr(self, "vision_end_token_id", None) is not None
35
+ else []
36
+ )
cross_attention.py ADDED
@@ -0,0 +1,396 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Modular cross-attention module"""
2
+
3
+ import bisect
4
+ import copy
5
+ from itertools import accumulate
6
+ from typing import TYPE_CHECKING, Literal, TypedDict, TypeVar
7
+ from typing import cast as type_cast
8
+
9
+ import torch
10
+
11
+ from .utils import StreamingModule, StreamingState
12
+
13
+ if TYPE_CHECKING:
14
+ from transformers.configuration_utils import PretrainedConfig
15
+
16
+
17
+ try:
18
+ from flash_attn import flash_attn_varlen_func
19
+ except ImportError:
20
+ flash_attn_varlen_func = None # type: ignore
21
+
22
+
23
+ AttentionBaseT = TypeVar("AttentionBaseT", bound=StreamingModule)
24
+
25
+
26
+ WindowsComputeKwargs = TypedDict(
27
+ "WindowsComputeKwargs",
28
+ {
29
+ "num_post_image_tokens": int,
30
+ "num_pre_image_tokens": int,
31
+ },
32
+ total=False,
33
+ )
34
+
35
+
36
+ def get_sample_lengths_for_xa(
37
+ image_embeds_insertion_points: list[torch.Tensor],
38
+ image_embeds: torch.Tensor | list[torch.Tensor] | None,
39
+ total_seq_len: int,
40
+ attention_mask: torch.Tensor | None = None,
41
+ **kwargs: WindowsComputeKwargs,
42
+ ) -> tuple[list[tuple[int, bool]], list[int], torch.Tensor | None]:
43
+ """Sample lengths for cross-attention. Compared to other functions in this file,
44
+ it also returns a mask on the text tokens to mark tokens which do not relate
45
+ to any images (e.g. squashed text-only-samples, or BoS in prefix_after_bos)
46
+ """
47
+ num_post_image_tokens = type_cast(int, kwargs.get("num_post_image_tokens", 0))
48
+ num_pre_image_tokens = type_cast(int, kwargs.get("num_pre_image_tokens", 0))
49
+ squashed_samples_lengths = type_cast(
50
+ list[list[int]] | None, kwargs.get("squashed_samples_lengths", None)
51
+ )
52
+ if squashed_samples_lengths is not None:
53
+ assert len(squashed_samples_lengths) == len(image_embeds_insertion_points)
54
+
55
+ def __insert_next_sample__(
56
+ batch_idx: int, insrt_pt: int, last_insrt_pt: int, end_of_batch_sample: bool = False
57
+ ) -> None:
58
+ nonlocal attention_mask, active_tokens
59
+ nonlocal text_sample_lengths, full_sample_lengths
60
+ nonlocal cum_samples_lengths, current_image_offset
61
+ # Add the sample between [last_insrt_pt, insrt_pt] with breaks in
62
+ # between any squashed samples we find on the way
63
+ nonlocal last_image_idx, current_image_idx, current_length
64
+ added_sample = False
65
+ start_pt = bisect.bisect_left(cum_samples_lengths, last_insrt_pt)
66
+ for end_of_sample in cum_samples_lengths[start_pt:]:
67
+ # we will break the loop at the end when end_of_sample = insrt_pt
68
+ end_of_sample = min(end_of_sample, insrt_pt)
69
+
70
+ # Add between [last_insrt_pt, end_of_sample]
71
+ current_length = end_of_sample - last_insrt_pt
72
+ num_padding_tokens = 0
73
+ if attention_mask is not None:
74
+ num_padding_tokens = int(
75
+ torch.sum(~attention_mask[batch_idx, last_insrt_pt:end_of_sample]).item()
76
+ )
77
+ current_length -= num_padding_tokens
78
+ num_image_tokens = 0
79
+ if current_length > 0:
80
+ # add image tokens to current_length
81
+ added_sample = True
82
+ if current_image_idx > 0 and image_embeds is not None:
83
+ images_in_sample = [
84
+ img_idx
85
+ for img_idx in range(last_image_idx, current_image_idx)
86
+ if img_idx < len(image_embeds_insertion_points[batch_idx])
87
+ and last_insrt_pt
88
+ <= image_embeds_insertion_points[batch_idx][img_idx]
89
+ < end_of_sample
90
+ ]
91
+ if len(images_in_sample) > 0:
92
+ num_image_tokens = sum(
93
+ _x.shape[0]
94
+ for _x in image_embeds[
95
+ current_image_offset + images_in_sample[0] : current_image_offset
96
+ + images_in_sample[-1]
97
+ + 1
98
+ ]
99
+ )
100
+ # If no image, we should not insert and instead make it as inactive
101
+ if num_image_tokens > 0:
102
+ text_sample_lengths.append(
103
+ (current_length, end_of_batch_sample and insrt_pt == end_of_sample)
104
+ )
105
+ full_sample_lengths.append(current_length + num_image_tokens)
106
+ # Active tokens
107
+ active_tokens += [int(num_image_tokens > 0)] * (current_length + num_padding_tokens)
108
+
109
+ # prepare for next loop
110
+ last_insrt_pt = end_of_sample
111
+ if end_of_sample == insrt_pt:
112
+ break
113
+ # End of loop: catching edge case where we end up on a span full of padding
114
+ if end_of_batch_sample:
115
+ assert added_sample, "Weird edge case. Don't do that, thank you"
116
+ text_sample_lengths[-1] = (text_sample_lengths[-1][0], True)
117
+
118
+ current_image_offset = 0
119
+ text_sample_lengths, full_sample_lengths = [], []
120
+ cum_samples_lengths: list[int] = []
121
+ active_tokens = []
122
+ current_length, last_insrt_pt, last_image_idx, current_image_idx = 0, 0, 0, 0
123
+ for batch_idx, pts in enumerate(image_embeds_insertion_points):
124
+ if squashed_samples_lengths is not None:
125
+ cum_samples_lengths = list(accumulate(squashed_samples_lengths[batch_idx]))
126
+ else:
127
+ cum_samples_lengths = [total_seq_len]
128
+ for current_image_idx, insrt_pt in enumerate(pts.cpu().tolist()):
129
+ # check if the images are consecutive in which way we want
130
+ # them to belong to the same window
131
+ if current_image_idx >= 1 and insrt_pt == (
132
+ image_embeds_insertion_points[batch_idx][current_image_idx - 1]
133
+ + num_pre_image_tokens
134
+ + num_post_image_tokens
135
+ ):
136
+ continue
137
+ # Otherwise, we found a new sample
138
+ # not very important but for completeness: the insertion points come *after*
139
+ # the pre-image tokens per design but for the document-id mask it is more consistent to
140
+ # have them correspond to the same image
141
+ insrt_pt -= num_pre_image_tokens
142
+
143
+ # Compute length between the two insertion points
144
+ current_length = insrt_pt - last_insrt_pt
145
+ if attention_mask is not None:
146
+ current_length -= int(
147
+ torch.sum(~attention_mask[batch_idx, last_insrt_pt:insrt_pt]).item()
148
+ )
149
+
150
+ # Update text and full sample lengths
151
+ if insrt_pt > last_insrt_pt:
152
+ __insert_next_sample__(
153
+ batch_idx, insrt_pt, last_insrt_pt, end_of_batch_sample=False
154
+ )
155
+ last_image_idx = current_image_idx
156
+ last_insrt_pt = insrt_pt
157
+
158
+ # End of batch: add sample in progress and reset
159
+ current_image_idx += 1
160
+ if cum_samples_lengths[-1] > last_insrt_pt:
161
+ __insert_next_sample__(
162
+ batch_idx, cum_samples_lengths[-1], last_insrt_pt, end_of_batch_sample=True
163
+ )
164
+ current_length, last_insrt_pt, last_image_idx, current_image_idx = 0, 0, 0, 0
165
+ current_image_offset += len(pts)
166
+
167
+ # Sample lengths
168
+ if image_embeds is None:
169
+ return text_sample_lengths, full_sample_lengths, None
170
+
171
+ return (
172
+ text_sample_lengths,
173
+ full_sample_lengths,
174
+ torch.tensor(active_tokens, dtype=torch.bool, device=image_embeds[0].device),
175
+ )
176
+
177
+
178
+ class CrossAttentionHandler:
179
+ def __init__(
180
+ self,
181
+ inputs_embeds: torch.Tensor,
182
+ image_embeds: torch.Tensor | list[torch.Tensor],
183
+ image_embeds_insertion_points: list[torch.Tensor] | None,
184
+ # info for building text->image windows link
185
+ ca_windows_info: None | WindowsComputeKwargs = None,
186
+ training: bool = True,
187
+ ):
188
+ if image_embeds_insertion_points is None:
189
+ image_embeds_insertion_points = [
190
+ torch.tensor(
191
+ [0] * len(image_embeds), # type: ignore[arg-type]
192
+ dtype=torch.long,
193
+ device=image_embeds[0].device, # type: ignore[index]
194
+ )
195
+ ]
196
+
197
+ # Create cu_seq_lens for queries (text tokens)
198
+ # Compute sample lengths based on image insertion points to get cu_seq_lens
199
+ text_sample_lengths, full_sample_lengths, self.active_tokens_mask = (
200
+ get_sample_lengths_for_xa(
201
+ image_embeds_insertion_points=image_embeds_insertion_points,
202
+ image_embeds=image_embeds,
203
+ total_seq_len=inputs_embeds.shape[1],
204
+ **(ca_windows_info or {}), # pyright: ignore[reportArgumentType]
205
+ )
206
+ )
207
+
208
+ if self.active_tokens_mask is None:
209
+ self.active_tokens_mask = torch.zeros(
210
+ (inputs_embeds.shape[0] * inputs_embeds.shape[1],),
211
+ dtype=torch.bool,
212
+ device=inputs_embeds.device,
213
+ )
214
+ assert sum(_x[0] for _x in text_sample_lengths) == int(
215
+ torch.sum(self.active_tokens_mask).item()
216
+ ), "Sanity check"
217
+
218
+ self.cu_seqlens_q = torch.Tensor(
219
+ list(accumulate([_x[0] for _x in text_sample_lengths], initial=0))
220
+ ).to(dtype=torch.int32, device=inputs_embeds.device)
221
+ self.max_seqlen_q = max(_x[0] for _x in text_sample_lengths)
222
+
223
+ # Create cu_seq_lens for keys values (the image tokens) while grouping the
224
+ # images which are consecutive
225
+ image_lens = [(l2 - l1) for (l1, _), l2 in zip(text_sample_lengths, full_sample_lengths)]
226
+ self.cu_seqlens_kv = torch.Tensor(list(accumulate(image_lens, initial=0))).to(
227
+ dtype=torch.int32, device=inputs_embeds.device
228
+ )
229
+ self.max_seqlen_kv = max(image_lens)
230
+ self.image_embeds = torch.cat([_x for _x in image_embeds], dim=0)[None, :, :]
231
+
232
+ def get_active_tokens(self, hidden_states: torch.Tensor) -> torch.Tensor:
233
+ """Return tokens to be used as queries while ignoring
234
+ tokens who have nothing to do with an image"""
235
+ channels = hidden_states.shape[-1]
236
+ src = hidden_states.flatten(0, 1)
237
+ if self.active_tokens_mask is None:
238
+ return src.reshape((1, -1, channels))
239
+ return torch.masked_select(src, self.active_tokens_mask[:, None]).reshape((1, -1, channels))
240
+
241
+ def replace_active_tokens(
242
+ self, token_updates: torch.Tensor, hidden_states_in: torch.Tensor
243
+ ) -> torch.Tensor:
244
+ if self.active_tokens_mask is None:
245
+ return token_updates
246
+ updates = torch.zeros_like(hidden_states_in.flatten(0, 1))
247
+ updates.masked_scatter_(source=token_updates, mask=self.active_tokens_mask[:, None])
248
+ return updates
249
+
250
+
251
+ def tie_qkvo_projections(self_attn: torch.nn.Module, cross_attn: "CrossAttention") -> None:
252
+ """Alias the q/k/v/o projections of `cross_attn` onto those of `self_attn`.
253
+
254
+ Used in the `xa_share_qkvo` setting so the cross-attention reuses the host
255
+ self-attention's projection weights (a single set of parameters). The two attention
256
+ calls and their independent softmaxes are otherwise unchanged.
257
+
258
+ :param self_attn: the host self-attention module owning the q/k/v/o projections
259
+ :param cross_attn: the cross-attention module whose projections are aliased
260
+ """
261
+ for proj in ("q_proj", "k_proj", "v_proj", "o_proj"):
262
+ setattr(cross_attn, proj, getattr(self_attn, proj))
263
+
264
+
265
+ class CrossAttention(StreamingModule[StreamingState]):
266
+ """Attention module between images and text tokens"""
267
+
268
+ def __init__(
269
+ self,
270
+ config: "PretrainedConfig",
271
+ layer_idx: int | None,
272
+ input_layernorm: torch.nn.Module | None = None,
273
+ ):
274
+ super().__init__(StreamingState)
275
+ self.head_dim = config.head_dim
276
+ self.config = config
277
+
278
+ self.is_first_ca_layer = layer_idx == (min(config.xa_layers) if config.xa_layers else 0)
279
+
280
+ # When weights are shared with the host self-attention, the q/k/v/o projections
281
+ # are not owned by this module; they are aliased onto the self-attn projections
282
+ # by the host model (see tie_qkvo_projections).
283
+ self.xa_share_qkvo: bool = getattr(config, "xa_share_qkvo", False)
284
+ if not self.xa_share_qkvo:
285
+ self.q_proj = self.init_from_config_proj("q", config)
286
+ self.k_proj = self.init_from_config_proj("k", config)
287
+ self.v_proj = self.init_from_config_proj("v", config)
288
+ self.o_proj = self.init_from_config_proj("o", config)
289
+
290
+ self.norm: torch.nn.Module | None = copy.deepcopy(input_layernorm)
291
+ self.cross_attention_handler = None
292
+ # (source image embeddings, projected keys, projected values); the projections
293
+ # are reused across decoding steps as long as the source tensor is unchanged
294
+ self._cached_image_kv: tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None = None
295
+
296
+ def init_from_config_proj(
297
+ self, key: Literal["q", "o", "k", "v"], config: "PretrainedConfig"
298
+ ) -> torch.nn.Linear:
299
+ """Initialize the Linear proj in this module"""
300
+ num_heads = config.num_key_value_heads if key in {"k", "v"} else config.num_attention_heads
301
+ return torch.nn.Linear(
302
+ config.hidden_size,
303
+ num_heads * config.head_dim,
304
+ bias=config.attention_bias if key != "o" else False,
305
+ )
306
+
307
+ def reset_streaming(self):
308
+ super().reset_streaming()
309
+ self._cached_image_kv = None
310
+
311
+ def forward( # pyright: ignore[reportIncompatibleMethodOverride]
312
+ self, hidden_states: torch.Tensor, cross_attention_handler: CrossAttentionHandler | None
313
+ ) -> torch.Tensor | None:
314
+ if self.is_streaming:
315
+ if self.cross_attention_handler is None:
316
+ self.cross_attention_handler = cross_attention_handler
317
+ else:
318
+ # extend the shared handler
319
+ cross_attention_handler = self.cross_attention_handler
320
+ if self.is_first_ca_layer:
321
+ cross_attention_handler.active_tokens_mask = None
322
+ cross_attention_handler.cu_seqlens_q = torch.tensor(
323
+ range(0, hidden_states.shape[0] + 1),
324
+ dtype=cross_attention_handler.cu_seqlens_q.dtype,
325
+ device=cross_attention_handler.cu_seqlens_q.device,
326
+ )
327
+
328
+ # Case of text-only samples (or inference when no handler was cached)
329
+ # in this case we just want to skip cross-attention
330
+ if cross_attention_handler is None:
331
+ return None
332
+
333
+ # Currently, we assume that images never change/get added on the fly at inference
334
+ if self.is_streaming and self.streaming_state.offset > 0:
335
+ assert hidden_states.shape[0] == (len(cross_attention_handler.cu_seqlens_kv) - 1)
336
+
337
+ og_dtype = hidden_states.dtype
338
+ og_shape = hidden_states.shape
339
+
340
+ # kv inputs: (1, num_total_image_tokens, dim)
341
+ q_inputs = cross_attention_handler.get_active_tokens(hidden_states)
342
+ kv_inputs = cross_attention_handler.image_embeds
343
+
344
+ if self.norm is not None:
345
+ q_inputs = self.norm(q_inputs)
346
+ assert q_inputs.shape[0] == kv_inputs.shape[0] == 1
347
+
348
+ # Compute QKV for the blockwise attention
349
+ bs = 1
350
+ hidden_shape_q = (bs, q_inputs.shape[1], -1, self.head_dim)
351
+ query_states = self.q_proj(q_inputs).view(*hidden_shape_q)
352
+
353
+ # The image keys/values are identical at every decoding step, so at inference
354
+ # they are computed on the first call and reused until the images change
355
+ if self._cached_image_kv is not None and self._cached_image_kv[0] is kv_inputs:
356
+ _, key_states, value_states = self._cached_image_kv
357
+ else:
358
+ normed_kv = self.norm(kv_inputs) if self.norm is not None else kv_inputs
359
+ hidden_shape_kv = (bs, kv_inputs.shape[1], -1, self.head_dim)
360
+ key_states = self.k_proj(normed_kv).view(*hidden_shape_kv)
361
+ value_states = self.v_proj(normed_kv).view(*hidden_shape_kv)
362
+ if self.is_streaming:
363
+ self._cached_image_kv = (kv_inputs, key_states, value_states)
364
+
365
+ assert flash_attn_varlen_func is not None, (
366
+ "flash_attention is not installed but required for block-wise attention"
367
+ )
368
+
369
+ assert cross_attention_handler.cu_seqlens_q[-1] == query_states.shape[1], (
370
+ f"{cross_attention_handler.cu_seqlens_q[-1]} != {query_states.shape[1]}"
371
+ )
372
+ attn_output: torch.Tensor = flash_attn_varlen_func(
373
+ query_states[0].to(torch.bfloat16),
374
+ key_states[0].to(torch.bfloat16),
375
+ value_states[0].to(torch.bfloat16),
376
+ cu_seqlens_q=cross_attention_handler.cu_seqlens_q,
377
+ cu_seqlens_k=cross_attention_handler.cu_seqlens_kv,
378
+ max_seqlen_q=cross_attention_handler.max_seqlen_q,
379
+ max_seqlen_k=cross_attention_handler.max_seqlen_kv,
380
+ dropout_p=0.0,
381
+ # No need for causality when cross-attending to image tokens since
382
+ # image tokens are never padded
383
+ causal=False,
384
+ ).to(og_dtype)
385
+
386
+ attn_output = attn_output.reshape(hidden_shape_q[1], -1).contiguous()
387
+ attn_output = self.o_proj(attn_output)
388
+
389
+ # Reshape from flattened to non-flattened
390
+ attn_output = cross_attention_handler.replace_active_tokens(attn_output, hidden_states)
391
+ attn_output = attn_output.reshape(og_shape)
392
+
393
+ if self.is_streaming:
394
+ self.streaming_state.offset += attn_output.shape[1]
395
+
396
+ return attn_output
generation_config.json ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "bos_token_id": 151643,
4
+ "pad_token_id": 151643,
5
+ "eos_token_id": [
6
+ 151643,
7
+ 151645
8
+ ],
9
+ "transformers_version": "4.51.3"
10
+ }
image_encoder.py ADDED
@@ -0,0 +1,90 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Qwen2.5VL encoder with delayed normalization"""
2
+
3
+ import torch
4
+ from einops import rearrange
5
+ from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
6
+ Qwen2_5_VisionTransformerPretrainedModel,
7
+ )
8
+
9
+
10
+ class _LinearPatchEmbedProj(torch.nn.Module):
11
+ """Replaces Conv3d in Qwen's patch embed after collapsing the temporal dim.
12
+
13
+ The training codebase replaces the Conv3d with a plain Linear for efficiency
14
+ when temporal_patch_size=1. The checkpoint therefore stores the weight under
15
+ ``patch_embed.proj.linear.weight`` rather than ``patch_embed.proj.weight``,
16
+ so this wrapper is required for state-dict key compatibility.
17
+ """
18
+
19
+ def __init__(self, in_features: int, out_features: int) -> None:
20
+ super().__init__()
21
+ self.linear = torch.nn.Linear(in_features, out_features, bias=False)
22
+
23
+ @property
24
+ def weight(self) -> torch.Tensor:
25
+ return self.linear.weight
26
+
27
+ def forward(self, x: torch.Tensor) -> torch.Tensor:
28
+ n = x.shape[0]
29
+ return self.linear(x.reshape(n, -1)).view(n, -1, 1, 1, 1)
30
+
31
+
32
+ def _replace_patch_embed_proj(visual: "Qwen2_5_VisionTransformerPretrainedModel") -> None:
33
+ """Swap the Conv3d patch-embed proj to a Linear so parameter names match the checkpoint."""
34
+ proj = visual.patch_embed.proj
35
+ if not isinstance(proj, torch.nn.Conv3d):
36
+ return
37
+ _, in_ch, _t, p, _ = proj.weight.shape
38
+ visual.patch_embed.temporal_patch_size = 1
39
+ visual.patch_embed.proj = _LinearPatchEmbedProj(in_ch * p * p, proj.out_channels) # type: ignore[assignment]
40
+
41
+
42
+ def prepare_for_qwen_encoder(
43
+ x: torch.Tensor | list[torch.Tensor], mean: torch.Tensor, std: torch.Tensor
44
+ ) -> tuple[torch.Tensor, torch.Tensor]:
45
+ """
46
+ Preprocessing for Qwen encoder
47
+ Image mean and std come from processor.image_processor.image_mean and image_std
48
+ """
49
+ grid_thw = torch.Tensor([[1, img.shape[0], img.shape[1]] for img in x]).to(x[0].device)
50
+ hws_flatten_shape = torch.prod(grid_thw, dim=-1)
51
+ x = torch.cat(
52
+ [img.reshape((int(hws_flatten_shape[idx].item()), -1)) for idx, img in enumerate(x)],
53
+ dim=0,
54
+ )
55
+ assert x.min() >= 0.0 and x.max() <= 1.0
56
+ og_shape = x.shape
57
+ x = rearrange(x, "L (c d) -> L c d", c=3)
58
+ x = (x - mean) / std
59
+ x = x.view(og_shape).to(torch.bfloat16)
60
+ return x, grid_thw
61
+
62
+
63
+ class Qwen25VLEncoder(torch.nn.Module):
64
+ """Qwen2.5 VL encoder with pre/post processing compatible with our CA implementation"""
65
+
66
+ def __init__(
67
+ self,
68
+ visual: "Qwen2_5_VisionTransformerPretrainedModel",
69
+ ):
70
+ super().__init__()
71
+ self.visual = visual
72
+ # Match training checkpoint: Conv3d is replaced by a Linear for t=1 images
73
+ _replace_patch_embed_proj(self.visual)
74
+ self.image_mean = torch.tensor(self.visual.config.image_mean).view(1, 3, 1)
75
+ self.image_std = torch.tensor(self.visual.config.image_std).view(1, 3, 1)
76
+
77
+ def forward(
78
+ self, x: torch.Tensor | list[torch.Tensor]
79
+ ) -> dict[str, torch.Tensor | list[torch.Tensor]]:
80
+ x, grid_thw = prepare_for_qwen_encoder(
81
+ x, mean=self.image_mean.to(x[0].device), std=self.image_std.to(x[0].device)
82
+ )
83
+
84
+ grid_thw = grid_thw.type(torch.int)
85
+ assert len(x) == grid_thw.prod(dim=1).sum()
86
+ out = self.visual(x, grid_thw=grid_thw)
87
+
88
+ split_sizes = (grid_thw.prod(dim=-1) // self.visual.spatial_merge_size**2).tolist()
89
+ embeds = list(torch.split(out, split_sizes, dim=0)) # Ni * (seq, C)
90
+ return {"image_embeds": embeds, "grid_thw": grid_thw}
language_qwen2_5vl_ca.py ADDED
@@ -0,0 +1,128 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Literal, Optional
2
+
3
+ import torch
4
+ from transformers.cache_utils import Cache
5
+ from transformers.configuration_utils import PretrainedConfig
6
+ from transformers.models.qwen2.modeling_qwen2 import Qwen2RMSNorm
7
+ from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
8
+ Qwen2_5_VLDecoderLayer,
9
+ Qwen2_5_VLFlashAttention2,
10
+ )
11
+
12
+ from .cross_attention import (
13
+ CrossAttention,
14
+ CrossAttentionHandler,
15
+ tie_qkvo_projections,
16
+ )
17
+ from .configuration_qwen2_5vl_ca import Qwen2_5_VLCAConfig
18
+
19
+
20
+ class QwenCrossAttention(CrossAttention):
21
+ """A CrossAttention layer compatible with Qwen's projection conventions"""
22
+
23
+ def __init__(
24
+ self,
25
+ config: Qwen2_5_VLCAConfig,
26
+ layer_idx: int | None,
27
+ ):
28
+ super().__init__(config, layer_idx) # pyright: ignore[reportArgumentType]
29
+ self.norm = Qwen2RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
30
+ assert config.rope_scaling is not None
31
+ self.mrope_section = config.rope_scaling["mrope_section"] * 2
32
+
33
+ def init_from_config_proj(
34
+ self, key: Literal["q", "o", "k", "v"], config: PretrainedConfig
35
+ ) -> torch.nn.Linear:
36
+ """Follows modeling_qwen2_5_vl.py initialization"""
37
+ head_dim = config.hidden_size // config.num_attention_heads
38
+ if key == "q":
39
+ return torch.nn.Linear(
40
+ config.hidden_size, config.num_attention_heads * head_dim, bias=True
41
+ )
42
+ if key in {"k", "v"}:
43
+ return torch.nn.Linear(
44
+ config.hidden_size, config.num_key_value_heads * head_dim, bias=True
45
+ )
46
+ if key == "o":
47
+ return torch.nn.Linear(
48
+ config.num_attention_heads * config.head_dim, config.hidden_size, bias=False
49
+ )
50
+ raise NotImplementedError(f"Unknown key {key}")
51
+
52
+
53
+ class Qwen2_5_VLAttention_CrossAttention(Qwen2_5_VLFlashAttention2):
54
+ """
55
+ Qwen Attention with extra CrossAttention layer
56
+ """
57
+
58
+ def __init__(
59
+ self,
60
+ config: Qwen2_5_VLCAConfig,
61
+ layer_idx: Optional[int] = None,
62
+ input_layernorm: torch.nn.Module | None = None,
63
+ ):
64
+ super().__init__(config, layer_idx) # pyright: ignore[reportArgumentType]
65
+ self.cross_attn = QwenCrossAttention(config, layer_idx=layer_idx)
66
+ self.cross_attention_handler: CrossAttentionHandler | None = None
67
+ if getattr(config, "xa_share_qkvo", False):
68
+ tie_qkvo_projections(self, self.cross_attn)
69
+
70
+ @classmethod
71
+ def from_qwen2_5_vl_attention(
72
+ cls, attention: Qwen2_5_VLFlashAttention2, input_layernorm: torch.nn.Module | None
73
+ ):
74
+ """Init this layer from an existing Qwen Attention layer"""
75
+ layer_idx = attention.layer_idx
76
+ assert layer_idx is not None
77
+ new_attention = cls(attention.config, layer_idx=layer_idx, input_layernorm=input_layernorm) # pyright: ignore
78
+ new_attention.load_state_dict(attention.state_dict(), strict=False)
79
+ return new_attention
80
+
81
+ def forward( # pyright: ignore[reportIncompatibleMethodOverride]
82
+ self,
83
+ hidden_states: torch.Tensor,
84
+ attention_mask: Optional[torch.Tensor] = None,
85
+ position_ids: Optional[torch.LongTensor] = None,
86
+ past_key_value: Optional[Cache] = None,
87
+ output_attentions: bool = False,
88
+ use_cache: bool = False,
89
+ cache_position: Optional[torch.LongTensor] = None,
90
+ position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None,
91
+ ):
92
+ attn_output, attn_weights, past_key_values = super().forward(
93
+ hidden_states,
94
+ attention_mask,
95
+ position_ids,
96
+ past_key_value,
97
+ output_attentions,
98
+ use_cache,
99
+ cache_position,
100
+ position_embeddings,
101
+ )
102
+ if self.cross_attn is not None:
103
+ ca_out = self.cross_attn(
104
+ hidden_states=hidden_states,
105
+ cross_attention_handler=self.cross_attention_handler,
106
+ )
107
+ # ca_out is None when there is no handler (text-only or streaming non-first call)
108
+ if ca_out is not None:
109
+ attn_output = ca_out + attn_output
110
+ return attn_output, attn_weights, past_key_values
111
+
112
+
113
+ def maybe_replace_with_cross_attention_layers(
114
+ m: torch.nn.Module, xa_layers: tuple[int, ...] | None, reindex: bool = False
115
+ ):
116
+ """Replace Attention layer by CrossAttention layer as needed"""
117
+ if isinstance(m, Qwen2_5_VLDecoderLayer):
118
+ layer_idx = m.self_attn.layer_idx
119
+ assert layer_idx is not None
120
+ if xa_layers is None or len(xa_layers) == 0 or layer_idx in xa_layers:
121
+ m.self_attn = Qwen2_5_VLAttention_CrossAttention.from_qwen2_5_vl_attention(
122
+ m.self_attn, input_layernorm=m.input_layernorm
123
+ )
124
+ elif reindex:
125
+ # shift left by number of cross-attention layers before this one
126
+ logical_idx = layer_idx - sum(j < layer_idx for j in (xa_layers or ()))
127
+ assert logical_idx >= 0
128
+ m.self_attn.layer_idx = logical_idx
merges.txt ADDED
The diff for this file is too large to render. See raw diff
 
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0831fdee11174fc55e3bf54d8b107d3efc299cb5bfb2c373f1230894f0be8721
3
+ size 8187682152
modeling_qwen2_5vl_ca.py ADDED
@@ -0,0 +1,304 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from functools import partial
2
+ from typing import Any, Sequence
3
+ from typing import cast as type_cast
4
+
5
+ import torch
6
+ from transformers.cache_utils import DynamicCache
7
+ from transformers.generation.utils import GenerateOutput
8
+ from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
9
+ Qwen2_5_VLCausalLMOutputWithPast,
10
+ Qwen2_5_VLForConditionalGeneration,
11
+ )
12
+
13
+ from .cross_attention import CrossAttentionHandler, tie_qkvo_projections
14
+ from .image_encoder import Qwen25VLEncoder
15
+ from .configuration_qwen2_5vl_ca import Qwen2_5_VLCAConfig
16
+ from .language_qwen2_5vl_ca import (
17
+ Qwen2_5_VLAttention_CrossAttention,
18
+ QwenCrossAttention,
19
+ maybe_replace_with_cross_attention_layers,
20
+ )
21
+
22
+
23
+ class V2Qwen2_5VL(Qwen2_5_VLForConditionalGeneration): # pyright: ignore[reportIncompatibleMethodOverride]
24
+ config_class = Qwen2_5_VLCAConfig
25
+
26
+ def __init__(self, config: Qwen2_5_VLCAConfig, **kwargs: Any) -> None:
27
+ del kwargs
28
+ super().__init__(config)
29
+ # Wrap the Qwen visual encoder so its output matches our CA interface
30
+ self.image_prefix = Qwen25VLEncoder(self.visual) # type: ignore[assignment]
31
+ self.visual = None
32
+ self.model.apply(
33
+ partial(maybe_replace_with_cross_attention_layers, xa_layers=self.config.xa_layers)
34
+ )
35
+ # The cross-attention layers are swapped in after the base post_init, so register
36
+ # the shared-weight alias keys and (re-)tie now that the CA modules exist.
37
+ if config.xa_share_qkvo:
38
+ shared_keys: list[str] = []
39
+ for i, layer in enumerate(self.model.layers):
40
+ if isinstance(layer.self_attn, Qwen2_5_VLAttention_CrossAttention):
41
+ for proj, biased in (
42
+ ("q_proj", True),
43
+ ("k_proj", True),
44
+ ("v_proj", True),
45
+ ("o_proj", False),
46
+ ):
47
+ prefix = f"model.layers.{i}.self_attn.cross_attn.{proj}"
48
+ shared_keys.append(f"{prefix}.weight")
49
+ if biased:
50
+ shared_keys.append(f"{prefix}.bias")
51
+ self._tied_weights_keys = list(self._tied_weights_keys or []) + shared_keys
52
+ self.tie_weights()
53
+
54
+ def _tie_weights(self) -> None:
55
+ if not getattr(self.config, "xa_share_qkvo", False):
56
+ return
57
+ for layer in self.model.layers:
58
+ if isinstance(layer.self_attn, Qwen2_5_VLAttention_CrossAttention):
59
+ tie_qkvo_projections(layer.self_attn, layer.self_attn.cross_attn)
60
+
61
+ def get_device(self) -> str:
62
+ """Return the device type of the model"""
63
+ return next(self.parameters()).device.type
64
+
65
+ @property
66
+ def token_dim(self) -> int:
67
+ """Returns the number of dimensions for the token representation"""
68
+ return self.config.hidden_size
69
+
70
+ def _update_model_kwargs_for_generation(
71
+ self,
72
+ outputs: Any,
73
+ model_kwargs: dict[str, Any],
74
+ is_encoder_decoder: bool = False,
75
+ num_new_tokens: int = 1,
76
+ ):
77
+ """Override to handle multi-turn generation and propagate updated attention masks"""
78
+ if (am := outputs.get("updated_attention_mask", None)) is not None:
79
+ model_kwargs["attention_mask"] = am
80
+ if "updated_cache_position" in outputs:
81
+ model_kwargs["cache_position"] = outputs.get("updated_cache_position")
82
+ else:
83
+ start = 0
84
+ if (kv := model_kwargs.get("past_key_values", None)) is not None:
85
+ start = kv._seen_tokens - am.shape[1]
86
+ model_kwargs["cache_position"] = torch.arange(
87
+ start,
88
+ start + am.shape[1],
89
+ dtype=model_kwargs["cache_position"].dtype,
90
+ device=model_kwargs["cache_position"].device,
91
+ )
92
+
93
+ # Call parent to get default updates
94
+ model_kwargs = super()._update_model_kwargs_for_generation(
95
+ outputs, model_kwargs, is_encoder_decoder, num_new_tokens
96
+ )
97
+ # Used by prepare_inputs_for_generation
98
+ model_kwargs["__is_first_gen_call__"] = False
99
+ return model_kwargs
100
+
101
+ def prepare_inputs_for_generation( # pyright: ignore[reportIncompatibleMethodOverride]
102
+ self,
103
+ input_ids: torch.Tensor,
104
+ past_key_values: DynamicCache | None = None,
105
+ **kwargs: Any,
106
+ ):
107
+ """Override to avoid Qwen erasing pixel_values on subsequent generation calls"""
108
+ backup = None
109
+ __is_first_gen_call__ = kwargs.get("__is_first_gen_call__", True)
110
+ if __is_first_gen_call__:
111
+ backup = kwargs.get("pixel_values", None)
112
+ if past_key_values is not None and (
113
+ kwargs.get("cache_position") is None
114
+ or type_cast(torch.Tensor, kwargs.get("cache_position")).shape[0] == 0
115
+ ):
116
+ # We're continuing from a cached state
117
+ past_length = past_key_values._seen_tokens
118
+ kwargs["cache_position"] = torch.arange(
119
+ past_length,
120
+ past_length + (input_ids.shape[1] if __is_first_gen_call__ else 1),
121
+ dtype=torch.long,
122
+ device=input_ids.device,
123
+ )
124
+ out = super().prepare_inputs_for_generation(
125
+ input_ids,
126
+ past_key_values=past_key_values,
127
+ **kwargs,
128
+ )
129
+ if backup is not None:
130
+ out["pixel_values"] = backup
131
+ return out
132
+
133
+ def prepare_multimodal_inputs(
134
+ self,
135
+ input_ids: torch.Tensor | None = None,
136
+ inputs_embeds: torch.Tensor | None = None,
137
+ attention_mask: torch.Tensor | None = None,
138
+ image_embeds_insertion_points: list[torch.Tensor] | None = None,
139
+ labels: torch.Tensor | None = None,
140
+ pixel_values: torch.Tensor | list[torch.Tensor] | None = None,
141
+ pre_image_tokens: list[int] | None = None,
142
+ post_image_tokens: list[int] | None = None,
143
+ **_kwargs: Any,
144
+ ) -> dict:
145
+ """Get a batch data mixing text and image data"""
146
+ del _kwargs
147
+
148
+ processed_inputs: dict = {
149
+ "input_ids": input_ids,
150
+ "inputs_embeds": inputs_embeds,
151
+ "labels": labels,
152
+ "attention_mask": attention_mask,
153
+ "image_embeds_insertion_points": image_embeds_insertion_points,
154
+ }
155
+ if pixel_values is not None:
156
+ processed_inputs.update(self.image_prefix(pixel_values))
157
+ image_embeds = processed_inputs.get("image_embeds")
158
+ assert image_embeds is not None
159
+ assert (isinstance(image_embeds, torch.Tensor) and image_embeds.ndim == 3) or (
160
+ isinstance(image_embeds, list) and all(_x.ndim == 2 for _x in image_embeds)
161
+ )
162
+
163
+ # Add kwargs necessary to compute cu_seqlens windows for CA
164
+ processed_inputs["ca_windows_info"] = {
165
+ "num_post_image_tokens": 0 if post_image_tokens is None else len(post_image_tokens),
166
+ "num_pre_image_tokens": 0 if pre_image_tokens is None else len(pre_image_tokens),
167
+ }
168
+
169
+ return processed_inputs
170
+
171
+ def forward( # type: ignore[override] # pylint: disable=W0221
172
+ self,
173
+ input_ids: torch.Tensor | None = None,
174
+ inputs_embeds: torch.Tensor | None = None,
175
+ attention_mask: torch.Tensor | None = None,
176
+ pixel_values: torch.Tensor | list[torch.Tensor] | None = None,
177
+ return_loss: bool = True,
178
+ labels: torch.Tensor | None = None,
179
+ image_embeds_insertion_points: list[torch.Tensor] | None = None,
180
+ pre_image_tokens: list[int] | None = None,
181
+ post_image_tokens: list[int] | None = None,
182
+ **kwargs: Any,
183
+ ) -> tuple | Qwen2_5_VLCausalLMOutputWithPast:
184
+ """Multi-modal forward pass"""
185
+ if self.training:
186
+ assert return_loss is True, (
187
+ "Qwen2.5VL always computes its own labels/losses in train mode"
188
+ )
189
+
190
+ if inputs_embeds is None:
191
+ assert input_ids is not None
192
+ inputs_embeds = type_cast(torch.Tensor, self.model.embed_tokens(input_ids))
193
+
194
+ # Case 1: First generation call — compute image embeddings and set up CA handler
195
+ if kwargs.pop("__is_first_gen_call__", True):
196
+ processed_inputs = self.prepare_multimodal_inputs(
197
+ input_ids=input_ids,
198
+ inputs_embeds=inputs_embeds,
199
+ attention_mask=attention_mask,
200
+ image_embeds_insertion_points=image_embeds_insertion_points,
201
+ pixel_values=pixel_values,
202
+ labels=labels,
203
+ pre_image_tokens=pre_image_tokens,
204
+ post_image_tokens=post_image_tokens,
205
+ )
206
+ image_embeds = processed_inputs.get("image_embeds", None)
207
+ inst_points = processed_inputs.get("image_embeds_insertion_points", None)
208
+
209
+ # Only build a handler when images are actually present
210
+ cross_attention_handler: CrossAttentionHandler | None = None
211
+ if image_embeds is not None and len(image_embeds) > 0:
212
+ cross_attention_handler = CrossAttentionHandler(
213
+ inputs_embeds=torch.zeros_like(inputs_embeds),
214
+ image_embeds=image_embeds,
215
+ image_embeds_insertion_points=inst_points,
216
+ ca_windows_info=processed_inputs.pop("ca_windows_info", None),
217
+ training=self.training,
218
+ )
219
+ self.update_cross_attention_states(cross_attention_handler)
220
+
221
+ # Run Qwen with the attention layers replaced to use cross-attention
222
+ assert inputs_embeds is not None, "Could not compute input embeddings!"
223
+ out = super().forward(
224
+ inputs_embeds=inputs_embeds, # type: ignore[arg-type]
225
+ attention_mask=attention_mask,
226
+ pixel_values=None,
227
+ **kwargs,
228
+ )
229
+
230
+ return out
231
+
232
+ @property
233
+ def default_generation_eos_token_id(self) -> int | Sequence[int] | None:
234
+ return self.generation_config.eos_token_id if self.generation_config is not None else None
235
+
236
+ @torch.no_grad()
237
+ def generate_from_image( # pyright: ignore[reportInconsistentOverload]
238
+ self,
239
+ reset_streaming: bool = True,
240
+ temperature: float | None = 0.0,
241
+ eos_token_id: int | Sequence[int] | None = None,
242
+ **kwargs: Any,
243
+ ) -> GenerateOutput | torch.LongTensor:
244
+ """Custom generate function"""
245
+ # init self-attention KVCache
246
+ if kwargs.get("past_key_values", None) is None:
247
+ kwargs["past_key_values"] = DynamicCache()
248
+
249
+ if eos_token_id is None:
250
+ eos_token_id = self.default_generation_eos_token_id
251
+ # To avoid generate warning
252
+ if kwargs.get("pad_token_id", None) is None:
253
+ kwargs["pad_token_id"] = kwargs.get("eos_token_id", None)
254
+ if isinstance(kwargs["pad_token_id"], (list, tuple)):
255
+ kwargs["pad_token_id"] = kwargs["pad_token_id"][0]
256
+ if "pre_image_tokens" not in kwargs:
257
+ kwargs["pre_image_tokens"] = list(self.config.pre_image_tokens)
258
+ if "post_image_tokens" not in kwargs:
259
+ kwargs["post_image_tokens"] = list(self.config.post_image_tokens)
260
+
261
+ if not kwargs.get("do_sample", False):
262
+ temperature = None
263
+ kwargs.pop("top_p", None)
264
+ kwargs.pop("top_k", None)
265
+
266
+ # Generate
267
+ self.start_ca_streaming_states()
268
+ outputs = self.generate(
269
+ use_cache=True,
270
+ eos_token_id=eos_token_id,
271
+ temperature=temperature,
272
+ **kwargs,
273
+ )
274
+ if reset_streaming:
275
+ self.reset_ca_streaming_states()
276
+ return outputs
277
+
278
+ def update_cross_attention_states(self, handler: CrossAttentionHandler | None):
279
+ """Push the new handler into all CA attention layers"""
280
+
281
+ def __update__(m: torch.nn.Module):
282
+ nonlocal handler
283
+ if isinstance(m, Qwen2_5_VLAttention_CrossAttention):
284
+ m.cross_attention_handler = handler
285
+
286
+ self.apply(__update__)
287
+
288
+ def reset_ca_streaming_states(self) -> None:
289
+ def __reset__(m: torch.nn.Module):
290
+ if isinstance(m, QwenCrossAttention):
291
+ m._set_streaming(False, ())
292
+ m.reset_streaming()
293
+ if hasattr(m, "cross_attention_handler"):
294
+ del m.cross_attention_handler
295
+ m.cross_attention_handler = None
296
+
297
+ self.apply(__reset__)
298
+
299
+ def start_ca_streaming_states(self) -> None:
300
+ def __start__(m: torch.nn.Module):
301
+ if isinstance(m, QwenCrossAttention):
302
+ m._set_streaming(True, ())
303
+
304
+ self.apply(__start__)
processing.py ADDED
@@ -0,0 +1,494 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # pylint: disable=no-member # avoid weird pylint warnings from SentencePieceProcessor
2
+ """Text and Image processor for CA models using Qwen2.5_VL image encoder"""
3
+
4
+ from math import ceil
5
+ from typing import TYPE_CHECKING, Any, Literal, TypedDict, cast, overload
6
+ from typing import cast as type_cast
7
+
8
+ import torch
9
+ import torchvision.transforms.v2 as T
10
+ from einops import rearrange
11
+ from PIL import Image
12
+ from torchvision.transforms import InterpolationMode
13
+ from torchvision.transforms.functional import to_tensor as pil_to_tensor
14
+ from torchvision.transforms.v2 import functional as F
15
+ from transformers.image_processing_utils import BaseImageProcessor
16
+ from transformers.processing_utils import ProcessorMixin
17
+
18
+ if TYPE_CHECKING:
19
+ from transformers.models.qwen2.tokenization_qwen2 import Qwen2Tokenizer
20
+ from transformers.tokenization_utils_fast import PreTrainedTokenizerFast
21
+
22
+
23
+ ImageMessage = TypedDict(
24
+ "ImageMessage",
25
+ {
26
+ "type": Literal["image"],
27
+ "image": str | Image.Image | None,
28
+ },
29
+ )
30
+
31
+ TextMessage = TypedDict(
32
+ "TextMessage",
33
+ {
34
+ "type": Literal["text"],
35
+ "text": str,
36
+ },
37
+ )
38
+
39
+ MessageContent = list[ImageMessage | TextMessage]
40
+
41
+ Message = TypedDict(
42
+ "Message",
43
+ {
44
+ "role": Literal["system", "user", "assistant"],
45
+ "content": MessageContent,
46
+ },
47
+ )
48
+
49
+ ProcessorInput = list[list[Message]] | list[Message]
50
+
51
+ __INTERP_NAME_TO_MODE__ = {
52
+ "nearest": InterpolationMode.NEAREST,
53
+ "bilinear": InterpolationMode.BILINEAR,
54
+ "bicubic": InterpolationMode.BICUBIC,
55
+ "lanczos": InterpolationMode.LANCZOS,
56
+ }
57
+
58
+ __INTERP_INT_TO_MODE__ = {
59
+ 0: InterpolationMode.NEAREST,
60
+ 2: InterpolationMode.BILINEAR,
61
+ 3: InterpolationMode.BICUBIC,
62
+ 4: InterpolationMode.BOX,
63
+ 5: InterpolationMode.HAMMING,
64
+ 1: InterpolationMode.LANCZOS,
65
+ }
66
+
67
+
68
+ @overload
69
+ def universal_resize(
70
+ img: Image.Image,
71
+ size: tuple[int, int],
72
+ interpolation: str | InterpolationMode | int = "bilinear",
73
+ antialias: bool = True,
74
+ ) -> Image.Image: ...
75
+ @overload
76
+ def universal_resize(
77
+ img: torch.Tensor,
78
+ size: tuple[int, int],
79
+ interpolation: str | InterpolationMode | int = "bilinear",
80
+ antialias: bool = True,
81
+ ) -> torch.Tensor: ...
82
+ def universal_resize(
83
+ img: Image.Image | torch.Tensor,
84
+ size: tuple[int, int],
85
+ interpolation: str | InterpolationMode | int = "bilinear",
86
+ antialias: bool = True,
87
+ ) -> Image.Image | torch.Tensor:
88
+ """Resize that works for PIL.Image, CHW tensor, or BCHW tensor"""
89
+ if isinstance(interpolation, str):
90
+ interpolation = __INTERP_NAME_TO_MODE__[interpolation]
91
+ elif isinstance(interpolation, int):
92
+ interpolation = __INTERP_INT_TO_MODE__[interpolation]
93
+
94
+ return F.resize(
95
+ img, size, interpolation=type_cast(InterpolationMode, interpolation), antialias=antialias
96
+ )
97
+
98
+
99
+ @overload
100
+ def convert_to_rgb(img: Image.Image) -> Image.Image: ...
101
+ @overload
102
+ def convert_to_rgb(img: torch.Tensor) -> torch.Tensor: ...
103
+ def convert_to_rgb(img: Image.Image | torch.Tensor) -> Image.Image | torch.Tensor:
104
+ """Convert any image to RGB in a way that does not throw PIL warning"""
105
+ if isinstance(img, torch.Tensor):
106
+ return img
107
+ if img.mode == "RGB": # no changes
108
+ return img
109
+ if img.mode == "P": # palette images need to be converted to RGBA first
110
+ return img.convert("RGBA").convert("RGB")
111
+ return img.convert("RGB")
112
+
113
+
114
+ class QwenImageProcessor(BaseImageProcessor):
115
+ """Resizing for the Qwen2.5VL encoder. Note that the normalization is
116
+ handled in the image_encoder in the model forward"""
117
+
118
+ def __init__(
119
+ self,
120
+ img_size: int = 448,
121
+ interpolation: Literal["bicubic", "bilinear", "nearest", "nearest_exact"] = "bicubic",
122
+ max_ratio: int = 10,
123
+ round_to_patch_size: int = 56,
124
+ use_fast: bool = True,
125
+ **kwargs: Any,
126
+ ) -> None:
127
+ self._num_target_channels = 588
128
+ self._merge_size = 2
129
+ self._patch_size = 14
130
+ super().__init__(
131
+ use_fast=use_fast,
132
+ do_normalize=False,
133
+ **kwargs,
134
+ )
135
+ self.img_size = img_size
136
+ self.interpolation = interpolation
137
+ self.max_ratio = max_ratio
138
+ self.round_to_patch_size = round_to_patch_size
139
+
140
+ def resize_transform(
141
+ self, img: Image.Image | torch.Tensor, img_size: int | None = None
142
+ ) -> Image.Image | torch.Tensor:
143
+ if img_size is None:
144
+ img_size = self.img_size
145
+ max_area = img_size**2
146
+ if isinstance(img, Image.Image):
147
+ img = convert_to_rgb(img)
148
+ w_og, h_og = img.size
149
+ else:
150
+ h_og, w_og = img.shape[-2:]
151
+ w, h = w_og, h_og
152
+
153
+ # Qwen requires max ratio of 10 between max and min sizes
154
+ if self.max_ratio > 0:
155
+ w, h = max(w, h // self.max_ratio), max(h, w // self.max_ratio)
156
+
157
+ # resize to max area
158
+ current_area = w * h
159
+ if current_area > max_area:
160
+ scale = (max_area / current_area) ** 0.5
161
+ w, h = int(w * scale), int(h * scale)
162
+
163
+ # resize to patch size
164
+ if self.round_to_patch_size > 0:
165
+ w = ceil(w / self.round_to_patch_size) * self.round_to_patch_size
166
+ h = ceil((h / self.round_to_patch_size)) * self.round_to_patch_size
167
+
168
+ # resize
169
+ if w != w_og or h != h_og:
170
+ img = universal_resize(img, (h, w), self.interpolation)
171
+ if isinstance(img, torch.Tensor):
172
+ img = T.ToDtype(torch.float32, scale=True)(T.ToImage()(img))
173
+ return img
174
+
175
+ def __process_one__(
176
+ self, video_or_img: Image.Image | torch.Tensor, img_size: int | None = None
177
+ ) -> torch.Tensor:
178
+ """Same operation as __process_one_with_processor__ but without going through numpy"""
179
+ video_or_img = self.resize_transform(video_or_img, img_size)
180
+ if isinstance(video_or_img, Image.Image):
181
+ video_or_img = pil_to_tensor(video_or_img)
182
+ assert isinstance(video_or_img, torch.Tensor)
183
+ if video_or_img.ndim == 3:
184
+ video_or_img = video_or_img[None]
185
+ assert video_or_img.ndim == 4 and video_or_img.shape[1] == 3, (
186
+ f"Invalid shape {video_or_img.shape}."
187
+ )
188
+ t, c, h, w = video_or_img.shape
189
+ p = self._patch_size
190
+ m = self._merge_size
191
+
192
+ # Convert to RGB
193
+ if c == 1:
194
+ video_or_img = video_or_img.expand((-1, 3, -1, -1))
195
+ if c == 4:
196
+ video_or_img = video_or_img[:, :3]
197
+ c = video_or_img.shape[1]
198
+ assert c == 3, "Expecting RGB image in QwenNormalize"
199
+
200
+ # Reshape to t h w c' format
201
+ h, w = video_or_img.shape[2] // p, video_or_img.shape[3] // p
202
+ rearrange_dict = dict(p1=p, p2=p, m1=m, m2=m)
203
+
204
+ video_or_img = rearrange(
205
+ video_or_img,
206
+ "t c (h m1 p1) (w m2 p2) -> (t h w m1 m2) (c p1 p2)",
207
+ **rearrange_dict,
208
+ )
209
+ assert video_or_img.shape[-1] == self._num_target_channels, (
210
+ f"{video_or_img.shape[-1]} != {self._num_target_channels}"
211
+ )
212
+ video_or_img = video_or_img.view((-1, h, w, self._num_target_channels))
213
+
214
+ return video_or_img
215
+
216
+ @overload
217
+ def process_images(
218
+ self, image: Image.Image | torch.Tensor, img_size: int | None = None
219
+ ) -> torch.Tensor: ...
220
+ @overload
221
+ def process_images(
222
+ self, image: list[Image.Image] | list[torch.Tensor], img_size: int | None = None
223
+ ) -> list[torch.Tensor]: ...
224
+ def process_images(
225
+ self,
226
+ image: Image.Image | torch.Tensor | list[Image.Image] | list[torch.Tensor],
227
+ img_size: int | None = None,
228
+ ) -> torch.Tensor | list[torch.Tensor]:
229
+ if isinstance(image, list):
230
+ return [self.__process_one__(_x, img_size) for _x in image]
231
+ return self.__process_one__(image, img_size)
232
+
233
+
234
+ class ProcessorOutput(dict):
235
+ input_ids: torch.Tensor
236
+ attention_mask: torch.Tensor
237
+ image_embeds_insertion_points: list[torch.Tensor] | None
238
+ pixel_values: torch.Tensor | list[torch.Tensor] | None
239
+
240
+ def to(
241
+ self, device: torch.device | str, dtype: torch.dtype = torch.bfloat16
242
+ ) -> "ProcessorOutput":
243
+ return ProcessorOutput(
244
+ {
245
+ "input_ids": self["input_ids"].to(device),
246
+ "attention_mask": self["attention_mask"].to(device),
247
+ "image_embeds_insertion_points": self["image_embeds_insertion_points"],
248
+ "pixel_values": (
249
+ self["pixel_values"].to(dtype).to(device)
250
+ if isinstance(self["pixel_values"], torch.Tensor)
251
+ else [x.to(dtype).to(device) for x in self["pixel_values"]]
252
+ if self["pixel_values"] is not None
253
+ else None
254
+ ),
255
+ }
256
+ )
257
+
258
+
259
+ class BaseProcessor(ProcessorMixin):
260
+ def __init__(
261
+ self,
262
+ tokenizer: "PreTrainedTokenizerFast | Qwen2Tokenizer",
263
+ pre_image_tokens: tuple[int, ...] = (),
264
+ post_image_tokens: tuple[int, ...] = (),
265
+ system_start_tokens: tuple[int, ...] = (),
266
+ system_end_tokens: tuple[int, ...] = (),
267
+ user_start_tokens: tuple[int, ...] = (),
268
+ user_end_tokens: tuple[int, ...] = (),
269
+ asst_start_tokens: tuple[int, ...] = (),
270
+ asst_end_tokens: tuple[int, ...] = (),
271
+ allow_system_prompt: bool = True,
272
+ pad_token: int = 0,
273
+ bos_token: int | None = None,
274
+ ) -> None:
275
+ self.pre_image_tokens = list(pre_image_tokens)
276
+ self.post_image_tokens = list(post_image_tokens)
277
+ self.system_start_tokens = list(system_start_tokens)
278
+ self.system_end_tokens = list(system_end_tokens)
279
+ self.user_start_tokens = list(user_start_tokens)
280
+ self.user_end_tokens = list(user_end_tokens)
281
+ self.asst_start_tokens = list(asst_start_tokens)
282
+ self.asst_end_tokens = list(asst_end_tokens)
283
+ self._allow_system_prompt = allow_system_prompt
284
+ self.tokenizer = tokenizer
285
+ self._image_processor = None
286
+ self._pad_token = pad_token
287
+ self.bos_token = bos_token
288
+
289
+ @property
290
+ def image_processor(self) -> QwenImageProcessor:
291
+ assert self._image_processor is not None
292
+ return self._image_processor
293
+
294
+ def _process_content(
295
+ self,
296
+ message_content: MessageContent,
297
+ role: Literal["system", "user", "assistant"],
298
+ tokenized_messages: list[torch.Tensor],
299
+ insertion_points: list[int],
300
+ image_list: list[torch.Tensor | None],
301
+ token_count: int,
302
+ img_size: int | None = None,
303
+ **kwargs: Any,
304
+ ) -> int:
305
+ mapping = {
306
+ "user": (self.user_start_tokens, self.user_end_tokens),
307
+ "assistant": (self.asst_start_tokens, self.asst_end_tokens),
308
+ "system": (self.system_start_tokens, self.system_end_tokens),
309
+ }
310
+ if role.lower() not in mapping:
311
+ raise ValueError(f"Unknown role '{role}' encountered in messages.")
312
+ start_tokens, end_tokens = mapping[role.lower()]
313
+ # 1) Add the start tokens
314
+ if start_tokens:
315
+ tokenized_messages.append(torch.Tensor(start_tokens).flatten().to(torch.long))
316
+ token_count += len(start_tokens)
317
+ # 2) Process the message content one by one (potentially interleaved image and text)
318
+ for part in message_content:
319
+ elt_type = part["type"]
320
+ if elt_type == "image":
321
+ part = cast(ImageMessage, part)
322
+ self._process_image_message(
323
+ part,
324
+ tokenized_messages,
325
+ image_list,
326
+ img_size=img_size,
327
+ )
328
+ token_count += len(self.pre_image_tokens)
329
+ insertion_points.append(token_count)
330
+ token_count += len(self.post_image_tokens)
331
+ else:
332
+ part = cast(TextMessage, part)
333
+ self._process_text_message(
334
+ part["text"],
335
+ role=role,
336
+ token_list=tokenized_messages,
337
+ **kwargs,
338
+ )
339
+ token_count += tokenized_messages[-1].size(0)
340
+ # 3) Add the end tokens
341
+ if end_tokens:
342
+ tokenized_messages.append(torch.Tensor(end_tokens).flatten().to(torch.long))
343
+ token_count += len(end_tokens)
344
+ return token_count
345
+
346
+ def _process_text_message(
347
+ self,
348
+ message: str,
349
+ role: Literal["system", "user", "assistant"],
350
+ token_list: list[torch.Tensor],
351
+ **kwargs: Any,
352
+ ) -> None:
353
+ if role.lower() == "system" and not self._allow_system_prompt:
354
+ raise ValueError("System prompts are not allowed in this tokenizer configuration.")
355
+ tokens = self.tokenizer.encode(
356
+ message, add_special_tokens=False, return_tensors="pt", **kwargs
357
+ )
358
+ tokens = cast(torch.Tensor, tokens)
359
+ token_list.append(tokens.flatten().to(torch.long))
360
+
361
+ def _process_image_message(
362
+ self,
363
+ message: ImageMessage,
364
+ token_list: list[torch.Tensor],
365
+ image_list: list[torch.Tensor | None],
366
+ img_size: int | None = None,
367
+ ) -> None:
368
+ img = message["image"]
369
+ if img is None:
370
+ image_list.append(None)
371
+ else:
372
+ image_list.append(
373
+ self.image_processor.process_images(
374
+ self._load_image(img), img_size=img_size
375
+ ).squeeze(0)
376
+ )
377
+ if self.pre_image_tokens:
378
+ token_list.append(torch.Tensor(self.pre_image_tokens).flatten().to(torch.long))
379
+
380
+ if self.post_image_tokens:
381
+ token_list.append(torch.Tensor(self.post_image_tokens).flatten().to(torch.long))
382
+
383
+ def _load_image(self, image_path_or_image: str | Image.Image) -> Image.Image:
384
+ if isinstance(image_path_or_image, str):
385
+ return Image.open(image_path_or_image).convert("RGB")
386
+ return image_path_or_image
387
+
388
+ def _maybe_pad(self, tokens: torch.Tensor, pad_len: int, pad_value: int) -> torch.Tensor:
389
+ return torch.nn.functional.pad(
390
+ tokens,
391
+ (0, pad_len) if self.tokenizer.padding_side == "right" else (pad_len, 0),
392
+ value=pad_value,
393
+ )
394
+
395
+ def pad_tokenized_messages(
396
+ self,
397
+ tokenized_messages_batch: list[torch.Tensor],
398
+ image_insertion_points_batch: list[torch.Tensor] | None = None,
399
+ ) -> tuple[torch.Tensor, torch.Tensor, list[torch.Tensor] | None]:
400
+ max_len = max(len(x) for x in tokenized_messages_batch)
401
+ if image_insertion_points_batch is not None and self.tokenizer.padding_side == "left":
402
+ image_insertion_points_batch = [
403
+ x + max_len - len(tokenized_messages_batch[idx])
404
+ for idx, x in enumerate(image_insertion_points_batch)
405
+ ]
406
+ input_ids = torch.stack(
407
+ [
408
+ self._maybe_pad(s, max_len - s.size(0), self._pad_token)
409
+ for s in tokenized_messages_batch
410
+ ],
411
+ dim=0,
412
+ )
413
+ attention_mask = torch.stack(
414
+ [
415
+ self._maybe_pad(torch.ones_like(s), max_len - s.size(0), 0)
416
+ for s in tokenized_messages_batch
417
+ ],
418
+ dim=0,
419
+ )
420
+ return input_ids, attention_mask, image_insertion_points_batch
421
+
422
+ def tokenize_messages(
423
+ self,
424
+ messages: ProcessorInput,
425
+ suppress_bos_token: bool = False,
426
+ **kwargs: Any,
427
+ ) -> ProcessorOutput | None:
428
+ """Tokenize a batch of messages into token IDs suitable for the CA model.
429
+
430
+ Args:
431
+ messages: Batch of message lists (or single list of messages),
432
+ where each message is a dict with 'role' and 'content' keys.
433
+ suppress_bos_token: If True, the BOS token will not be added.
434
+ **kwargs: Additional keyword arguments passed to the underlying encode method.
435
+ """
436
+ if not messages:
437
+ return None
438
+ if isinstance(messages[0], dict):
439
+ messages = [messages] # type: ignore[assignment]
440
+
441
+ messages = cast(list[list[Message]], messages)
442
+ image_insertion_points_batch = []
443
+ tokenized_messages_batch = []
444
+ image_list: list[torch.Tensor | None] = []
445
+ for msgs in messages:
446
+ tokenized_messages = []
447
+ if not suppress_bos_token and self.bos_token is not None:
448
+ tokenized_messages.append(torch.tensor([self.bos_token], dtype=torch.long))
449
+ insertion_points = []
450
+ token_count = 0
451
+ for msg in msgs:
452
+ token_count = self._process_content(
453
+ msg["content"],
454
+ role=msg["role"],
455
+ tokenized_messages=tokenized_messages,
456
+ insertion_points=insertion_points,
457
+ image_list=image_list,
458
+ token_count=token_count,
459
+ **kwargs,
460
+ )
461
+ tokenized_messages_batch.append(torch.cat(tokenized_messages, dim=0).to(torch.long))
462
+ image_insertion_points_batch.append(torch.tensor(insertion_points, dtype=torch.long))
463
+
464
+ if msgs and self.asst_end_tokens and msgs[-1]["role"].lower() == "assistant":
465
+ # Remove the assistant end tokens from the final message
466
+ end_token_len = len(self.asst_end_tokens)
467
+ tokenized_messages_batch[-1] = tokenized_messages_batch[-1][:-end_token_len]
468
+ if msgs and self.asst_start_tokens and msgs[-1]["role"].lower() == "user":
469
+ tokenized_messages_batch[-1] = torch.cat(
470
+ [
471
+ tokenized_messages_batch[-1],
472
+ torch.Tensor(self.asst_start_tokens).to(torch.long),
473
+ ]
474
+ )
475
+
476
+ input_ids, attention_mask, image_embeds_insertion_points = self.pad_tokenized_messages(
477
+ tokenized_messages_batch, image_insertion_points_batch
478
+ )
479
+
480
+ if image_list:
481
+ assert sum(img is None for img in image_list) % len(image_list) == 0, (
482
+ "Either all or no image must be None."
483
+ )
484
+ pixel_values: None | torch.Tensor | list[torch.Tensor]
485
+ if image_list[0] is None:
486
+ pixel_values = None
487
+ else:
488
+ pixel_values = cast(list[torch.Tensor], image_list)
489
+ return ProcessorOutput(
490
+ input_ids=input_ids,
491
+ image_embeds_insertion_points=image_embeds_insertion_points,
492
+ attention_mask=attention_mask,
493
+ pixel_values=pixel_values,
494
+ )
processing_qwen2_5vl_ca.py ADDED
@@ -0,0 +1,39 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from typing import Any
2
+
3
+ from transformers.models.qwen2.tokenization_qwen2 import Qwen2Tokenizer
4
+
5
+ from .processing import BaseProcessor, QwenImageProcessor
6
+
7
+
8
+ class QwenCASAProcessor(BaseProcessor):
9
+ attributes = ["tokenizer"]
10
+ tokenizer_class = "Qwen2Tokenizer"
11
+
12
+ def __init__(
13
+ self,
14
+ tokenizer: Qwen2Tokenizer,
15
+ pre_image_tokens: tuple[int, ...] = (151652,),
16
+ post_image_tokens: tuple[int, ...] = (151653,),
17
+ system_start_tokens: tuple[int, ...] = (151644, 8948, 198),
18
+ system_end_tokens: tuple[int, ...] = (151645, 198),
19
+ user_start_tokens: tuple[int, ...] = (151644, 872, 198),
20
+ user_end_tokens: tuple[int, ...] = (151645, 198),
21
+ asst_start_tokens: tuple[int, ...] = (151644, 77091, 198),
22
+ asst_end_tokens: tuple[int, ...] = (151645, 198),
23
+ image_size: int = 448,
24
+ **kwargs: Any,
25
+ ):
26
+ del kwargs
27
+ super().__init__(
28
+ tokenizer=tokenizer,
29
+ pre_image_tokens=pre_image_tokens,
30
+ post_image_tokens=post_image_tokens,
31
+ system_start_tokens=system_start_tokens,
32
+ system_end_tokens=system_end_tokens,
33
+ user_start_tokens=user_start_tokens,
34
+ user_end_tokens=user_end_tokens,
35
+ asst_start_tokens=asst_start_tokens,
36
+ asst_end_tokens=asst_end_tokens,
37
+ )
38
+
39
+ self._image_processor = QwenImageProcessor(img_size=image_size)
processor_config.json ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "auto_map": {
3
+ "AutoProcessor": "processing_qwen2_5vl_ca.QwenCASAProcessor"
4
+ },
5
+ "image_size": 896,
6
+ "post_image_tokens": [
7
+ 151653
8
+ ],
9
+ "pre_image_tokens": [
10
+ 151652
11
+ ],
12
+ "processor_class": "QwenCASAProcessor"
13
+ }
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer_config.json ADDED
@@ -0,0 +1,208 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "add_prefix_space": false,
3
+ "added_tokens_decoder": {
4
+ "151643": {
5
+ "content": "<|endoftext|>",
6
+ "lstrip": false,
7
+ "normalized": false,
8
+ "rstrip": false,
9
+ "single_word": false,
10
+ "special": true
11
+ },
12
+ "151644": {
13
+ "content": "<|im_start|>",
14
+ "lstrip": false,
15
+ "normalized": false,
16
+ "rstrip": false,
17
+ "single_word": false,
18
+ "special": true
19
+ },
20
+ "151645": {
21
+ "content": "<|im_end|>",
22
+ "lstrip": false,
23
+ "normalized": false,
24
+ "rstrip": false,
25
+ "single_word": false,
26
+ "special": true
27
+ },
28
+ "151646": {
29
+ "content": "<|object_ref_start|>",
30
+ "lstrip": false,
31
+ "normalized": false,
32
+ "rstrip": false,
33
+ "single_word": false,
34
+ "special": true
35
+ },
36
+ "151647": {
37
+ "content": "<|object_ref_end|>",
38
+ "lstrip": false,
39
+ "normalized": false,
40
+ "rstrip": false,
41
+ "single_word": false,
42
+ "special": true
43
+ },
44
+ "151648": {
45
+ "content": "<|box_start|>",
46
+ "lstrip": false,
47
+ "normalized": false,
48
+ "rstrip": false,
49
+ "single_word": false,
50
+ "special": true
51
+ },
52
+ "151649": {
53
+ "content": "<|box_end|>",
54
+ "lstrip": false,
55
+ "normalized": false,
56
+ "rstrip": false,
57
+ "single_word": false,
58
+ "special": true
59
+ },
60
+ "151650": {
61
+ "content": "<|quad_start|>",
62
+ "lstrip": false,
63
+ "normalized": false,
64
+ "rstrip": false,
65
+ "single_word": false,
66
+ "special": true
67
+ },
68
+ "151651": {
69
+ "content": "<|quad_end|>",
70
+ "lstrip": false,
71
+ "normalized": false,
72
+ "rstrip": false,
73
+ "single_word": false,
74
+ "special": true
75
+ },
76
+ "151652": {
77
+ "content": "<|vision_start|>",
78
+ "lstrip": false,
79
+ "normalized": false,
80
+ "rstrip": false,
81
+ "single_word": false,
82
+ "special": true
83
+ },
84
+ "151653": {
85
+ "content": "<|vision_end|>",
86
+ "lstrip": false,
87
+ "normalized": false,
88
+ "rstrip": false,
89
+ "single_word": false,
90
+ "special": true
91
+ },
92
+ "151654": {
93
+ "content": "<|vision_pad|>",
94
+ "lstrip": false,
95
+ "normalized": false,
96
+ "rstrip": false,
97
+ "single_word": false,
98
+ "special": true
99
+ },
100
+ "151655": {
101
+ "content": "<|image_pad|>",
102
+ "lstrip": false,
103
+ "normalized": false,
104
+ "rstrip": false,
105
+ "single_word": false,
106
+ "special": true
107
+ },
108
+ "151656": {
109
+ "content": "<|video_pad|>",
110
+ "lstrip": false,
111
+ "normalized": false,
112
+ "rstrip": false,
113
+ "single_word": false,
114
+ "special": true
115
+ },
116
+ "151657": {
117
+ "content": "<tool_call>",
118
+ "lstrip": false,
119
+ "normalized": false,
120
+ "rstrip": false,
121
+ "single_word": false,
122
+ "special": false
123
+ },
124
+ "151658": {
125
+ "content": "</tool_call>",
126
+ "lstrip": false,
127
+ "normalized": false,
128
+ "rstrip": false,
129
+ "single_word": false,
130
+ "special": false
131
+ },
132
+ "151659": {
133
+ "content": "<|fim_prefix|>",
134
+ "lstrip": false,
135
+ "normalized": false,
136
+ "rstrip": false,
137
+ "single_word": false,
138
+ "special": false
139
+ },
140
+ "151660": {
141
+ "content": "<|fim_middle|>",
142
+ "lstrip": false,
143
+ "normalized": false,
144
+ "rstrip": false,
145
+ "single_word": false,
146
+ "special": false
147
+ },
148
+ "151661": {
149
+ "content": "<|fim_suffix|>",
150
+ "lstrip": false,
151
+ "normalized": false,
152
+ "rstrip": false,
153
+ "single_word": false,
154
+ "special": false
155
+ },
156
+ "151662": {
157
+ "content": "<|fim_pad|>",
158
+ "lstrip": false,
159
+ "normalized": false,
160
+ "rstrip": false,
161
+ "single_word": false,
162
+ "special": false
163
+ },
164
+ "151663": {
165
+ "content": "<|repo_name|>",
166
+ "lstrip": false,
167
+ "normalized": false,
168
+ "rstrip": false,
169
+ "single_word": false,
170
+ "special": false
171
+ },
172
+ "151664": {
173
+ "content": "<|file_sep|>",
174
+ "lstrip": false,
175
+ "normalized": false,
176
+ "rstrip": false,
177
+ "single_word": false,
178
+ "special": false
179
+ }
180
+ },
181
+ "additional_special_tokens": [
182
+ "<|endoftext|>",
183
+ "<|im_start|>",
184
+ "<|im_end|>",
185
+ "<|object_ref_start|>",
186
+ "<|object_ref_end|>",
187
+ "<|box_start|>",
188
+ "<|box_end|>",
189
+ "<|quad_start|>",
190
+ "<|quad_end|>",
191
+ "<|vision_start|>",
192
+ "<|vision_end|>",
193
+ "<|vision_pad|>",
194
+ "<|image_pad|>",
195
+ "<|video_pad|>"
196
+ ],
197
+ "bos_token": null,
198
+ "chat_template": "{% set image_count = namespace(value=0) %}{% set video_count = namespace(value=0) %}{% for message in messages %}{% if loop.first and message['role'] != 'system' %}<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n{% endif %}<|im_start|>{{ message['role'] }}\n{% if message['content'] is string %}{{ message['content'] }}<|im_end|>\n{% else %}{% for content in message['content'] %}{% if content['type'] == 'image' or 'image' in content or 'image_url' in content %}{% set image_count.value = image_count.value + 1 %}{% if add_vision_id %}Picture {{ image_count.value }}: {% endif %}<|vision_start|><|image_pad|><|vision_end|>{% elif content['type'] == 'video' or 'video' in content %}{% set video_count.value = video_count.value + 1 %}{% if add_vision_id %}Video {{ video_count.value }}: {% endif %}<|vision_start|><|video_pad|><|vision_end|>{% elif 'text' in content %}{{ content['text'] }}{% endif %}{% endfor %}<|im_end|>\n{% endif %}{% endfor %}{% if add_generation_prompt %}<|im_start|>assistant\n{% endif %}",
199
+ "clean_up_tokenization_spaces": false,
200
+ "eos_token": "<|im_end|>",
201
+ "errors": "replace",
202
+ "model_max_length": 131072,
203
+ "pad_token": "<|endoftext|>",
204
+ "split_special_tokens": false,
205
+ "tokenizer_class": "Qwen2Tokenizer",
206
+ "unk_token": null,
207
+ "add_bos_token": false
208
+ }
utils.py ADDED
@@ -0,0 +1,337 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # pylint: disable=protected-access
2
+ """Utils to handle CA layers construction"""
3
+
4
+ from contextlib import contextmanager
5
+ from dataclasses import dataclass, fields
6
+ from typing import Any, Callable, Generic, Literal, Sequence, TypeVar, overload
7
+ from typing import cast as type_cast
8
+
9
+ import torch
10
+
11
+
12
+ def __split_n_merge__(
13
+ x: torch.Tensor,
14
+ sample_lengths: list[int],
15
+ padding_side: Literal["left", "right"] = "right",
16
+ pad_value: int | float | bool = 0,
17
+ ) -> torch.Tensor:
18
+ max_sample_length = max(sample_lengths)
19
+ pad_tuple = tuple(0 for _ in range((x.ndim - 1) * 2))
20
+ return torch.stack(
21
+ [
22
+ torch.nn.functional.pad(
23
+ _x,
24
+ pad_tuple + (0, max_sample_length - _x.shape[0])
25
+ if padding_side == "right"
26
+ else pad_tuple + (max_sample_length - _x.shape[0], 0),
27
+ value=pad_value,
28
+ )
29
+ for _x in torch.split(x, sample_lengths, dim=0)
30
+ ],
31
+ dim=0,
32
+ )
33
+
34
+
35
+ @overload
36
+ def insert_image_tokens(
37
+ inputs_embeds: torch.Tensor,
38
+ image_embeds: torch.Tensor | Sequence[torch.Tensor],
39
+ image_embeds_insertion_points: list[torch.Tensor],
40
+ recover_batch_dim: Literal[True],
41
+ attention_mask: torch.Tensor | None = None,
42
+ padding_side: Literal["left", "right"] = "right",
43
+ keep_only_attended: bool = False,
44
+ pad_output: int | float | bool = 0.0,
45
+ ) -> tuple[
46
+ torch.Tensor,
47
+ None,
48
+ torch.Tensor | None,
49
+ torch.Tensor,
50
+ ]: ...
51
+ @overload
52
+ def insert_image_tokens(
53
+ inputs_embeds: torch.Tensor,
54
+ image_embeds: torch.Tensor | Sequence[torch.Tensor],
55
+ image_embeds_insertion_points: list[torch.Tensor],
56
+ recover_batch_dim: Literal[False],
57
+ attention_mask: torch.Tensor | None = None,
58
+ padding_side: Literal["left", "right"] = "right",
59
+ keep_only_attended: bool = False,
60
+ pad_output: int | float | bool = 0.0,
61
+ ) -> tuple[
62
+ torch.Tensor,
63
+ list[int],
64
+ torch.Tensor | None,
65
+ torch.Tensor,
66
+ ]: ...
67
+ def insert_image_tokens(
68
+ inputs_embeds: torch.Tensor,
69
+ image_embeds: torch.Tensor | Sequence[torch.Tensor],
70
+ image_embeds_insertion_points: list[torch.Tensor],
71
+ recover_batch_dim: bool = True,
72
+ attention_mask: torch.Tensor | None = None,
73
+ padding_side: Literal["left", "right"] = "right",
74
+ keep_only_attended: bool = False,
75
+ pad_output: int | float | bool = 0.0,
76
+ ) -> tuple[
77
+ torch.Tensor | torch.Tensor,
78
+ list[int] | None,
79
+ torch.Tensor | torch.Tensor | None,
80
+ torch.Tensor | torch.Tensor,
81
+ ]:
82
+ """
83
+ Insert image embeddings into text embeddings
84
+
85
+ Args:
86
+ inputs_embeds (torch.Tensor): (B, S, D) input token embeddings.
87
+ image_embeds (torch.Tensor | list[torch.Tensor]): (N_images, Nt, D) | List[(Nt, D)] image token embeddings.
88
+ image_embeds_insertion_points (list[torch.Tensor]): Insertion indices.
89
+ attention_mask (torch.Tensor, optional): (B, S) attention mask.
90
+ padding_side (Literal["left", "right"]): Padding scheme. Controls behavior for padded images.
91
+ return_indices (bool): Whether to return gather indices or the fused sequence directly.
92
+ keep_only_attended: This is only applicable when recover_batch_dim is False; whether to
93
+ remove any non-attended tokens in the whole array. In this case, the attention
94
+ mask returned is **still the original one**, so we can remember which indices have been
95
+ removed
96
+ Returns:
97
+ output (torch.Tensor): (B, S + Ni * Nt) gather indices or (B, S + Ni * Nt, D) fused sequence
98
+ image_embeds (torch.Tensor): (B, Ni * Nt) image embeds, padded and batch if input was a list
99
+ attention_mask (torch.Tensor): Same shape, 1 for real tokens, 0 for image and text padding.
100
+ image_tokens_mask (torch.Tensor): (B, S + Ni * Nt, 1), marks image token positions.
101
+ """
102
+ if isinstance(image_embeds, list) and len(image_embeds) == 0:
103
+ batch_size, text_seq_length, token_dim = inputs_embeds.shape
104
+ if recover_batch_dim:
105
+ return (
106
+ inputs_embeds,
107
+ None,
108
+ attention_mask,
109
+ torch.zeros((batch_size, text_seq_length, 1), dtype=torch.bool),
110
+ )
111
+ else:
112
+ flattened_seq_length = inputs_embeds.shape[0] * inputs_embeds.shape[1]
113
+ return (
114
+ torch.reshape(inputs_embeds, (flattened_seq_length, inputs_embeds.shape[2])),
115
+ [text_seq_length] * inputs_embeds.shape[0],
116
+ attention_mask.flatten() if attention_mask is not None else None,
117
+ torch.zeros((flattened_seq_length, 1), dtype=torch.bool),
118
+ )
119
+
120
+ # Sanity checks
121
+ if isinstance(image_embeds, torch.Tensor):
122
+ assert inputs_embeds.shape[-1] == image_embeds.shape[-1]
123
+ else:
124
+ assert all(inputs_embeds.shape[-1] == _x.shape[-1] for _x in image_embeds)
125
+
126
+ batch_size, text_seq_length, token_dim = inputs_embeds.shape
127
+ image_seq_length = [x.shape[0] for x in image_embeds]
128
+
129
+ # Flatten insertion points
130
+ insertion_offset = []
131
+ counter, offset_from_text, offset_from_image = 0, 0, 0
132
+ for sample in image_embeds_insertion_points:
133
+ for pt in sample:
134
+ insertion_offset.append(pt + offset_from_image + offset_from_text)
135
+ offset_from_image += image_seq_length[counter]
136
+ counter += 1
137
+ offset_from_text += text_seq_length
138
+ image_insert_positions = [
139
+ x for idx, pt in enumerate(insertion_offset) for x in range(pt, pt + image_seq_length[idx])
140
+ ]
141
+
142
+ # Flatten image embeds
143
+ if isinstance(image_embeds, list):
144
+ image_embeds = torch.cat(image_embeds, dim=0)
145
+ else:
146
+ image_embeds = type_cast(torch.Tensor, image_embeds)
147
+ image_embeds = torch.reshape(image_embeds, (-1, token_dim))
148
+
149
+ # Flatten text embeds across batch dim (B x S, D)
150
+ inputs_embeds = torch.reshape(inputs_embeds, (-1, token_dim))
151
+ flattened_seq_length = inputs_embeds.shape[0] + sum(image_seq_length)
152
+ text_insert_positions = sorted(
153
+ set(range(flattened_seq_length)).difference(set(image_insert_positions))
154
+ )
155
+
156
+ # Scatter image embeds in the flattened dict
157
+ # scatter text related stuff
158
+ output = torch.empty(
159
+ (flattened_seq_length, token_dim),
160
+ device=inputs_embeds.device,
161
+ dtype=inputs_embeds.dtype,
162
+ )
163
+ txt_positions_tensor = torch.Tensor(text_insert_positions).to(
164
+ dtype=torch.long, device=inputs_embeds.device
165
+ )
166
+ output.scatter_(0, txt_positions_tensor[:, None].expand(-1, token_dim), inputs_embeds)
167
+ attention_mask_new: torch.Tensor | None = None
168
+ if attention_mask is not None:
169
+ attention_mask_new = torch.ones(
170
+ (flattened_seq_length,), dtype=torch.bool, device=inputs_embeds.device
171
+ )
172
+ attention_mask_new.scatter_(
173
+ 0, txt_positions_tensor, attention_mask.flatten().to(torch.bool)
174
+ )
175
+
176
+ # scatter image related stuff
177
+ image_tokens_mask = torch.zeros(
178
+ (flattened_seq_length,), dtype=torch.bool, device=inputs_embeds.device
179
+ )
180
+ img_positions_tensor = torch.Tensor(image_insert_positions).to(
181
+ device=inputs_embeds.device, dtype=torch.long
182
+ )
183
+ output.scatter_(0, img_positions_tensor[:, None].expand(-1, token_dim), image_embeds)
184
+ image_tokens_mask.scatter_(0, img_positions_tensor, True)
185
+
186
+ # Compute expected sample length, taking into account the real batch
187
+ # i.e. recover the batch dimension of image embeddings
188
+ sample_lengths = []
189
+ counter = 0
190
+ for sample_idx, pts in enumerate(image_embeds_insertion_points):
191
+ num_image_tokens = 0
192
+ for _ in pts:
193
+ num_image_tokens += image_seq_length[counter]
194
+ counter += 1
195
+ if keep_only_attended and attention_mask is not None:
196
+ attended_seq_length = torch.sum(attention_mask[sample_idx]).cpu().item()
197
+ sample_lengths.append(attended_seq_length + num_image_tokens)
198
+ else:
199
+ sample_lengths.append(text_seq_length + num_image_tokens)
200
+
201
+ # For CA attention, we can keep stuff flatten and return
202
+ # the sample_lengths for the blockwise attention
203
+ if not recover_batch_dim:
204
+ if keep_only_attended and attention_mask_new is not None:
205
+ output = output[attention_mask_new]
206
+ image_tokens_mask = image_tokens_mask[attention_mask_new]
207
+ return output, sample_lengths, attention_mask_new, image_tokens_mask[..., None]
208
+
209
+ # Otherwise, time to (pad) and reshape
210
+ # Easy case: everything has the same length
211
+ if all(x == sample_lengths[0] for x in sample_lengths):
212
+ output = torch.reshape(output, (batch_size, sample_lengths[0], token_dim))
213
+ image_tokens_mask = torch.reshape(image_tokens_mask, (batch_size, sample_lengths[0], 1))
214
+ if attention_mask_new is not None:
215
+ attention_mask_new = torch.reshape(attention_mask_new, (batch_size, sample_lengths[0]))
216
+ # if there is any size mismatch we break into a
217
+ # list and pad again
218
+ else:
219
+ # split and merge
220
+ output = __split_n_merge__(output, sample_lengths, padding_side, pad_value=pad_output)
221
+ # note that the extra padding tokens are also marked as image tokens to be removed later
222
+ image_tokens_mask = __split_n_merge__(
223
+ image_tokens_mask, sample_lengths, padding_side, True
224
+ )[:, :, None]
225
+ if attention_mask_new is not None:
226
+ attention_mask_new = __split_n_merge__(
227
+ attention_mask_new, sample_lengths, padding_side, 0
228
+ )
229
+ # Return
230
+ return output, sample_lengths, attention_mask_new, image_tokens_mask
231
+
232
+
233
+ class SharedModuleType(type):
234
+ """Wrapper to build shared Pytorch modules. This can be used as a metaclass to build shared
235
+ modules; see an example in attention.py"""
236
+
237
+ _instances = {}
238
+
239
+ def __call__(cls, *args: Any, **kwargs: Any) -> Any:
240
+ if cls not in cls._instances:
241
+ cls._instances[cls] = super(SharedModuleType, cls).__call__(*args, **kwargs)
242
+ return cls._instances[cls]
243
+
244
+
245
+ @dataclass
246
+ class StreamingState:
247
+ """Streaming State used by CA layers at inference to save
248
+ e.g. the offset and other persistent states"""
249
+
250
+ offset: int = 0
251
+
252
+ def _is_valid_field(self, key: str) -> bool:
253
+ return key in {x.name for x in fields(self)}
254
+
255
+ def _init_field(self, key: str) -> None:
256
+ """Init function for non-argument dependent defaults"""
257
+ assert self._is_valid_field(key)
258
+ if key == "offset":
259
+ self.offset = 0
260
+ else:
261
+ # for fields which should be set explicitly and cannot be auto-initialized
262
+ setattr(self, key, None)
263
+
264
+ def init(self) -> None:
265
+ for key in [x.name for x in fields(self)]:
266
+ self._init_field(key)
267
+
268
+ def _reset_field(self, name: str) -> None:
269
+ """Resets the given field"""
270
+ self._init_field(name)
271
+
272
+ def reset(self) -> None:
273
+ for f in fields(self):
274
+ self._reset_field(f.name)
275
+
276
+ def _get_field(self, f: str) -> Any:
277
+ """Get field and init if not"""
278
+ assert self._is_valid_field(f)
279
+ if getattr(self, f) is None:
280
+ self._init_field(f)
281
+ return getattr(self, f)
282
+
283
+ def _set_field(self, f: str, value: Any) -> None:
284
+ assert self._is_valid_field(f)
285
+ setattr(self, f, value)
286
+
287
+
288
+ StreamingStateT = TypeVar("StreamingStateT", bound=StreamingState)
289
+
290
+
291
+ class StreamingModule(torch.nn.Module, Generic[StreamingStateT]): # pylint: disable=abstract-method
292
+ """Streaming-aware module base class"""
293
+
294
+ def __init__(self, state_class: type) -> None:
295
+ torch.nn.Module.__init__(self)
296
+ self.is_streaming: bool = False
297
+ self.enable_viz: tuple[str, ...] = ()
298
+ self._streaming_state: StreamingStateT = state_class()
299
+
300
+ @property
301
+ def streaming_state(self) -> StreamingStateT:
302
+ return self._streaming_state
303
+
304
+ def _apply_named_streaming(self, fn: Callable):
305
+ """Apply function to all streaming modules"""
306
+ for name, module in self.named_modules():
307
+ if isinstance(module, StreamingModule):
308
+ fn(name, module)
309
+
310
+ def reset_streaming(self):
311
+ """Reset the streaming state."""
312
+
313
+ def _reset(_: str, module: StreamingModule):
314
+ module._streaming_state.reset()
315
+
316
+ self._apply_named_streaming(_reset)
317
+
318
+ def _set_streaming(self, streaming: bool, viz: tuple[str, ...] = ()):
319
+ """Set all streaming modules in streaming mode"""
320
+
321
+ def _set_streaming(_, module: StreamingModule) -> None:
322
+ module.is_streaming = streaming
323
+ module.enable_viz = viz
324
+ if streaming:
325
+ module.streaming_state.init()
326
+
327
+ self._apply_named_streaming(_set_streaming)
328
+
329
+ @contextmanager
330
+ def streaming(self, stream: bool = True, viz: tuple[str, ...] = ()):
331
+ """Context manager to enter streaming mode. Reset streaming state on exit."""
332
+ self._set_streaming(stream, viz)
333
+ try:
334
+ yield
335
+ finally:
336
+ self._set_streaming(False, ())
337
+ self.reset_streaming()
vocab.json ADDED
The diff for this file is too large to render. See raw diff