MCPcopy Create free account
hub / github.com/HKUDS/PromptMM / forward

Method forward

codes/Models.py:169–240  ·  view source on GitHub ↗
(self, ui_graph, iu_graph, prompt_module=None)

Source from the content-addressed store, hash-verified

167 # self.prompt_item = soft_token_i
168
169 def forward(self, ui_graph, iu_graph, prompt_module=None):
170
171 # def forward(self, ui_graph, iu_graph):
172
173 prompt_user, prompt_item = prompt_module() # [n*32]
174 # ----feature prompt----
175 # feat_prompt_user = torch.mean( torch.stack((torch.mm(prompt_user, torch.mm(prompt_user.T, self.image_feats)), torch.mm(prompt_user, torch.mm(prompt_user.T, self.text_feats)))), dim=0 )
176 # feat_prompt_user = torch.mm(prompt_user, torch.mm(prompt_user.T, self.text_feats))
177 feat_prompt_item_image = torch.mm(prompt_item, torch.mm(prompt_item.T, self.image_feats))
178 feat_prompt_item_text = torch.mm(prompt_item, torch.mm(prompt_item.T, self.text_feats))
179 # feat_prompt_image_item = torch.mm(prompt_item, torch.mm(prompt_item.T, self.image_feats))
180 # feat_prompt_text_item = torch.mm(prompt_item, torch.mm(prompt_item.T, self.text_feats))
181 # ----feature prompt----
182 # image_feats = image_item_feats = self.dropout(self.image_trans(self.image_feats + feat_prompt_item_image ))
183 # text_feats = text_item_feats = self.dropout(self.text_trans(self.text_feats + feat_prompt_item_text ))
184
185 # image_feats = image_item_feats = self.dropout(self.image_trans(self.image_feats + F.normalize(feat_prompt_item_image, p=2, dim=1) ))
186 # text_feats = text_item_feats = self.dropout(self.text_trans(self.text_feats + F.normalize(feat_prompt_item_text, p=2, dim=1) ))
187
188 image_feats = image_item_feats = self.dropout(self.image_trans(self.image_feats + args.feat_soft_token_rate*F.normalize(feat_prompt_item_image, p=2, dim=1) ))
189 text_feats = text_item_feats = self.dropout(self.text_trans(self.text_feats + args.feat_soft_token_rate*F.normalize(feat_prompt_item_text, p=2, dim=1) ))
190 # args.feat_soft_token_rate*F.normalize(feat_prompt_item_image, p=2, dim=1)
191 # args.feat_soft_token_rate*F.normalize(feat_prompt_item_text, p=2, dim=1)
192
193
194 # image_feats = image_item_feats = self.dropout(self.image_trans(self.image_feats))
195 # text_feats = text_item_feats = self.dropout(self.text_trans(self.text_feats))
196
197 for i in range(args.layers):
198 image_user_feats = self.mm(ui_graph, image_feats)
199 image_item_feats = self.mm(iu_graph, image_user_feats)
200 # image_user_id = self.mm(image_ui_graph, self.item_id_embedding.weight)
201 # image_item_id = self.mm(image_iu_graph, self.user_id_embedding.weight)
202
203 text_user_feats = self.mm(ui_graph, text_feats)
204 text_item_feats = self.mm(iu_graph, text_user_feats)
205
206 # text_user_id = self.mm(text_ui_graph, self.item_id_embedding.weight)
207 # text_item_id = self.mm(text_iu_graph, self.user_id_embedding.weight)
208
209 # self.embedding_dict['user']['image'] = image_user_id
210 # self.embedding_dict['user']['text'] = text_user_id
211 # self.embedding_dict['item']['image'] = image_item_id
212 # self.embedding_dict['item']['text'] = text_item_id
213 # user_z, att_u = self.multi_head_self_attention(self.weight_dict, self.embedding_dict['user'], self.embedding_dict['user'])
214 # item_z, att_i = self.multi_head_self_attention(self.weight_dict, self.embedding_dict['item'], self.embedding_dict['item'])
215 # user_emb = user_z.mean(0)
216 # item_emb = item_z.mean(0)
217 u_g_embeddings = self.user_id_embedding.weight + args.soft_token_rate*F.normalize(prompt_user, p=2, dim=1)
218 i_g_embeddings = self.item_id_embedding.weight + args.soft_token_rate*F.normalize(prompt_item, p=2, dim=1)
219 user_emb_list = [u_g_embeddings]
220 item_emb_list = [i_g_embeddings]
221 for i in range(self.n_ui_layers):
222 if i == (self.n_ui_layers-1):
223 u_g_embeddings = self.softmax( torch.mm(ui_graph, i_g_embeddings) )
224 i_g_embeddings = self.softmax( torch.mm(iu_graph, u_g_embeddings) )
225
226 else:

Callers

nothing calls this directly

Calls 1

mmMethod · 0.95

Tested by

no test coverage detected