Coverage for hopwise/model/general_recommender/sgl.py: 80%
153 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 : 2021/10/12
2# @Author : Tian Zhen
3# @Email : chenyuwuxinn@gmail.com
5r"""SGL
6################################################
7Reference:
8 Jiancan Wu et al. "SGL: Self-supervised Graph Learning for Recommendation" in SIGIR 2021.
10Reference code:
11 https://github.com/wujcan/SGL
12"""
14import numpy as np
15import scipy.sparse as sp
16import torch
17import torch.nn.functional as F
19from hopwise.model.abstract_recommender import GeneralRecommender
20from hopwise.model.init import xavier_uniform_initialization
21from hopwise.model.loss import EmbLoss
22from hopwise.utils import InputType
25class SGL(GeneralRecommender):
26 r"""SGL is a GCN-based recommender model.
28 SGL supplements the classical supervised task of recommendation with an auxiliary
29 self supervised task, which reinforces node representation learning via self-
30 discrimination.Specifically,SGL generates multiple views of a node, maximizing the
31 agreement between different views of the same node compared to that of other nodes.
32 SGL devises three operators to generate the views — node dropout, edge dropout, and
33 random walk — that change the graph structure in different manners.
35 We implement the model following the original author with a pairwise training mode.
36 """
38 input_type = InputType.PAIRWISE
40 def __init__(self, config, dataset):
41 super().__init__(config, dataset)
42 self._user = dataset.inter_feat[dataset.uid_field]
43 self._item = dataset.inter_feat[dataset.iid_field]
44 self.embed_dim = config["embedding_size"]
45 self.n_layers = int(config["n_layers"])
46 self.type = config["type"]
47 self.drop_ratio = config["drop_ratio"]
48 self.ssl_tau = config["ssl_tau"]
49 self.reg_weight = config["reg_weight"]
50 self.ssl_weight = config["ssl_weight"]
51 self.user_embedding = torch.nn.Embedding(self.n_users, self.embed_dim)
52 self.item_embedding = torch.nn.Embedding(self.n_items, self.embed_dim)
53 self.reg_loss = EmbLoss()
54 self.train_graph = self.csr2tensor(self.create_adjust_matrix(is_sub=False))
55 self.restore_user_e = None
56 self.restore_item_e = None
57 self.apply(xavier_uniform_initialization)
58 self.other_parameter_name = ["restore_user_e", "restore_item_e"]
60 def graph_construction(self):
61 r"""Devise three operators to generate the views — node dropout, edge dropout, and random walk of a node."""
62 self.sub_graph1 = []
63 if self.type in ("ND", "ED"):
64 self.sub_graph1 = self.csr2tensor(self.create_adjust_matrix(is_sub=True))
65 elif self.type == "RW":
66 for i in range(self.n_layers):
67 _g = self.csr2tensor(self.create_adjust_matrix(is_sub=True))
68 self.sub_graph1.append(_g)
70 self.sub_graph2 = []
71 if self.type in ("ND", "ED"):
72 self.sub_graph2 = self.csr2tensor(self.create_adjust_matrix(is_sub=True))
73 elif self.type == "RW":
74 for i in range(self.n_layers):
75 _g = self.csr2tensor(self.create_adjust_matrix(is_sub=True))
76 self.sub_graph2.append(_g)
78 def rand_sample(self, high, size=None, replace=True):
79 r"""Randomly discard some points or edges.
81 Args:
82 high (int): Upper limit of index value
83 size (int): Array size after sampling
85 Returns:
86 numpy.ndarray: Array index after sampling, shape: [size]
87 """
88 a = np.arange(high)
89 sample = np.random.choice(a, size=size, replace=replace)
90 return sample
92 def create_adjust_matrix(self, is_sub: bool):
93 r"""Get the normalized interaction matrix of users and items.
95 Construct the square matrix from the training data and normalize it
96 using the laplace matrix.If it is a subgraph, it may be processed by
97 node dropout or edge dropout.
99 .. math::
100 A_{hat} = D^{-0.5} \times A \times D^{-0.5}
102 Returns:
103 csr_matrix of the normalized interaction matrix.
104 """
105 matrix = None
106 if not is_sub:
107 ratings = np.ones_like(self._user, dtype=np.float32)
108 matrix = sp.csr_matrix(
109 (ratings, (self._user, self._item + self.n_users)),
110 shape=(self.n_users + self.n_items, self.n_users + self.n_items),
111 )
112 elif self.type == "ND":
113 drop_user = self.rand_sample(
114 self.n_users,
115 size=int(self.n_users * self.drop_ratio),
116 replace=False,
117 )
118 drop_item = self.rand_sample(
119 self.n_items,
120 size=int(self.n_items * self.drop_ratio),
121 replace=False,
122 )
123 R_user = np.ones(self.n_users, dtype=np.float32)
124 R_user[drop_user] = 0.0
125 R_item = np.ones(self.n_items, dtype=np.float32)
126 R_item[drop_item] = 0.0
127 R_user = sp.diags(R_user)
128 R_item = sp.diags(R_item)
129 R_G = sp.csr_matrix(
130 (
131 np.ones_like(self._user, dtype=np.float32),
132 (self._user, self._item),
133 ),
134 shape=(self.n_users, self.n_items),
135 )
136 res = R_user.dot(R_G)
137 res = res.dot(R_item)
139 user, item = res.nonzero()
140 ratings = res.data
141 matrix = sp.csr_matrix(
142 (ratings, (user, item + self.n_users)),
143 shape=(self.n_users + self.n_items, self.n_users + self.n_items),
144 )
146 elif self.type in ("ED", "RW"):
147 keep_item = self.rand_sample(
148 len(self._user),
149 size=int(len(self._user) * (1 - self.drop_ratio)),
150 replace=False,
151 )
152 user = self._user[keep_item]
153 item = self._item[keep_item]
155 matrix = sp.csr_matrix(
156 (np.ones_like(user), (user, item + self.n_users)),
157 shape=(self.n_users + self.n_items, self.n_users + self.n_items),
158 )
160 matrix = matrix + matrix.T
161 D = np.array(matrix.sum(axis=1)) + 1e-7
162 D = np.power(D, -0.5).flatten()
163 D = sp.diags(D)
164 return D.dot(matrix).dot(D)
166 def csr2tensor(self, matrix: sp.csr_matrix):
167 r"""Convert csr_matrix to tensor.
169 Args:
170 matrix (scipy.csr_matrix): Sparse matrix to be converted.
172 Returns:
173 torch.sparse.FloatTensor: Transformed sparse matrix.
174 """
175 matrix = matrix.tocoo()
176 x = torch.sparse.FloatTensor(
177 torch.LongTensor(np.array([matrix.row, matrix.col])),
178 torch.FloatTensor(matrix.data.astype(np.float32)),
179 matrix.shape,
180 ).to(self.device)
181 return x
183 def forward(self, graph):
184 main_ego = torch.cat([self.user_embedding.weight, self.item_embedding.weight])
185 all_ego = [main_ego]
186 if isinstance(graph, list):
187 for sub_graph in graph:
188 main_ego = torch.sparse.mm(sub_graph, main_ego)
189 all_ego.append(main_ego)
190 else:
191 for i in range(self.n_layers):
192 main_ego = torch.sparse.mm(graph, main_ego)
193 all_ego.append(main_ego)
194 all_ego = torch.stack(all_ego, dim=1)
195 all_ego = torch.mean(all_ego, dim=1, keepdim=False)
196 user_emd, item_emd = torch.split(all_ego, [self.n_users, self.n_items], dim=0)
198 return user_emd, item_emd
200 def calculate_loss(self, interaction):
201 if self.restore_user_e is not None or self.restore_item_e is not None:
202 self.restore_user_e, self.restore_item_e = None, None
204 user_list = interaction[self.USER_ID]
205 pos_item_list = interaction[self.ITEM_ID]
206 neg_item_list = interaction[self.NEG_ITEM_ID]
207 user_emd, item_emd = self.forward(self.train_graph)
208 user_sub1, item_sub1 = self.forward(self.sub_graph1)
209 user_sub2, item_sub2 = self.forward(self.sub_graph2)
210 total_loss = self.calc_bpr_loss(
211 user_emd, item_emd, user_list, pos_item_list, neg_item_list
212 ) + self.calc_ssl_loss(user_list, pos_item_list, user_sub1, user_sub2, item_sub1, item_sub2)
213 return total_loss
215 def calc_bpr_loss(self, user_emd, item_emd, user_list, pos_item_list, neg_item_list):
216 r"""Calculate the the pairwise Bayesian Personalized Ranking (BPR) loss and parameter regularization loss.
218 Args:
219 user_emd (torch.Tensor): Ego embedding of all users after forwarding.
220 item_emd (torch.Tensor): Ego embedding of all items after forwarding.
221 user_list (torch.Tensor): List of the user.
222 pos_item_list (torch.Tensor): List of positive examples.
223 neg_item_list (torch.Tensor): List of negative examples.
225 Returns:
226 torch.Tensor: Loss of BPR tasks and parameter regularization.
227 """
228 u_e = user_emd[user_list]
229 pi_e = item_emd[pos_item_list]
230 ni_e = item_emd[neg_item_list]
231 p_scores = torch.mul(u_e, pi_e).sum(dim=1)
232 n_scores = torch.mul(u_e, ni_e).sum(dim=1)
234 l1 = torch.sum(-F.logsigmoid(p_scores - n_scores))
236 u_e_p = self.user_embedding(user_list)
237 pi_e_p = self.item_embedding(pos_item_list)
238 ni_e_p = self.item_embedding(neg_item_list)
240 l2 = self.reg_loss(u_e_p, pi_e_p, ni_e_p)
242 return l1 + l2 * self.reg_weight
244 def calc_ssl_loss(self, user_list, pos_item_list, user_sub1, user_sub2, item_sub1, item_sub2):
245 r"""Calculate the loss of self-supervised tasks.
247 Args:
248 user_list (torch.Tensor): List of the user.
249 pos_item_list (torch.Tensor): List of positive examples.
250 user_sub1 (torch.Tensor): Ego embedding of all users in the first subgraph after forwarding.
251 user_sub2 (torch.Tensor): Ego embedding of all users in the second subgraph after forwarding.
252 item_sub1 (torch.Tensor): Ego embedding of all items in the first subgraph after forwarding.
253 item_sub2 (torch.Tensor): Ego embedding of all items in the second subgraph after forwarding.
255 Returns:
256 torch.Tensor: Loss of self-supervised tasks.
257 """
258 u_emd1 = F.normalize(user_sub1[user_list], dim=1)
259 u_emd2 = F.normalize(user_sub2[user_list], dim=1)
260 all_user2 = F.normalize(user_sub2, dim=1)
261 v1 = torch.sum(u_emd1 * u_emd2, dim=1)
262 v2 = u_emd1.matmul(all_user2.T)
263 v1 = torch.exp(v1 / self.ssl_tau)
264 v2 = torch.sum(torch.exp(v2 / self.ssl_tau), dim=1)
265 ssl_user = -torch.sum(torch.log(v1 / v2))
267 i_emd1 = F.normalize(item_sub1[pos_item_list], dim=1)
268 i_emd2 = F.normalize(item_sub2[pos_item_list], dim=1)
269 all_item2 = F.normalize(item_sub2, dim=1)
270 v3 = torch.sum(i_emd1 * i_emd2, dim=1)
271 v4 = i_emd1.matmul(all_item2.T)
272 v3 = torch.exp(v3 / self.ssl_tau)
273 v4 = torch.sum(torch.exp(v4 / self.ssl_tau), dim=1)
274 ssl_item = -torch.sum(torch.log(v3 / v4))
276 return (ssl_item + ssl_user) * self.ssl_weight
278 def predict(self, interaction):
279 if self.restore_user_e is None or self.restore_item_e is None:
280 self.restore_user_e, self.restore_item_e = self.forward(self.train_graph)
282 user = self.restore_user_e[interaction[self.USER_ID]]
283 item = self.restore_item_e[interaction[self.ITEM_ID]]
284 return torch.sum(user * item, dim=1)
286 def full_sort_predict(self, interaction):
287 if self.restore_user_e is None or self.restore_item_e is None:
288 self.restore_user_e, self.restore_item_e = self.forward(self.train_graph)
290 user = self.restore_user_e[interaction[self.USER_ID]]
291 return user.matmul(self.restore_item_e.T)
293 def train(self, mode: bool = True):
294 r"""Override train method of base class.The subgraph is reconstructed each time it is called."""
295 T = super().train(mode=mode)
296 if mode:
297 self.graph_construction()
298 return T