MCPcopy Create free account
hub / github.com/EmbodiedGPT/EmbodiedGPT_Pytorch / answer

Method answer

demo/script.py:367–407  ·  view source on GitHub ↗
(self, conversations, language_model_inputs, modal_type="image")

Source from the content-addressed store, hash-verified

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)

Callers 2

ask_questionMethod · 0.95
post_questionMethod · 0.95

Calls 2

batch_decodeMethod · 0.80
generateMethod · 0.45

Tested by

no test coverage detected