| 507 | |
| 508 | |
| 509 | class 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): |