Views
No views yet
1import itertools
2import os
3import numpy as np
4import onnxruntime
5
6from huggingface_hub import snapshot_download
7from transformers import AutoConfig, AutoProcessor, GenerationConfig
8
9
10def get_vision_position_ids(start_position, grid_thw, temp_merge_size=1, spatial_merge_size=1, time_interval=1):
11 llm_grid_t = grid_thw[0] // temp_merge_size
12 llm_grid_h = grid_thw[1] // spatial_merge_size
13 llm_grid_w = grid_thw[2] // spatial_merge_size
14
15 image_seq_length = llm_grid_h * llm_grid_w * llm_grid_t
16 position_width = np.tile(np.arange(start_position, start_position + llm_grid_w), llm_grid_h * llm_grid_t)
17 position_height = np.repeat(np.arange(start_position, start_position + llm_grid_h), llm_grid_w * llm_grid_t)
18 position_temporal = np.full(image_seq_length, start_position, dtype=np.int64) * time_interval
19 return np.stack([position_temporal, position_height, position_width], axis=0)
20
21
22def get_rope_index(input_ids, mm_token_type_ids, image_grid_thw=None, video_grid_thw=None,
23 second_per_grid_ts=None, attention_mask=None, spatial_merge_size=2, tokens_per_second=25):
24 batch_size, seq_len = input_ids.shape
25 position_ids = np.zeros((3, batch_size, seq_len), dtype=np.int64)
26 mrope_position_deltas = []
27
28 grid_iters = {
29 1: iter(image_grid_thw) if image_grid_thw is not None else None,
30 2: iter(video_grid_thw) if video_grid_thw is not None else None,
31 }
32 second_per_grid_ts_iter = iter(second_per_grid_ts) if second_per_grid_ts is not None else iter([1] * seq_len)
33
34 for batch_idx in range(batch_size):
35 current_input_ids = input_ids[batch_idx]
36 input_token_type = mm_token_type_ids[batch_idx]
37
38 if attention_mask is not None:
39 mask = attention_mask[batch_idx].astype(bool)
40 current_input_ids = current_input_ids[mask]
41 input_token_type = input_token_type[mask]
42
43 input_type_groups = []
44 for key, group in itertools.groupby(enumerate(input_token_type.tolist()), lambda x: x[1]):
45 group = list(group)
46 input_type_groups.append((key, group[0][0], group[-1][0] + 1))
47
48 current_pos = 0
49 llm_pos_ids_list = []
50 for modality_type, start_idx, end_idx in input_type_groups:
51 if modality_type == 0: # text
52 text_len = end_idx - start_idx
53 text_pos = np.arange(current_pos, current_pos + text_len, dtype=np.int64)
54 llm_pos_ids_list.append(np.tile(text_pos, (3, 1)))
55 current_pos += text_len
56 else: # image (1) or video (2)
57 grid_thw = next(grid_iters[modality_type])
58 time_interval = tokens_per_second * int(next(second_per_grid_ts_iter))
59 vision_pos = get_vision_position_ids(current_pos, grid_thw, 1, spatial_merge_size, time_interval)
60 llm_pos_ids_list.append(vision_pos)
61 current_pos += max(grid_thw[1], grid_thw[2]) // spatial_merge_size
62
63 llm_positions = np.concatenate(llm_pos_ids_list, axis=1) # (3, total_tokens)
64 if attention_mask is not None:
65 position_ids[:, batch_idx, attention_mask[batch_idx].astype(bool)] = llm_positions
66 else:
67 position_ids[:, batch_idx] = llm_positions
68
69 mrope_position_deltas.append(llm_positions.max() + 1 - len(current_input_ids))
70
71 mrope_position_deltas = np.array(mrope_position_deltas, dtype=np.int64).reshape(-1, 1)
72 return position_ids, mrope_position_deltas
73
74
75# 1. Load models
76## Define Model ID
77model_id = "onnx-community/Qwen2.5-VL-3B-Instruct-ONNX"
78
79## Load config, processor, and generation config
80config = AutoConfig.from_pretrained(model_id)
81processor = AutoProcessor.from_pretrained(model_id)
82generation_config = GenerationConfig.from_pretrained(model_id)
83
84## Select model precisions
85vision_encoder_path = "vision_encoder.onnx"
86embed_tokens_path = "embed_tokens.onnx"
87decoder_model_path = "decoder_model_merged_q4.onnx"
88
89## Download ONNX models
90print("Downloading ONNX models...")
91onnx_dir = snapshot_download(
92 repo_id=model_id,
93 allow_patterns=[
94 # Download requested graphs and weights
95 f"onnx/{vision_encoder_path}*",
96 f"onnx/{embed_tokens_path}*",
97 f"onnx/{decoder_model_path}*",
98 ]
99)
100
101## Load sessions
102vision_session = onnxruntime.InferenceSession(os.path.join(onnx_dir, "onnx", vision_encoder_path))
103embed_session = onnxruntime.InferenceSession(os.path.join(onnx_dir, "onnx", embed_tokens_path))
104decoder_session = onnxruntime.InferenceSession(os.path.join(onnx_dir, "onnx", decoder_model_path))
105
106## Set config values
107text_config = config.text_config
108num_key_value_heads = text_config.num_key_value_heads
109head_dim = text_config.hidden_size // text_config.num_attention_heads
110image_token_id = config.image_token_id
111eos_token_id = generation_config.eos_token_id
112spatial_merge_size = config.vision_config.spatial_merge_size
113tokens_per_second = config.vision_config.tokens_per_second
114
115# 2. Prepare inputs
116## Create input messages
117messages = [
118 {
119 "role": "user",
120 "content": [
121 {"type": "image", "image": "https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-VL/assets/demo.jpeg"},
122 {"type": "text", "text": "Describe this image."}
123 ]
124 },
125]
126
127## Apply chat template
128pt_inputs = processor.apply_chat_template(
129 messages,
130 add_generation_prompt=True,
131 tokenize=True,
132 return_dict=True,
133 return_tensors="pt",
134)
135inputs = {k: v.cpu().numpy() for k, v in pt_inputs.items()}
136
137## Calculate vision positions and RoPE deltas
138position_ids, rope_deltas = get_rope_index(
139 input_ids=inputs["input_ids"],
140 mm_token_type_ids=inputs.get("mm_token_type_ids"),
141 image_grid_thw=inputs.get("image_grid_thw"),
142 attention_mask=inputs["attention_mask"],
143 spatial_merge_size=spatial_merge_size,
144 tokens_per_second=tokens_per_second,
145)
146
147## Prepare decoder variables
148batch_size = inputs['input_ids'].shape[0]
149input_ids = inputs['input_ids']
150attention_mask = inputs['attention_mask']
151
152## Initialize past_key_values cache
153past_key_values = {
154 inp.name: np.zeros(
155 [batch_size, num_key_value_heads, 0, head_dim],
156 dtype=np.float16 if "float16" in inp.type else np.float32
157 )
158 for inp in decoder_session.get_inputs()
159 if inp.name.startswith("past_key_values")
160}
161
162# 3. Generation loop
163max_new_tokens = 512
164generated_tokens = np.array([[]], dtype=np.int64)
165image_features = None
166
167print("Generating...")
168for step in range(max_new_tokens):
169 ## Generate text embeddings
170 inputs_embeds = embed_session.run(None, {'input_ids': input_ids})[0]
171
172 ## Compute and inject vision features (only on the very first step)
173 if image_features is None and "pixel_values" in inputs:
174 vision_inputs = {"pixel_values": inputs["pixel_values"]}
175 vision_input_names = {inp.name for inp in vision_session.get_inputs()}
176
177 for optional_input in ["pixel_attention_mask", "spatial_shapes", "image_sizes", "image_grid_thw"]:
178 if optional_input in vision_input_names and optional_input in inputs:
179 vision_inputs[optional_input] = inputs[optional_input]
180
181 image_features = vision_session.run(None, vision_inputs)[0]
182
183 # Merge vision embeddings into the text embedding sequence
184 inputs_embeds[input_ids == image_token_id] = image_features.reshape(-1, image_features.shape[-1])
185
186 ## Run decoder step
187 outputs = decoder_session.run(None, dict(
188 inputs_embeds=inputs_embeds,
189 attention_mask=attention_mask,
190 position_ids=position_ids,
191 **past_key_values,
192 ))
193 logits, present_key_values = outputs[0], outputs[1:]
194
195 ## Update states for the next iteration
196 next_token = logits[:, -1].argmax(-1, keepdims=True)
197 input_ids = next_token
198
199 attention_mask = np.concatenate([attention_mask, np.ones((batch_size, 1), dtype=attention_mask.dtype)], axis=-1)
200
201 ## Re-calculate positional IDs and apply RoPE deltas for the generated token
202 text_positions = np.cumsum(attention_mask, axis=-1, dtype=np.int64) - 1
203 text_positions = np.clip(text_positions, 0, None)[:, -1:]
204 position_ids = np.broadcast_to(text_positions[None, ...], (3,) + text_positions.shape) + rope_deltas[None, ...]
205
206 for j, key in enumerate(past_key_values):
207 past_key_values[key] = present_key_values[j]
208
209 generated_tokens = np.concatenate([generated_tokens, input_ids], axis=-1)
210
211 ## (Optional) Streaming
212 print(processor.decode(input_ids[0]), end='', flush=True)
213
214 if np.isin(input_ids, eos_token_id).all():
215 break
216print()
217
218
219# 4. Output result
220print("\n--- Final Decoded Output ---")
221print(processor.decode(generated_tokens[0], skip_special_tokens=True))