Coverage for hopwise/model/general_recommender/simplex.py: 90%

113 statements  

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

1# @Time : 2022/3/25 13:38 

2# @Author : HaoJun Qin 

3# @Email : 18697951462@163.com 

4 

5r"""SimpleX 

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

7 

8Reference: 

9 Kelong Mao et al. "SimpleX: A Simple and Strong Baseline for Collaborative Filtering." in CIKM 2021. 

10 

11Reference code: 

12 https://github.com/xue-pai/TwoToweRS 

13""" 

14 

15import torch 

16import torch.nn.functional as F 

17from torch import nn 

18 

19from hopwise.model.abstract_recommender import GeneralRecommender 

20from hopwise.model.init import xavier_normal_initialization 

21from hopwise.model.loss import EmbLoss 

22from hopwise.utils import InputType 

23 

24 

25class SimpleX(GeneralRecommender): 

26 r"""SimpleX is a simple, unified collaborative filtering model. 

27 

28 SimpleX presents a simple and easy-to-understand model. Its advantage lies 

29 in its loss function, which uses a larger number of negative samples and 

30 sets a threshold to filter out less informative samples, it also uses 

31 relative weights to control the balance of positive-sample loss 

32 and negative-sample loss. 

33 

34 We implement the model following the original author with a pairwise training mode. 

35 """ 

36 

37 input_type = InputType.PAIRWISE 

38 

39 def __init__(self, config, dataset): 

40 super().__init__(config, dataset) 

41 

42 # Get user history interacted items 

43 self.history_item_id, _, self.history_item_len = dataset.history_item_matrix( 

44 max_history_len=config["history_len"] 

45 ) 

46 self.history_item_id = self.history_item_id.to(self.device) 

47 self.history_item_len = self.history_item_len.to(self.device) 

48 

49 # load parameters info 

50 self.embedding_size = config["embedding_size"] 

51 self.margin = config["margin"] 

52 self.negative_weight = config["negative_weight"] 

53 self.gamma = config["gamma"] 

54 self.neg_seq_len = config["train_neg_sample_args"]["sample_num"] 

55 self.reg_weight = config["reg_weight"] 

56 self.aggregator = config["aggregator"] 

57 if self.aggregator not in ["mean", "user_attention", "self_attention"]: 

58 raise ValueError("aggregator must be mean, user_attention or self_attention") 

59 self.history_len = torch.max(self.history_item_len, dim=0) 

60 

61 # user embedding matrix 

62 self.user_emb = nn.Embedding(self.n_users, self.embedding_size) 

63 # item embedding matrix 

64 self.item_emb = nn.Embedding(self.n_items, self.embedding_size, padding_idx=0) 

65 # feature space mapping matrix of user and item 

66 self.UI_map = nn.Linear(self.embedding_size, self.embedding_size, bias=False) 

67 if self.aggregator in ["user_attention", "self_attention"]: 

68 self.W_k = nn.Sequential(nn.Linear(self.embedding_size, self.embedding_size), nn.Tanh()) 

69 if self.aggregator == "self_attention": 

70 self.W_q = nn.Linear(self.embedding_size, 1, bias=False) 

71 # dropout 

72 self.dropout = nn.Dropout(config["dropout_prob"]) 

73 self.require_pow = config["require_pow"] 

74 # l2 regularization loss 

75 self.reg_loss = EmbLoss() 

76 

77 # parameters initialization 

78 self.apply(xavier_normal_initialization) 

79 # get the mask 

80 self.item_emb.weight.data[0, :] = 0 

81 

82 def get_UI_aggregation(self, user_e, history_item_e, history_len): 

83 r"""Get the combined vector of user and historically interacted items 

84 

85 Args: 

86 user_e (torch.Tensor): User's feature vector, shape: [user_num, embedding_size] 

87 history_item_e (torch.Tensor): History item's feature vector, 

88 shape: [user_num, max_history_len, embedding_size] 

89 history_len (torch.Tensor): User's history length, shape: [user_num] 

90 

91 Returns: 

92 torch.Tensor: Combined vector of user and item sequences, shape: [user_num, embedding_size] 

93 """ 

94 if self.aggregator == "mean": 

95 pos_item_sum = history_item_e.sum(dim=1) 

96 # [user_num, embedding_size] 

97 out = pos_item_sum / (history_len + 1.0e-10).unsqueeze(1) 

98 elif self.aggregator in ["user_attention", "self_attention"]: 

99 # [user_num, max_history_len, embedding_size] 

100 key = self.W_k(history_item_e) 

101 if self.aggregator == "user_attention": 

102 # [user_num, max_history_len] 

103 attention = torch.matmul(key, user_e.unsqueeze(2)).squeeze(2) 

104 elif self.aggregator == "self_attention": 

105 # [user_num, max_history_len] 

106 attention = self.W_q(key).squeeze(2) 

107 e_attention = torch.exp(attention) 

108 mask = (history_item_e.sum(dim=-1) != 0).int() 

109 e_attention = e_attention * mask 

110 # [user_num, max_history_len] 

111 attention_weight = e_attention / (e_attention.sum(dim=1, keepdim=True) + 1.0e-10) 

112 # [user_num, embedding_size] 

113 out = torch.matmul(attention_weight.unsqueeze(1), history_item_e).squeeze(1) 

114 # Combined vector of user and item sequences 

115 out = self.UI_map(out) 

116 g = self.gamma 

117 UI_aggregation_e = g * user_e + (1 - g) * out 

118 return UI_aggregation_e 

119 

120 def get_cos(self, user_e, item_e): 

