DennisHuang648 commited on
Commit
fbf124f
·
verified ·
1 Parent(s): 00302df

Upload modeling_minicpmv.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. modeling_minicpmv.py +625 -0
modeling_minicpmv.py ADDED
@@ -0,0 +1,625 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import math
2
+ from typing import List, Optional
3
+ import json
4
+ import os
5
+ from threading import Thread
6
+ from copy import deepcopy
7
+
8
+ import torch
9
+ import torch.nn as nn
10
+ import torchvision
11
+ import anndata as ad
12
+ from PIL import Image
13
+
14
+ from transformers import AutoProcessor, Qwen2PreTrainedModel, Qwen2ForCausalLM, TextIteratorStreamer
15
+
16
+ from .configuration_minicpm import MiniCPMVConfig
17
+ from .modeling_navit_siglip import SiglipVisionTransformer
18
+ from .resampler import Resampler
19
+ from .processing_minicpmv import MiniCPMVProcessor
20
+
21
+ # gene
22
+ from .modeling_nicheformer import NicheformerModel
23
+ from .configuration_nicheformer import NicheformerConfig
24
+ from .gene_projector_module import GeneProjector
25
+ from .gene_qformer_module import GeneQFormerBiomedBERT
26
+
27
+ def _is_debug_enabled() -> bool:
28
+ return os.getenv("DEBUG_GENE", "0") == "1"
29
+
30
+
31
+ def _assert_finite(x: torch.Tensor, name: str):
32
+ if not torch.is_tensor(x):
33
+ return
34
+ if not torch.isfinite(x).all():
35
+ # 打印一些统计,方便定位
36
+ with torch.no_grad():
37
+ finite_mask = torch.isfinite(x)
38
+ num_bad = (~finite_mask).sum().item()
39
+ msg = (
40
+ f"[NaN/Inf Detected] {name} has non-finite values. "
41
+ f"bad_count={num_bad}, dtype={x.dtype}, device={x.device}, shape={tuple(x.shape)}"
42
+ )
43
+ raise RuntimeError(msg)
44
+
45
+
46
+ class MiniCPMVPreTrainedModel(Qwen2PreTrainedModel):
47
+ config_class = MiniCPMVConfig
48
+
49
+
50
+ class MiniCPMV(MiniCPMVPreTrainedModel):
51
+ def __init__(self, config):
52
+ super().__init__(config)
53
+ self.llm = Qwen2ForCausalLM(config)
54
+
55
+ self.vpm = self.init_vision_module()
56
+ self.vision_dim = self.vpm.embed_dim
57
+ self.embed_dim = self.llm.config.hidden_size
58
+
59
+ self.resampler = self.init_resampler(self.embed_dim, self.vision_dim)
60
+
61
+ # ===== Gene modules =====
62
+ self.nicheformer = self.init_gene_module(config)
63
+ self.gene_dim = self.nicheformer.config.dim_model # e.g. 512
64
+
65
+ self.gene_qformer = GeneQFormerBiomedBERT(
66
+ gene_in_dim=self.gene_dim, # 512
67
+ hidden=768,
68
+ num_queries=32,
69
+ load_pretrained_bert=False,
70
+ )
71
+
72
+ # Project: 768 -> LLM hidden (e.g. 3584)
73
+ self.gene_projector = GeneProjector(in_dim=768, out_dim=self.embed_dim)
74
+
75
+ self.processor = None
76
+ self.terminators = ['<|im_end|>', '<|endoftext|>']
77
+ self._generate = self.generate
78
+
79
+ self._gene_fp32_forced = False
80
+
81
+ # self.post_init()
82
+
83
+ def init_gene_module(self, config):
84
+ if hasattr(config, "gene_config"):
85
+ return NicheformerModel(config.gene_config)
86
+ else:
87
+ nicheformer_config = NicheformerConfig()
88
+ return NicheformerModel(nicheformer_config)
89
+
90
+ def init_vision_module(self):
91
+ if self.config._attn_implementation == 'flash_attention_2':
92
+ self.config.vision_config._attn_implementation = 'flash_attention_2'
93
+ else:
94
+ self.config.vision_config._attn_implementation = 'eager'
95
+
96
+ model = SiglipVisionTransformer(self.config.vision_config)
97
+ if self.config.drop_vision_last_layer:
98
+ model.encoder.layers = model.encoder.layers[:-1]
99
+
100
+ setattr(model, 'embed_dim', model.embeddings.embed_dim)
101
+ setattr(model, 'patch_size', model.embeddings.patch_size)
102
+ return model
103
+
104
+ def init_resampler(self, embed_dim, vision_dim):
105
+ return Resampler(
106
+ num_queries=self.config.query_num,
107
+ embed_dim=embed_dim,
108
+ num_heads=embed_dim // 128,
109
+ kv_dim=vision_dim,
110
+ adaptive=True
111
+ )
112
+
113
+ def get_input_embeddings(self):
114
+ return self.llm.get_input_embeddings()
115
+
116
+ def set_input_embeddings(self, value):
117
+ self.llm.embed_tokens = value
118
+
119
+ def get_output_embeddings(self):
120
+ return self.llm.lm_head
121
+
122
+ def set_output_embeddings(self, new_embeddings):
123
+ self.llm.lm_head = new_embeddings
124
+
125
+ def set_decoder(self, decoder):
126
+ self.llm = decoder
127
+
128
+ def get_decoder(self):
129
+ return self.llm
130
+
131
+
132
+ def get_vllm_embedding(self, data):
133
+ dtype = self.llm.model.embed_tokens.weight.dtype
134
+ device = self.llm.model.embed_tokens.weight.device
135
+ self.gene_qformer = self.gene_qformer.float()
136
+ self.gene_projector = self.gene_projector.float()
137
+
138
+ # =========================
139
+ # 1) Vision
140
+ # =========================
141
+ if 'vision_hidden_states' not in data:
142
+ tgt_sizes = data['tgt_sizes']
143
+ pixel_values_list = data['pixel_values']
144
+ vision_hidden_states = []
145
+ all_pixel_values = []
146
+ img_cnt = []
147
+
148
+ for pixel_values in pixel_values_list:
149
+ img_cnt.append(len(pixel_values))
150
+ all_pixel_values.extend([i.flatten(end_dim=1).permute(1, 0) for i in pixel_values])
151
+
152
+ if all_pixel_values:
153
+ tgt_sizes = [tgt_size for tgt_size in tgt_sizes if isinstance(tgt_size, torch.Tensor)]
154
+ tgt_sizes = torch.vstack(tgt_sizes).type(torch.int32)
155
+
156
+ max_patches = torch.max(tgt_sizes[:, 0] * tgt_sizes[:, 1])
157
+
158
+ all_pixel_values = torch.nn.utils.rnn.pad_sequence(
159
+ all_pixel_values, batch_first=True, padding_value=0.0
160
+ )
161
+ B, L, _ = all_pixel_values.shape
162
+ all_pixel_values = all_pixel_values.permute(0, 2, 1).reshape(B, 3, -1, L)
163
+
164
+ patch_attn_mask = torch.zeros((B, 1, max_patches), dtype=torch.bool, device=device)
165
+ for i in range(B):
166
+ patch_attn_mask[i, 0, :tgt_sizes[i][0] * tgt_sizes[i][1]] = True
167
+
168
+ vision_batch_size = self.config.vision_batch_size
169
+ all_pixel_values = all_pixel_values.to(device=device, dtype=dtype)
170
+
171
+ if B > vision_batch_size:
172
+ hs = []
173
+ for i in range(0, B, vision_batch_size):
174
+ start_idx = i
175
+ end_idx = i + vision_batch_size
176
+ tmp_hs = self.vpm(
177
+ all_pixel_values[start_idx:end_idx],
178
+ patch_attention_mask=patch_attn_mask[start_idx:end_idx],
179
+ tgt_sizes=tgt_sizes[start_idx:end_idx]
180
+ ).last_hidden_state
181
+ hs.append(tmp_hs)
182
+ vision_embedding = torch.cat(hs, dim=0)
183
+ else:
184
+ vision_embedding = self.vpm(
185
+ all_pixel_values,
186
+ patch_attention_mask=patch_attn_mask,
187
+ tgt_sizes=tgt_sizes
188
+ ).last_hidden_state
189
+
190
+ vision_embedding = self.resampler(vision_embedding, tgt_sizes)
191
+
192
+ start = 0
193
+ for pixel_values in pixel_values_list:
194
+ c = len(pixel_values)
195
+ if c > 0:
196
+ vision_hidden_states.append(vision_embedding[start: start + c])
197
+ start += c
198
+ else:
199
+ vision_hidden_states.append([])
200
+ else:
201
+ # no image
202
+ if self.training:
203
+ dummy_image = torch.zeros((1, 3, 224, 224), device=device, dtype=dtype)
204
+ tgt_sizes_dummy = torch.Tensor(
205
+ [[(224 // self.config.patch_size), math.ceil(224 / self.config.patch_size)]]
206
+ ).type(torch.int32)
207
+ dummy_feature = self.resampler(self.vpm(dummy_image).last_hidden_state, tgt_sizes_dummy)
208
+ else:
209
+ dummy_feature = []
210
+ for _ in range(len(pixel_values_list)):
211
+ vision_hidden_states.append(dummy_feature)
212
+ else:
213
+ vision_hidden_states = data['vision_hidden_states']
214
+
215
+ # =========================
216
+ # 2) Gene
217
+ # =========================
218
+ bs = len(data['input_ids'])
219
+ gene_hidden_states = [None] * bs
220
+
221
+ if 'gene_input_ids' in data and data['gene_input_ids'] is not None:
222
+ gene_input_ids = data['gene_input_ids'].to(device)
223
+ gene_attention_mask = data.get('gene_attention_mask', None)
224
+ if gene_attention_mask is not None:
225
+ gene_attention_mask = gene_attention_mask.to(device)
226
+
227
+ # Nicheformer: expect [B, seq_len, gene_dim]
228
+ nicheformer_output = self.nicheformer.forward(
229
+ input_ids=gene_input_ids,
230
+ attention_mask=gene_attention_mask
231
+ )
232
+
233
+ # 丢掉前 3 个 special token(按你当前实现)
234
+ gene_tokens = nicheformer_output[:, 3:, :] # [B, L, 512]
235
+
236
+ gene_pad_mask = None
237
+ if gene_attention_mask is not None:
238
+ gene_pad_mask = (gene_attention_mask[:, 3:] == 0)
239
+
240
+ # dtype 对齐:以 gene_qformer 的参数 dtype 为准
241
+ qformer_dtype = next(self.gene_qformer.parameters()).dtype
242
+ gene_tokens = gene_tokens.to(dtype=qformer_dtype)
243
+
244
+ # Q-Former: [B, L, 512] -> [B, 32, 768]
245
+ q_tokens = self.gene_qformer(gene_tokens, gene_pad_mask=gene_pad_mask)
246
+
247
+ # Projector: [B, 32, 768] -> [B, 32, embed_dim]
248
+ proj_dtype = next(self.gene_projector.parameters()).dtype
249
+ q_tokens = q_tokens.to(dtype=proj_dtype)
250
+ gene_tokens_llm = self.gene_projector(q_tokens) # [B, 32, 3584]
251
+ _assert_finite(gene_tokens, "gene_tokens(after nicheformer)")
252
+ _assert_finite(q_tokens, "q_tokens(after qformer)")
253
+ _assert_finite(gene_tokens_llm, "gene_tokens_llm(after projector)")
254
+
255
+ # # 插入到对应位置(默认每个样本最多 1 个 <gene> span)
256
+ # gene_bounds = data.get('gene_bound', [[] for _ in range(bs)])
257
+ # for i, bounds in enumerate(gene_bounds):
258
+ # if not bounds:
259
+ # continue
260
+ # gene_hidden_states[i] = gene_tokens_llm[i] # [32, embed_dim]
261
+
262
+ gene_bounds = data.get('gene_bound', [None] * bs)
263
+
264
+ for i, bounds in enumerate(gene_bounds):
265
+ # bounds can be: None, [], or a Tensor of shape [N,2]
266
+ if bounds is None:
267
+ continue
268
+ if isinstance(bounds, list) and len(bounds) == 0:
269
+ continue
270
+ if torch.is_tensor(bounds) and bounds.numel() == 0:
271
+ continue
272
+
273
+ # 默认每个样本只用第一个 gene span(你当前设定)
274
+ gene_hidden_states[i] = gene_tokens_llm[i] # [32, embed_dim]
275
+
276
+ # =========================
277
+ # 3) Text token embeddings
278
+ # =========================
279
+ if hasattr(self.llm.config, 'scale_emb'):
280
+ vllm_embedding = self.llm.model.embed_tokens(data['input_ids']) * self.llm.config.scale_emb
281
+ else:
282
+ vllm_embedding = self.llm.model.embed_tokens(data['input_ids'])
283
+
284
+ new_vllm_embedding = vllm_embedding.clone()
285
+
286
+ # dtype/device align
287
+ vision_hidden_states = [
288
+ x.to(dtype=vllm_embedding.dtype, device=vllm_embedding.device) if torch.is_tensor(x) else x
289
+ for x in vision_hidden_states
290
+ ]
291
+ gene_hidden_states = [
292
+ x.to(dtype=vllm_embedding.dtype, device=vllm_embedding.device) if torch.is_tensor(x) else x
293
+ for x in gene_hidden_states
294
+ ]
295
+
296
+ # =========================
297
+ # 4) Scatter insert: image + gene
298
+ # =========================
299
+ for i in range(bs):
300
+ # ---- image ----
301
+ cur_vs_hs = vision_hidden_states[i]
302
+ if torch.is_tensor(cur_vs_hs) and cur_vs_hs.numel() > 0:
303
+ cur_vllm_emb = vllm_embedding[i]
304
+ cur_image_bound = data['image_bound'][i]
305
+ if len(cur_image_bound) > 0:
306
+ image_indices = torch.cat([
307
+ torch.arange(r[0], r[1], dtype=torch.long, device=vllm_embedding.device)
308
+ for r in cur_image_bound if (r[1] - r[0]) > 1
309
+ ])
310
+ new_vllm_embedding[i] = cur_vllm_emb.scatter(
311
+ 0,
312
+ image_indices.view(-1, 1).repeat(1, cur_vllm_emb.shape[-1]),
313
+ cur_vs_hs.view(-1, cur_vs_hs.shape[-1])
314
+ )
315
+ elif self.training:
316
+ new_vllm_embedding[i] += cur_vs_hs[0].mean() * 0
317
+
318
+ # ---- gene ----
319
+ cur_gene_hs = gene_hidden_states[i]
320
+ if cur_gene_hs is not None:
321
+ cur_gene_bound = data.get('gene_bound', [[] for _ in range(bs)])[i]
322
+ if len(cur_gene_bound) > 0:
323
+ r = cur_gene_bound[0] # [start, end)
324
+ gene_indices = torch.arange(
325
+ r[0], r[1], dtype=torch.long, device=vllm_embedding.device
326
+ )
327
+ span = gene_indices.numel()
328
+ if span != cur_gene_hs.shape[0]:
329
+ raise ValueError(
330
+ f"[GeneSpanMismatch] gene span={span}, gene tokens={cur_gene_hs.shape[0]} "
331
+ f"(expect 32). Check processor placeholder length."
332
+ )
333
+
334
+ cur_vllm_emb = new_vllm_embedding[i]
335
+ new_vllm_embedding[i] = cur_vllm_emb.scatter(
336
+ 0,
337
+ gene_indices.view(-1, 1).repeat(1, cur_vllm_emb.shape[-1]),
338
+ cur_gene_hs.to(cur_vllm_emb.dtype) # [32, embed_dim]
339
+ )
340
+
341
+ if _is_debug_enabled():
342
+ if "gene_bound" in data and len(data["gene_bound"]) > 0 and len(data["gene_bound"][0]) > 0:
343
+ gb0 = data["gene_bound"][0]
344
+ print("[DEBUG] gene_bound[0]:", gb0)
345
+ print("[DEBUG] gene_span[0]:", gb0[0][1] - gb0[0][0])
346
+
347
+ _assert_finite(new_vllm_embedding, "new_vllm_embedding(final inputs_embeds)")
348
+
349
+ return new_vllm_embedding, vision_hidden_states
350
+
351
+ def forward(self, data, **kwargs):
352
+ vllm_embedding, vision_hidden_states = self.get_vllm_embedding(data)
353
+
354
+ position_ids = data["position_ids"]
355
+ if position_ids.dtype != torch.int64:
356
+ position_ids = position_ids.long()
357
+
358
+ for key in ['input_ids', 'inputs_embeds', 'position_ids']:
359
+ if key in kwargs:
360
+ del kwargs[key]
361
+
362
+ return self.llm(
363
+ input_ids=None,
364
+ position_ids=position_ids,
365
+ inputs_embeds=vllm_embedding,
366
+ **kwargs
367
+ )
368
+
369
+ def _decode(self, inputs_embeds, tokenizer, attention_mask, decode_text=False, **kwargs):
370
+ terminators = [tokenizer.convert_tokens_to_ids(i) for i in self.terminators]
371
+ output = self.llm.generate(
372
+ inputs_embeds=inputs_embeds,
373
+ pad_token_id=0,
374
+ eos_token_id=terminators,
375
+ attention_mask=attention_mask,
376
+ **kwargs
377
+ )
378
+ if decode_text:
379
+ return self._decode_text(output, tokenizer)
380
+ return output
381
+
382
+ def _decode_stream(self, inputs_embeds, tokenizer, **kwargs):
383
+ terminators = [tokenizer.convert_tokens_to_ids(i) for i in self.terminators]
384
+ streamer = TextIteratorStreamer(tokenizer=tokenizer)
385
+ generation_kwargs = {
386
+ 'inputs_embeds': inputs_embeds,
387
+ 'pad_token_id': 0,
388
+ 'eos_token_id': terminators,
389
+ 'streamer': streamer
390
+ }
391
+ generation_kwargs.update(kwargs)
392
+
393
+ thread = Thread(target=self.llm.generate, kwargs=generation_kwargs)
394
+ thread.start()
395
+ return streamer
396
+
397
+ def _decode_text(self, result_ids, tokenizer):
398
+ terminators = [tokenizer.convert_tokens_to_ids(i) for i in self.terminators]
399
+ result_text = []
400
+ for result in result_ids:
401
+ result = result[result != 0]
402
+ if result[0] == tokenizer.bos_id:
403
+ result = result[1:]
404
+ if result[-1] in terminators:
405
+ result = result[:-1]
406
+ result_text.append(tokenizer.decode(result).strip())
407
+ return result_text
408
+
409
+ def generate(
410
+ self,
411
+ input_ids=None,
412
+ pixel_values=None,
413
+ tgt_sizes=None,
414
+ image_bound=None,
415
+ gene_input_ids=None,
416
+ gene_attention_mask=None,
417
+ gene_bound=None,
418
+ attention_mask=None,
419
+ tokenizer=None,
420
+ vision_hidden_states=None,
421
+ return_vision_hidden_states=False,
422
+ stream=False,
423
+ decode_text=False,
424
+ **kwargs
425
+ ):
426
+ assert input_ids is not None
427
+ if pixel_values is not None:
428
+ assert len(input_ids) == len(pixel_values)
429
+ if gene_input_ids is not None:
430
+ assert len(input_ids) == len(gene_input_ids)
431
+
432
+ model_inputs = {
433
+ "input_ids": input_ids,
434
+ "image_bound": image_bound,
435
+ "gene_input_ids": gene_input_ids,
436
+ "gene_attention_mask": gene_attention_mask,
437
+ "gene_bound": gene_bound,
438
+ }
439
+
440
+ if vision_hidden_states is None:
441
+ model_inputs["pixel_values"] = pixel_values
442
+ model_inputs['tgt_sizes'] = tgt_sizes
443
+ else:
444
+ model_inputs["vision_hidden_states"] = vision_hidden_states
445
+
446
+ with torch.inference_mode():
447
+ model_inputs["inputs_embeds"], vision_hidden_states = self.get_vllm_embedding(model_inputs)
448
+
449
+ if stream:
450
+ result = self._decode_stream(model_inputs["inputs_embeds"], tokenizer, **kwargs)
451
+ else:
452
+ result = self._decode(
453
+ model_inputs["inputs_embeds"],
454
+ tokenizer,
455
+ attention_mask,
456
+ decode_text=decode_text,
457
+ **kwargs
458
+ )
459
+
460
+ if return_vision_hidden_states:
461
+ return result, vision_hidden_states
462
+ return result
463
+
464
+ def chat(
465
+ self,
466
+ msgs,
467
+ tokenizer,
468
+ image=None,
469
+ gene_sequence=None,
470
+ processor=None,
471
+ vision_hidden_states=None,
472
+ max_new_tokens=2048,
473
+ min_new_tokens=0,
474
+ sampling=True,
475
+ max_inp_length=12000,
476
+ system_prompt='',
477
+ stream=False,
478
+ max_slice_nums=None,
479
+ use_image_id=None,
480
+ **kwargs
481
+ ):
482
+ if isinstance(msgs[0], list):
483
+ batched = True
484
+ else:
485
+ batched = False
486
+
487
+ msgs_list = msgs
488
+ images_list = image
489
+ gene_sequences_list = gene_sequence
490
+
491
+ if batched is False:
492
+ images_list, msgs_list = [images_list], [msgs_list]
493
+ gene_sequences_list = [gene_sequences_list]
494
+ else:
495
+ assert images_list is None, "Please integrate image to msgs when using batch inference."
496
+ images_list = [None] * len(msgs_list)
497
+
498
+ assert len(images_list) == len(msgs_list), "The batch dim of images_list and msgs_list should be the same."
499
+
500
+ if processor is None:
501
+ if self.processor is None:
502
+ self.processor = AutoProcessor.from_pretrained(self.config._name_or_path, trust_remote_code=True)
503
+ processor = self.processor
504
+
505
+ assert self.config.query_num == processor.image_processor.image_feature_size
506
+ assert self.config.patch_size == processor.image_processor.patch_size
507
+ assert self.config.use_image_id == processor.image_processor.use_image_id
508
+ assert self.config.slice_config.max_slice_nums == processor.image_processor.max_slice_nums
509
+ assert self.config.slice_mode == processor.image_processor.slice_mode
510
+
511
+ prompts_lists = []
512
+ input_images_lists = []
513
+ input_gene_sequences_lists = []
514
+
515
+ for image, gene_seq, msgs in zip(images_list, gene_sequences_list, msgs_list):
516
+ if isinstance(msgs, str):
517
+ msgs = json.loads(msgs)
518
+ copy_msgs = deepcopy(msgs)
519
+
520
+ assert len(msgs) > 0, "msgs is empty"
521
+ assert sampling or not stream, "if use stream mode, make sure sampling=True"
522
+
523
+ content_raw = copy_msgs[0]["content"]
524
+ new_content = []
525
+ if image is not None:
526
+ new_content.append(image)
527
+ if gene_seq is not None:
528
+ new_content.append(gene_seq)
529
+ if isinstance(content_raw, str):
530
+ new_content.append(content_raw)
531
+ elif isinstance(content_raw, list):
532
+ new_content.extend(content_raw)
533
+ copy_msgs[0]["content"] = new_content
534
+
535
+ images_in_msg = []
536
+ gene_in_msg = []
537
+ for i, msg in enumerate(copy_msgs):
538
+ role = msg["role"]
539
+ content = msg["content"]
540
+ assert role in ["user", "assistant"]
541
+ if i == 0:
542
+ assert role == "user", "The role of first msg should be user"
543
+ if not isinstance(content, list):
544
+ content = [content]
545
+
546
+ cur_msgs = []
547
+ for c in content:
548
+ if isinstance(c, Image.Image):
549
+ images_in_msg.append(c)
550
+ cur_msgs.append("(<image>./</image>)")
551
+ elif isinstance(c, ad.AnnData):
552
+ gene_in_msg.append(c)
553
+ cur_msgs.append("(<gene>./</gene>)")
554
+ elif isinstance(c, str):
555
+ cur_msgs.append(c)
556
+ else:
557
+ raise TypeError(f"Unsupported content type: {type(c)}")
558
+
559
+ msg["content"] = "\n".join(cur_msgs)
560
+
561
+ if system_prompt:
562
+ sys_msg = {'role': 'system', 'content': system_prompt}
563
+ copy_msgs = [sys_msg] + copy_msgs
564
+
565
+ prompts_lists.append(
566
+ processor.tokenizer.apply_chat_template(copy_msgs, tokenize=False, add_generation_prompt=True)
567
+ )
568
+ input_images_lists.append(images_in_msg)
569
+ input_gene_sequences_lists.append(gene_in_msg)
570
+
571
+ inputs = processor(
572
+ prompts_lists,
573
+ input_images_lists,
574
+ input_gene_sequences_lists,
575
+ max_slice_nums=max_slice_nums,
576
+ use_image_id=use_image_id,
577
+ return_tensors="pt",
578
+ max_length=max_inp_length
579
+ ).to(self.device)
580
+
581
+ if sampling:
582
+ generation_config = {
583
+ "top_p": 0.8,
584
+ "top_k": 100,
585
+ "temperature": 0.7,
586
+ "do_sample": True,
587
+ "repetition_penalty": 1.05
588
+ }
589
+ else:
590
+ generation_config = {
591
+ "num_beams": 3,
592
+ "repetition_penalty": 1.2,
593
+ }
594
+
595
+ if min_new_tokens > 0:
596
+ generation_config['min_new_tokens'] = min_new_tokens
597
+
598
+ generation_config.update((k, kwargs[k]) for k in generation_config.keys() & kwargs.keys())
599
+
600
+ inputs.pop("image_sizes")
601
+
602
+ with torch.inference_mode():
603
+ res = self.generate(
604
+ **inputs,
605
+ tokenizer=tokenizer,
606
+ max_new_tokens=max_new_tokens,
607
+ vision_hidden_states=vision_hidden_states,
608
+ stream=stream,
609
+ decode_text=True,
610
+ **generation_config
611
+ )
612
+
613
+ if stream:
614
+ def stream_gen():
615
+ for text in res:
616
+ for term in self.terminators:
617
+ text = text.replace(term, '')
618
+ yield text
619
+ return stream_gen()
620
+ else:
621
+ if batched:
622
+ answer = res
623
+ else:
624
+ answer = res[0]
625
+ return answer