Coverage for hopwise/model/sequential_recommender/fossil.py: 83%

95 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-30 13:25 +0000

1# @Time : 2020/11/21 20:00 

2# @Author : Shao Weiqi 

3# @Reviewer : Lin Kun 

4# @Email : shaoweiqi@ruc.edu.cn 

5 

6r"""FOSSIL 

7################################################ 

8 

9Reference: 

10 Ruining He et al. "Fusing Similarity Models with Markov Chains for Sparse Sequential Recommendation." in ICDM 2016. 

11 

12 

13""" 

14 

15import torch 

16from torch import nn 

17from torch.nn.init import xavier_normal_ 

18 

19from hopwise.model.abstract_recommender import SequentialRecommender 

20from hopwise.model.loss import BPRLoss 

21 

22 

23class FOSSIL(SequentialRecommender): 

24 r"""FOSSIL uses similarity of the items as main purpose and uses high MC as a way of sequential preference improve of 

25 ability of sequential recommendation 

26 

27 """ # noqa: E501 

28 

29 def __init__(self, config, dataset): 

30 super().__init__(config, dataset) 

31 

32 # load the dataset information 

33 self.n_users = dataset.num(self.USER_ID) 

34 self.device = config["device"] 

35 

36 # load the parameters 

37 self.embedding_size = config["embedding_size"] 

38 self.order_len = config["order_len"] 

39 assert self.order_len <= self.max_seq_length, "order_len can't longer than the max_seq_length" 

40 self.reg_weight = config["reg_weight"] 

41 self.alpha = config["alpha"] 

42 

43 # define the layers and loss type 

44 self.item_embedding = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0) 

45 self.user_lambda = nn.Embedding(self.n_users, self.order_len) 

46 self.lambda_ = nn.Parameter(torch.zeros(self.order_len)) 

47 

48 self.loss_type = config["loss_type"] 

49 if self.loss_type == "BPR": 

50 self.loss_fct = BPRLoss() 

51 elif self.loss_type == "CE": 

52 self.loss_fct = nn.CrossEntropyLoss() 

53 else: 

54 raise NotImplementedError("Make sure 'loss_type' in ['BPR', 'CE']!") 

55 

56 # init the parameters of the model 

57 self.apply(self.init_weights) 

58 

59 def inverse_seq_item_embedding(self, seq_item_embedding, seq_item_len): 

60 """Inverse seq_item_embedding like this (simple to 2-dim): 

61 

62 [1,2,3,0,0,0] -- ??? -- >> [0,0,0,1,2,3] 

63 

64 first: [0,0,0,0,0,0] concat [1,2,3,0,0,0] 

65 

66 using gather_indexes: to get one by one 

67 

68 first get 3,then 2,last 1 

69 """ 

70 zeros = torch.zeros_like(seq_item_embedding, dtype=torch.float).to(self.device) 

71 # batch_size * seq_len * embedding_size 

72 item_embedding_zeros = torch.cat([zeros, seq_item_embedding], dim=1) 

73 # batch_size * 2_mul_seq_len * embedding_size 

74 embedding_list = list() 

75 for i in range(self.order_len): 

76 embedding = self.gather_indexes( 

77 item_embedding_zeros, 

78 self.max_seq_length + seq_item_len - self.order_len + i, 

79 ) 

80 embedding_list.append(embedding.unsqueeze(1)) 

81 short_item_embedding = torch.cat(embedding_list, dim=1) 

82 # batch_size * short_len * embedding_size 

83 

84 return short_item_embedding 

85 

86 def reg_loss(self, user_embedding, item_embedding, seq_output): 

87 reg_1 = self.reg_weight 

88 loss_1 = ( 

89 reg_1 * torch.norm(user_embedding, p=2) 

90 + reg_1 * torch.norm(item_embedding, p=2) 

91 + reg_1 * torch.norm(seq_output, p=2) 

92 ) 

93 

94 return loss_1 

95 

96 def init_weights(self, module): 

97 if isinstance(module, nn.Embedding) or isinstance(module, nn.Linear): 

98 xavier_normal_(module.weight.data) 

