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

Method __init__

codes/Models.py:249–298  ·  view source on GitHub ↗
(self, image_feats=None, text_feats=None, ui_graph=None)

Source from the content-addressed store, hash-verified

247
248class 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):

Callers

nothing calls this directly

Calls 2

__init__Method · 0.45
mmMethod · 0.45

Tested by

no test coverage detected