121 r"""Get the cosine similarity between user and item 

122 

123 Args: 

124 user_e (torch.Tensor): User's feature vector, shape: [user_num, embedding_size] 

125 item_e (torch.Tensor): Item's feature vector, 

126 shape: [user_num, item_num, embedding_size] 

127 

128 Returns: 

129 torch.Tensor: Cosine similarity between user and item, shape: [user_num, item_num] 

130 """ 

131 user_e = F.normalize(user_e, dim=1) 

132 # [user_num, embedding_size, 1] 

133 user_e = user_e.unsqueeze(2) 

134 item_e = F.normalize(item_e, dim=2) 

135 UI_cos = torch.matmul(item_e, user_e) 

136 return UI_cos.squeeze(2) 

137 

138 def forward(self, user, pos_item, history_item, history_len, neg_item_seq): 

139 r"""Get the loss 

140 

141 Args: 

142 user (torch.Tensor): User's id, shape: [user_num] 

143 pos_item (torch.Tensor): Positive item's id, shape: [user_num] 

144 history_item (torch.Tensor): Id of historty item, shape: [user_num, max_history_len] 

145 history_len (torch.Tensor): History item's length, shape: [user_num] 

146 neg_item_seq (torch.Tensor): Negative item seq's id, shape: [user_num, neg_seq_len] 

147 

148 Returns: 

149 torch.Tensor: Loss, shape: [] 

150 """ 

151 # [user_num, embedding_size] 

152 user_e = self.user_emb(user) 

153 # [user_num, embedding_size] 

154 pos_item_e = self.item_emb(pos_item) 

155 # [user_num, max_history_len, embedding_size] 

156 history_item_e = self.item_emb(history_item) 

157 # [nuser_num, neg_seq_len, embedding_size] 

158 neg_item_seq_e = self.item_emb(neg_item_seq) 

159 

160 # [user_num, embedding_size] 

161 UI_aggregation_e = self.get_UI_aggregation(user_e, history_item_e, history_len) 

162 UI_aggregation_e = self.dropout(UI_aggregation_e) 

163 

164 pos_cos = self.get_cos(UI_aggregation_e, pos_item_e.unsqueeze(1)) 

165 neg_cos = self.get_cos(UI_aggregation_e, neg_item_seq_e) 

166 

167 # CCL loss 

168 pos_loss = torch.relu(1 - pos_cos) 

169 neg_loss = torch.relu(neg_cos - self.margin) 

170 neg_loss = neg_loss.mean(1, keepdim=True) * self.negative_weight 

171 CCL_loss = (pos_loss + neg_loss).mean() 

172 

173 # l2 regularization loss 

174 reg_loss = self.reg_loss( 

175 user_e, 

176 pos_item_e, 

177 history_item_e, 

178 neg_item_seq_e, 

179 require_pow=self.require_pow, 

180 ) 

181 

182 loss = CCL_loss + self.reg_weight * reg_loss.sum() 

183 return loss 

184 

185 def calculate_loss(self, interaction): 

186 r"""Data processing and call function forward(), return loss 

187 

188 To use SimpleX, a user must have a historical transaction record, 

189 a pos item and a sequence of neg items. Based on the hopwise 

190 framework, the data in the interaction object is ordered, so 

191 we can get the data quickly. 

192 """ 

193 user = interaction[self.USER_ID] 

194 pos_item = interaction[self.ITEM_ID] 

195 neg_item = interaction[self.NEG_ITEM_ID] 

196 

197 # get the sequence of neg items 

198 neg_item_seq = neg_item.reshape((self.neg_seq_len, -1)) 

199 neg_item_seq = neg_item_seq.T 

200 user_number = int(len(user) / self.neg_seq_len) 

201 # user's id 

202 user = user[0:user_number] 

203 # historical transaction record 

204 history_item = self.history_item_id[user] 

205 # positive item's id 

206 pos_item = pos_item[0:user_number] 

207 # history_len 

208 history_len = self.history_item_len[user] 

209 

210 loss = self.forward(user, pos_item, history_item, history_len, neg_item_seq) 

211 return loss 

212 

213 def predict(self, interaction): 

214 user = interaction[self.USER_ID] 

215 history_item = self.history_item_id[user] 

216 history_len = self.history_item_len[user] 

217 test_item = interaction[self.ITEM_ID] 

218 

219 # [user_num, embedding_size] 

220 user_e = self.user_emb(user) 

221 # [user_num, embedding_size] 

222 test_item_e = self.item_emb(test_item) 

223 # [user_num, max_history_len, embedding_size] 

224 history_item_e = self.item_emb(history_item) 

225 

226 # [user_num, embedding_size] 

227 UI_aggregation_e = self.get_UI_aggregation(user_e, history_item_e, history_len) 

228 

229 UI_cos = self.get_cos(UI_aggregation_e, test_item_e.unsqueeze(1)) 

230 return UI_cos.squeeze(1) 

231 

232 def full_sort_predict(self, interaction): 

233 user = interaction[self.USER_ID] 

234 history_item = self.history_item_id[user] 

235 history_len = self.history_item_len[user] 

236 

237 # [user_num, embedding_size] 

238 user_e = self.user_emb(user) 

239 # [user_num, max_history_len, embedding_size] 

240 history_item_e = self.item_emb(history_item) 

241 

242 # [user_num, embedding_size] 

243 UI_aggregation_e = self.get_UI_aggregation(user_e, history_item_e, history_len) 

244 

245 UI_aggregation_e = F.normalize(UI_aggregation_e, dim=1) 

246 all_item_emb = self.item_emb.weight 

247 all_item_emb = F.normalize(all_item_emb, dim=1) 

248 UI_cos = torch.matmul(UI_aggregation_e, all_item_emb.T) 

249 return UI_cos