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
« 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
5r"""SimpleX
6################################################
8Reference:
9 Kelong Mao et al. "SimpleX: A Simple and Strong Baseline for Collaborative Filtering." in CIKM 2021.
11Reference code:
12 https://github.com/xue-pai/TwoToweRS
13"""
15import torch
16import torch.nn.functional as F
17from torch import nn
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
25class SimpleX(GeneralRecommender):
26 r"""SimpleX is a simple, unified collaborative filtering model.
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.
34 We implement the model following the original author with a pairwise training mode.
35 """
37 input_type = InputType.PAIRWISE
39 def __init__(self, config, dataset):
40 super().__init__(config, dataset)
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)
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)
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()
77 # parameters initialization
78 self.apply(xavier_normal_initialization)
79 # get the mask
80 self.item_emb.weight.data[0, :] = 0
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
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]
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
120 def get_cos(self, user_e, item_e):
121 r"""Get the cosine similarity between user and item
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]
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)
138 def forward(self, user, pos_item, history_item, history_len, neg_item_seq):
139 r"""Get the loss
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]
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)
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)
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)
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()
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 )
182 loss = CCL_loss + self.reg_weight * reg_loss.sum()
183 return loss
185 def calculate_loss(self, interaction):
186 r"""Data processing and call function forward(), return loss
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]
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]
210 loss = self.forward(user, pos_item, history_item, history_len, neg_item_seq)
211 return loss
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]
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)
226 # [user_num, embedding_size]
227 UI_aggregation_e = self.get_UI_aggregation(user_e, history_item_e, history_len)
229 UI_cos = self.get_cos(UI_aggregation_e, test_item_e.unsqueeze(1))
230 return UI_cos.squeeze(1)
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]
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)
242 # [user_num, embedding_size]
243 UI_aggregation_e = self.get_UI_aggregation(user_e, history_item_e, history_len)
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