MCPcopy Create free account
hub / github.com/alipay/Ant-Multi-Modal-Framework / forward

Method forward

prj/M2_Encoder/ms_wrapper.py:32–59  ·  view source on GitHub ↗
(self, forward_params)

Source from the content-addressed store, hash-verified

30 self._tokenizer, self._img_processor = processors
31
32 def forward(self, forward_params):
33 rets = {}
34 if "text" in forward_params:
35 text = forward_params.get("text")
36 txt_encoding = self._tokenizer(
37 text,
38 padding="max_length",
39 truncation=True,
40 max_length=self.model.hparams.config["max_text_len"],
41 return_special_tokens_mask=True,
42 )
43 txt_data = {
44 "text_ids": torch.tensor(txt_encoding["input_ids"]).to(self._device),
45 "text_masks": torch.tensor(txt_encoding["attention_mask"]).to(
46 self._device
47 ),
48 "text_labels": None,
49 }
50 txt_feats = self.model.infer_text(txt_data)["cls_vlffn_feats"]
51 rets.update({"text_embedding": txt_feats.detach()})
52 if "img" in forward_params:
53 input_img = forward_params["img"]
54 img = self._img_processor(input_img).unsqueeze(0)
55 img_data = {"image": [img.to(self._device)]}
56 img_feats = self.model.infer_image(img_data)["cls_vlffn_feats"]
57 rets.update({"img_embedding": img_feats.detach()})
58
59 return rets
60
61
62@PREPROCESSORS.register_module(

Callers 1

forwardMethod · 0.45

Calls 5

infer_textMethod · 0.80
infer_imageMethod · 0.80
getMethod · 0.45
toMethod · 0.45
updateMethod · 0.45

Tested by

no test coverage detected