Coverage for hopwise/model/general_recommender/bpr.py: 100%
43 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
1# @Time : 2020/6/25
2# @Author : Shanlei Mu
3# @Email : slmu@ruc.edu.cn
5# UPDATE:
6# @Time : 2020/9/16
7# @Author : Shanlei Mu
8# @Email : slmu@ruc.edu.cn
10r"""BPR
11################################################
12Reference:
13 Steffen Rendle et al. "BPR: Bayesian Personalized Ranking from Implicit Feedback." in UAI 2009.
14"""
16import torch
17from torch import nn
19from hopwise.model.abstract_recommender import GeneralRecommender
20from hopwise.model.init import xavier_normal_initialization
21from hopwise.model.loss import BPRLoss
22from hopwise.utils import InputType
25class BPR(GeneralRecommender):
26 r"""BPR is a basic matrix factorization model that be trained in the pairwise way."""
28 input_type = InputType.PAIRWISE
30 def __init__(self, config, dataset):
31 super().__init__(config, dataset)
33 # load parameters info
34 self.embedding_size = config["embedding_size"]
36 # define layers and loss
37 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size)
38 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size)
39 self.loss = BPRLoss()
41 # parameters initialization
42 self.apply(xavier_normal_initialization)
44 def get_user_embedding(self, user):
45 r"""Get a batch of user embedding tensor according to input user's id.
47 Args:
48 user (torch.LongTensor): The input tensor that contains user's id, shape: [batch_size, ]
50 Returns:
51 torch.FloatTensor: The embedding tensor of a batch of user, shape: [batch_size, embedding_size]
52 """
53 return self.user_embedding(user)
55 def get_item_embedding(self, item):
56 r"""Get a batch of item embedding tensor according to input item's id.
58 Args:
59 item (torch.LongTensor): The input tensor that contains item's id, shape: [batch_size, ]
61 Returns:
62 torch.FloatTensor: The embedding tensor of a batch of item, shape: [batch_size, embedding_size]
63 """
64 return self.item_embedding(item)
66 def forward(self, user, item):
67 user_e = self.get_user_embedding(user)
68 item_e = self.get_item_embedding(item)
69 return user_e, item_e
71 def calculate_loss(self, interaction):
72 user = interaction[self.USER_ID]
73 pos_item = interaction[self.ITEM_ID]
74 neg_item = interaction[self.NEG_ITEM_ID]
76 user_e, pos_e = self.forward(user, pos_item)
77 neg_e = self.get_item_embedding(neg_item)
78 pos_item_score, neg_item_score = (
79 torch.mul(user_e, pos_e).sum(dim=1),
80 torch.mul(user_e, neg_e).sum(dim=1),
81 )
82 loss = self.loss(pos_item_score, neg_item_score)
83 return loss
85 def predict(self, interaction):
86 user = interaction[self.USER_ID]
87 item = interaction[self.ITEM_ID]
88 user_e, item_e = self.forward(user, item)
89 return torch.mul(user_e, item_e).sum(dim=1)
91 def full_sort_predict(self, interaction):
92 user = interaction[self.USER_ID]
93 user_e = self.get_user_embedding(user)
94 all_item_e = self.item_embedding.weight
95 score = torch.matmul(user_e, all_item_e.transpose(0, 1))
96 return score.view(-1)