Coverage for hopwise/model/sequential_recommender/transrec.py: 87%
76 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/9/14 17:01
2# @Author : Hui Wang
3# @Email : hui.wang@ruc.edu.cn
5r"""TransRec
6################################################
8Reference:
9 Ruining He et al. "Translation-based Recommendation." In RecSys 2017.
11"""
13import torch
14from torch import nn
16from hopwise.model.abstract_recommender import SequentialRecommender
17from hopwise.model.init import xavier_normal_initialization
18from hopwise.model.loss import BPRLoss, EmbLoss, RegLoss
19from hopwise.utils import InputType
22class TransRec(SequentialRecommender):
23 r"""TransRec is translation-based model for sequential recommendation.
24 It assumes that the `prev. item` + `user` = `next item`.
25 We use the Euclidean Distance to calculate the similarity in this implementation.
26 """
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 # load dataset info
37 self.n_users = dataset.user_num
39 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size, padding_idx=0)
40 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0)
41 self.bias = nn.Embedding(self.n_items, 1, padding_idx=0) # Beta popularity bias
42 self.T = nn.Parameter(torch.zeros(self.embedding_size)) # average user representation 'global'
44 self.bpr_loss = BPRLoss()
45 self.emb_loss = EmbLoss()
46 self.reg_loss = RegLoss()
48 # parameters initialization
49 self.apply(xavier_normal_initialization)
51 def _l2_distance(self, x, y):
52 return torch.sqrt(torch.sum((x - y) ** 2, dim=-1, keepdim=True)) # [B 1]
54 def gather_last_items(self, item_seq, gather_index):
55 """Gathers the last_item at the specific positions over a minibatch"""
56 gather_index = gather_index.view(-1, 1)
57 last_items = item_seq.gather(index=gather_index, dim=1) # [B 1]
58 return last_items.squeeze(-1) # [B]
60 def forward(self, user, item_seq, item_seq_len):
61 # the last item at the last position
62 last_items = self.gather_last_items(item_seq, item_seq_len - 1) # [B]
63 user_emb = self.user_embedding(user) # [B H]
64 last_items_emb = self.item_embedding(last_items) # [B H]
65 T = self.T.expand_as(user_emb) # [B H]
66 seq_output = user_emb + T + last_items_emb # [B H]
67 return seq_output
69 def calculate_loss(self, interaction):
70 user = interaction[self.USER_ID] # [B]
71 item_seq = interaction[self.ITEM_SEQ] # [B Len]
72 item_seq_len = interaction[self.ITEM_SEQ_LEN]
74 seq_output = self.forward(user, item_seq, item_seq_len) # [B H]
76 pos_items = interaction[self.POS_ITEM_ID] # [B]
77 neg_items = interaction[self.NEG_ITEM_ID] # [B] sample 1 negative item
79 pos_items_emb = self.item_embedding(pos_items) # [B H]
80 neg_items_emb = self.item_embedding(neg_items)
82 pos_bias = self.bias(pos_items) # [B 1]
83 neg_bias = self.bias(neg_items)
85 pos_score = pos_bias - self._l2_distance(seq_output, pos_items_emb)
86 neg_score = neg_bias - self._l2_distance(seq_output, neg_items_emb)
88 bpr_loss = self.bpr_loss(pos_score, neg_score)
89 item_emb_loss = self.emb_loss(self.item_embedding(pos_items).detach())
90 user_emb_loss = self.emb_loss(self.user_embedding(user).detach())
91 bias_emb_loss = self.emb_loss(self.bias(pos_items).detach())
93 reg_loss = self.reg_loss(self.T)
94 return bpr_loss + item_emb_loss + user_emb_loss + bias_emb_loss + reg_loss
96 def predict(self, interaction):
97 user = interaction[self.USER_ID] # [B]
98 item_seq = interaction[self.ITEM_SEQ] # [B Len]
99 item_seq_len = interaction[self.ITEM_SEQ_LEN]
100 test_item = interaction[self.ITEM_ID]
102 seq_output = self.forward(user, item_seq, item_seq_len) # [B H]
103 test_item_emb = self.item_embedding(test_item) # [B H]
104 test_bias = self.bias(test_item) # [B 1]
106 scores = test_bias - self._l2_distance(seq_output, test_item_emb) # [B 1]
107 scores = scores.squeeze(-1) # [B]
108 return scores
110 def full_sort_predict(self, interaction):
111 user = interaction[self.USER_ID] # [B]
112 item_seq = interaction[self.ITEM_SEQ] # [B Len]
113 item_seq_len = interaction[self.ITEM_SEQ_LEN]
115 seq_output = self.forward(user, item_seq, item_seq_len) # [B H]
117 test_items_emb = self.item_embedding.weight # [item_num H]
118 test_items_emb = test_items_emb.repeat(seq_output.size(0), 1, 1) # [user_num item_num H]
120 user_hidden = seq_output.unsqueeze(1).expand_as(test_items_emb) # [user_num item_num H]
121 test_bias = self.bias.weight # [item_num 1]
122 test_bias = test_bias.repeat(user_hidden.size(0), 1, 1) # [user_num item_num 1]
124 scores = test_bias - self._l2_distance(user_hidden, test_items_emb) # [user_num item_num 1]
125 scores = scores.squeeze(-1) # [B n_items]
126 return scores