99 

100 def forward(self, seq_item, seq_item_len, user): 

101 seq_item_embedding = self.item_embedding(seq_item) 

102 

103 high_order_seq_item_embedding = self.inverse_seq_item_embedding(seq_item_embedding, seq_item_len) 

104 # batch_size * order_len * embedding 

105 

106 high_order = self.get_high_order_Markov(high_order_seq_item_embedding, user) 

107 similarity = self.get_similarity(seq_item_embedding, seq_item_len) 

108 

109 return high_order + similarity 

110 

111 def get_high_order_Markov(self, high_order_item_embedding, user): 

112 """In order to get the inference of past items and the user's taste to the current predict item""" 

113 user_lambda = self.user_lambda(user).unsqueeze(dim=2) 

114 # batch_size * order_len * 1 

115 lambda_ = self.lambda_.unsqueeze(dim=0).unsqueeze(dim=2) 

116 # 1 * order_len * 1 

117 lambda_ = torch.add(user_lambda, lambda_) 

118 # batch_size * order_len * 1 

119 high_order_item_embedding = torch.mul(high_order_item_embedding, lambda_) 

120 # batch_size * order_len * embedding_size 

121 high_order_item_embedding = high_order_item_embedding.sum(dim=1) 

122 # batch_size * embedding_size 

123 

124 return high_order_item_embedding 

125 

126 def get_similarity(self, seq_item_embedding, seq_item_len): 

127 """In order to get the inference of past items to the current predict item""" 

128 coeff = torch.pow(seq_item_len.unsqueeze(1), -self.alpha).float() 

129 # batch_size * 1 

130 similarity = torch.mul(coeff, seq_item_embedding.sum(dim=1)) 

131 # batch_size * embedding_size 

132 

133 return similarity 

134 

135 def calculate_loss(self, interaction): 

136 seq_item = interaction[self.ITEM_SEQ] 

137 user = interaction[self.USER_ID] 

138 seq_item_len = interaction[self.ITEM_SEQ_LEN] 

139 seq_output = self.forward(seq_item, seq_item_len, user) 

140 pos_items = interaction[self.POS_ITEM_ID] 

141 pos_items_emb = self.item_embedding(pos_items) 

142 

143 user_lambda = self.user_lambda(user) 

144 pos_items_embedding = self.item_embedding(pos_items) 

145 if self.loss_type == "BPR": 

146 neg_items = interaction[self.NEG_ITEM_ID] 

147 neg_items_emb = self.item_embedding(neg_items) 

148 pos_score = torch.sum(seq_output * pos_items_emb, dim=-1) 

149 neg_score = torch.sum(seq_output * neg_items_emb, dim=-1) 

150 loss = self.loss_fct(pos_score, neg_score) 

151 return loss + self.reg_loss(user_lambda, pos_items_embedding, seq_output) 

152 else: # self.loss_type = 'CE' 

153 test_item_emb = self.item_embedding.weight 

154 logits = torch.matmul(seq_output, test_item_emb.transpose(0, 1)) 

155 loss = self.loss_fct(logits, pos_items) 

156 return loss + self.reg_loss(user_lambda, pos_items_embedding, seq_output) 

157 

158 def predict(self, interaction): 

159 item_seq = interaction[self.ITEM_SEQ] 

160 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

161 test_item = interaction[self.ITEM_ID] 

162 user = interaction[self.USER_ID] 

163 seq_output = self.forward(item_seq, item_seq_len, user) 

164 test_item_emb = self.item_embedding(test_item) 

165 scores = torch.mul(seq_output, test_item_emb).sum(dim=1) 

166 return scores 

167 

168 def full_sort_predict(self, interaction): 

169 item_seq = interaction[self.ITEM_SEQ] 

170 user = interaction[self.USER_ID] 

171 item_seq_len = interaction[self.ITEM_SEQ_LEN] 

172 seq_output = self.forward(item_seq, item_seq_len, user) 

173 test_items_emb = self.item_embedding.weight 

174 scores = torch.matmul(seq_output, test_items_emb.transpose(0, 1)) 

175 return scores