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

Method answer

demo/inference.py:319–359  ·  view source on GitHub ↗
(self, conversations, language_model_inputs, modal_type="image")

Source from the content-addressed store, hash-verified

317
318 @torch.no_grad()
319 def answer(self, conversations, language_model_inputs, modal_type="image"):
320 model_inputs = self.tokenizer(
321 conversations,
322 return_tensors="pt",
323 )
324 model_inputs.pop("token_type_ids", None)
325
326 input_ids = model_inputs["input_ids"].to(self.device)
327 attention_mask = model_inputs["attention_mask"].to(self.device)
328
329 if modal_type == "text":
330 generation_output = self.model.language_model.generate(
331 input_ids=input_ids,
332 attention_mask=attention_mask,
333 generation_config=self.generation_config,
334 return_dict_in_generate=True,
335 output_scores=True
336 )
337 else:
338 pixel_values = model_inputs.pop("pixel_values", None)
339 if pixel_values is not None:
340 pixel_values = pixel_values.to(self.device)
341
342 generation_output = self.model.generate(
343 pixel_values=pixel_values,
344 input_ids=input_ids,
345 attention_mask=attention_mask,
346 language_model_inputs=language_model_inputs,
347 generation_config=self.generation_config,
348 return_dict_in_generate=True,
349 output_scores=True
350 )
351
352 preds = generation_output.sequences
353 outputs = self.tokenizer.batch_decode(preds, skip_special_tokens=True)[0]
354
355 if modal_type == "text":
356 skip_echo_len = len(conversations[0]) - conversations[0].count("</s>") * 3
357 outputs = outputs[skip_echo_len:].strip()
358
359 return outputs
360
361if __name__ == '__main__':
362 # model_path = "/mnt/petrelfs/zhangqinglong/Documents/Husky/work_dirs/husky_v3/EmbodiedGPT/pretrain_0727"

Callers 1

inference.pyFile · 0.45

Calls 2

batch_decodeMethod · 0.80
generateMethod · 0.45

Tested by

no test coverage detected