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

Class Student_MLP

codes/Models.py:509–569  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

507
508
509class Student_MLP(nn.Module):
510 def __init__(self):
511 super(Student_MLP, self).__init__()
512 # self.n_users = n_users
513 # self.n_items = n_items
514 # self.embedding_dim = embedding_dim
515
516 # self.uEmbeds = nn.Parameter(init(torch.empty(args.user, args.latdim)))
517 # self.iEmbeds = nn.Parameter(init(torch.empty(args.item, args.latdim)))
518
519 self.user_trans = nn.Linear(args.embed_size, args.embed_size)
520 self.item_trans = nn.Linear(args.embed_size, args.embed_size)
521 nn.init.xavier_uniform_(self.user_trans.weight)
522 nn.init.xavier_uniform_(self.item_trans.weight)
523
524 self.MLP = BLMLP()
525 # self.overallTime = datetime.timedelta(0)
526
527
528 def get_embedding(self):
529 return self.user_id_embedding, self.item_id_embedding
530
531
532 def forward(self, pre_user, pre_item, ):
533 # pre_user, pre_item = self.user_id_embedding.weight, self.item_id_embedding.weight
534 user_embed = self.user_trans(pre_user)
535 item_embed = self.user_trans(pre_item)
536
537 return user_embed, item_embed
538 # return pre_user, pre_item
539
540 def init_user_item_embed(self, pre_u_embed, pre_i_embed):
541 self.user_id_embedding = nn.Embedding.from_pretrained(pre_u_embed, freeze=False)
542 self.item_id_embedding = nn.Embedding.from_pretrained(pre_i_embed, freeze=False)
543
544 def pointPosPredictwEmbeds(self, uEmbeds, iEmbeds, ancs, poss):
545 ancEmbeds = uEmbeds[ancs]
546 posEmbeds = iEmbeds[poss]
547 nume = self.MLP.pairPred(ancEmbeds, posEmbeds)
548 return nume
549
550 def pointNegPredictwEmbeds(self, embeds1, embeds2, nodes1, temp=1.0):
551 pckEmbeds1 = embeds1[nodes1]
552 preds = self.MLP.crossPred(pckEmbeds1, embeds2)
553 return torch.exp(preds / temp).sum(-1)
554
555 def pairPredictwEmbeds(self, uEmbeds, iEmbeds, ancs, poss, negs):
556 ancEmbeds = uEmbeds[ancs]
557 posEmbeds = iEmbeds[poss]
558 negEmbeds = iEmbeds[negs]
559 posPreds = self.MLP.pairPred(ancEmbeds, posEmbeds)
560 negPreds = self.MLP.pairPred(ancEmbeds, negEmbeds)
561 return posPreds - negPreds
562
563 def predAll(self, pckUEmbeds, iEmbeds):
564 return self.MLP.crossPred(pckUEmbeds, iEmbeds)
565
566 def testPred(self, usr, trnMask):

Callers 1

trainMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected