| 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: |