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

1# @Time : 2021/10/12 

2# @Author : Tian Zhen 

3# @Email : chenyuwuxinn@gmail.com 

4 

5r"""SGL 

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

7Reference: 

8 Jiancan Wu et al. "SGL: Self-supervised Graph Learning for Recommendation" in SIGIR 2021. 

9 

10Reference code: 

11 https://github.com/wujcan/SGL 

12""" 

13 

14import numpy as np 

15import scipy.sparse as sp 

16import torch 

17import torch.nn.functional as F 

18 

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 

23 

24 

25class SGL(GeneralRecommender): 

26 r"""SGL is a GCN-based recommender model. 

27 

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. 

34 

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

36 """ 

37 

38 input_type = InputType.PAIRWISE 

39 

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

59 

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) 

69 

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) 

77 

78 def rand_sample(self, high, size=None, replace=True): 

79 r"""Randomly discard some points or edges. 

80 

81 Args: 

82 high (int): Upper limit of index value 

83 size (int): Array size after sampling 

84 

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 

91 

92 def create_adjust_matrix(self, is_sub: bool): 

93 r"""Get the normalized interaction matrix of users and items. 

94 

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. 

98 

99 .. math:: 

100 A_{hat} = D^{-0.5} \times A \times D^{-0.5} 

101 

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) 

138 

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 ) 

145 

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] 

154 

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 ) 

159 

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) 

165 

166 def csr2tensor(self, matrix: sp.csr_matrix): 

167 r"""Convert csr_matrix to tensor. 

168 

169 Args: 

170 matrix (scipy.csr_matrix): Sparse matrix to be converted. 

171 

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 

182 

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) 

197 

198 return user_emd, item_emd 

199 

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 

203 

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 

214 

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. 

217 

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. 

224 

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) 

233 

234 l1 = torch.sum(-F.logsigmoid(p_scores - n_scores)) 

235 

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) 

239 

240 l2 = self.reg_loss(u_e_p, pi_e_p, ni_e_p) 

241 

242 return l1 + l2 * self.reg_weight 

243 

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. 

246 

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. 

254 

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

266 

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

275 

276 return (ssl_item + ssl_user) * self.ssl_weight 

277 

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) 

281 

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) 

285 

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) 

289 

290 user = self.restore_user_e[interaction[self.USER_ID]] 

291 return user.matmul(self.restore_item_e.T) 

292 

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