(self, conversations, language_model_inputs, modal_type="image")
| 365 | |
| 366 | @torch.no_grad() |
| 367 | def answer(self, conversations, language_model_inputs, modal_type="image"): |
| 368 | model_inputs = self.tokenizer( |
| 369 | conversations, |
| 370 | return_tensors="pt", |
| 371 | ) |
| 372 | model_inputs.pop("token_type_ids", None) |
| 373 | |
| 374 | input_ids = model_inputs["input_ids"].to(self.device) |
| 375 | attention_mask = model_inputs["attention_mask"].to(self.device) |
| 376 | |
| 377 | if modal_type == "text": |
| 378 | generation_output = self.model.language_model.generate( |
| 379 | input_ids=input_ids, |
| 380 | attention_mask=attention_mask, |
| 381 | generation_config=self.generation_config, |
| 382 | return_dict_in_generate=True, |
| 383 | output_scores=True |
| 384 | ) |
| 385 | else: |
| 386 | pixel_values = model_inputs.pop("pixel_values", None) |
| 387 | if pixel_values is not None: |
| 388 | pixel_values = pixel_values.to(self.device) |
| 389 | |
| 390 | generation_output = self.model.generate( |
| 391 | pixel_values=pixel_values, |
| 392 | input_ids=input_ids, |
| 393 | attention_mask=attention_mask, |
| 394 | language_model_inputs=language_model_inputs, |
| 395 | generation_config=self.generation_config, |
| 396 | return_dict_in_generate=True, |
| 397 | output_scores=True |
| 398 | ) |
| 399 | |
| 400 | preds = generation_output.sequences |
| 401 | outputs = self.tokenizer.batch_decode(preds, skip_special_tokens=True)[0] |
| 402 | |
| 403 | if modal_type == "text": |
| 404 | skip_echo_len = len(conversations[0]) - conversations[0].count("</s>") * 3 |
| 405 | outputs = outputs[skip_echo_len:].strip() |
| 406 | |
| 407 | return outputs |
| 408 | |
| 409 | def merge_box(self,dict1,dict2,dict3): |
| 410 | combined_dict = defaultdict(list) |
no test coverage detected