| 247 | |
| 248 | class PromptLearner(nn.Module): |
| 249 | def __init__(self, image_feats=None, text_feats=None, ui_graph=None): |
| 250 | super().__init__() |
| 251 | self.ui_graph = ui_graph |
| 252 | |
| 253 | |
| 254 | if args.hard_token_type=='pca': |
| 255 | try: |
| 256 | t1 = time() |
| 257 | hard_token_image = pickle.load(open(args.data_path + args.dataset + '/hard_token_image_pca','rb')) |
| 258 | hard_token_text = pickle.load(open(args.data_path + args.dataset + '/hard_token_text_pca','rb')) |
| 259 | print('already load hard token', time() - t1) |
| 260 | except Exception: |
| 261 | hard_token_image = PCA(n_components=args.embed_size).fit_transform(image_feats) |
| 262 | hard_token_text = PCA(n_components=args.embed_size).fit_transform(text_feats) |
| 263 | pickle.dump(hard_token_image, open(args.data_path + args.dataset + '/hard_token_image_pca','wb')) |
| 264 | pickle.dump(hard_token_text, open(args.data_path + args.dataset + '/hard_token_text_pca','wb')) |
| 265 | elif args.hard_token_type=='ica': |
| 266 | try: |
| 267 | t1 = time() |
| 268 | hard_token_image = pickle.load(open(args.data_path + args.dataset + '/hard_token_image_ica','rb')) |
| 269 | hard_token_text = pickle.load(open(args.data_path + args.dataset + '/hard_token_text_ica','rb')) |
| 270 | print('already load hard token', time() - t1) |
| 271 | except Exception: |
| 272 | hard_token_image = FastICA(n_components=args.embed_size, random_state=12).fit_transform(image_feats) |
| 273 | hard_token_text = FastICA(n_components=args.embed_size, random_state=12).fit_transform(text_feats) |
| 274 | pickle.dump(hard_token_image, open(args.data_path + args.dataset + '/hard_token_image_ica','wb')) |
| 275 | pickle.dump(hard_token_text, open(args.data_path + args.dataset + '/hard_token_text_ica','wb')) |
| 276 | elif args.hard_token_type=='isomap': |
| 277 | hard_token_image = manifold.Isomap(n_neighbors=5, n_components=args.embed_size, n_jobs=-1).fit_transform(image_feats) |
| 278 | hard_token_text = manifold.Isomap(n_neighbors=5, n_components=args.embed_size, n_jobs=-1).fit_transform(text_feats) |
| 279 | # elif args.hard_token_type=='tsne': |
| 280 | # hard_token_image = TSNE(n_components=args.embed_size, n_iter=300).fit_transform(image_feats) |
| 281 | # hard_token_text = TSNE(n_components=args.embed_size, n_iter=300).fit_transform(text_feats) |
| 282 | # elif args.hard_token_type=='lda': |
| 283 | # hard_token_image = LinearDiscriminantAnalysis(n_components=args.embed_size).fit_transform(image_feats) |
| 284 | # hard_token_text = LinearDiscriminantAnalysis(n_components=args.embed_size).fit_transform(text_feats) |
| 285 | |
| 286 | # self.item_hard_token = nn.Embedding.from_pretrained(torch.mean((torch.stack((torch.tensor(hard_token_image).float(), torch.tensor(hard_token_text).float()))), dim=0), freeze=False).cuda().weight |
| 287 | # self.user_hard_token = nn.Embedding.from_pretrained(torch.mm(ui_graph, self.item_hard_token), freeze=False).cuda().weight |
| 288 | |
| 289 | self.item_hard_token = torch.mean((torch.stack((torch.tensor(hard_token_image).float(), torch.tensor(hard_token_text).float()))), dim=0).cuda() |
| 290 | self.user_hard_token = torch.mm(ui_graph, self.item_hard_token).cuda() |
| 291 | |
| 292 | self.trans_user = nn.Linear(args.embed_size, args.embed_size).cuda() |
| 293 | self.trans_item = nn.Linear(args.embed_size, args.embed_size).cuda() |
| 294 | # nn.init.xavier_uniform_(self.gnn_trans_user.weight) |
| 295 | # nn.init.xavier_uniform_(self.gnn_trans_item.weight) |
| 296 | # self.gnn_trans_user = self.gnn_trans_user.cuda() |
| 297 | # self.gnn_trans_item = self.gnn_trans_item.cuda() |
| 298 | # self.item_hard_token = torch.mean((torch.stack((torch.tensor(hard_token_image).float(), torch.tensor(hard_token_text).float()))), dim=0).cuda() |
| 299 | |
| 300 | |
| 301 | def forward(self): |