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

1# @Time : 2020/9/14 17:01 

2# @Author : Hui Wang 

3# @Email : hui.wang@ruc.edu.cn 

4 

5r"""TransRec 

6################################################ 

7 

8Reference: 

9 Ruining He et al. "Translation-based Recommendation." In RecSys 2017. 

10 

11""" 

12 

13import torch 

14from torch import nn 

15 

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 

20 

21 

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 """ 

27 

28 input_type = InputType.PAIRWISE 

29 

30 def __init__(self, config, dataset): 

31 super().__init__(config, dataset) 

32 

33 # load parameters info 

34 self.embedding_size = config["embedding_size"] 

35 

36 # load dataset info 

37 self.n_users = dataset.user_num 

38 

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' 

43 

44 self.bpr_loss = BPRLoss() 

45 self.emb_loss = EmbLoss() 

46 self.reg_loss = RegLoss() 

47 

48 # parameters initialization 

49 self.apply(xavier_normal_initialization) 

50 

51 def _l2_distance(self, x, y): 

52 return torch.sqrt(torch.sum((x - y) ** 2, dim=-1, keepdim=True)) # [B 1] 

53 

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] 

59 

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 

68 

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] 

73 

74 seq_output = self.forward(user, item_seq, item_seq_len) # [B H] 

75 

76 pos_items = interaction[self.POS_ITEM_ID] # [B] 

77 neg_items = interaction[self.NEG_ITEM_ID] # [B] sample 1 negative item 

78 

79 pos_items_emb = self.item_embedding(pos_items) # [B H] 

80 neg_items_emb = self.item_embedding(neg_items) 

81 

82 pos_bias = self.bias(pos_items) # [B 1] 

83 neg_bias = self.bias(neg_items) 

84 

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) 

87 

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()) 

92 

93 reg_loss = self.reg_loss(self.T) 

94 return bpr_loss + item_emb_loss + user_emb_loss + bias_emb_loss + reg_loss 

95 

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] 

101 

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] 

105 

106 scores = test_bias - self._l2_distance(seq_output, test_item_emb) # [B 1] 

107 scores = scores.squeeze(-1) # [B] 

108 return scores 

109 

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] 

114 

115 seq_output = self.forward(user, item_seq, item_seq_len) # [B H] 

116 

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] 

119 

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] 

123 

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