Coverage for hopwise/model/knowledge_aware_recommender/kgcn.py: 95%

145 statements  

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

1# @Time : 2020/10/6 

2# @Author : Changxin Tian 

3# @Email : cx.tian@outlook.com 

4 

5r"""KGCN 

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

7 

8Reference: 

9 Hongwei Wang et al. "Knowledge graph convolution networks for recommender systems." in WWW 2019. 

10 

11Reference code: 

12 https://github.com/hwwang55/KGCN 

13""" 

14 

15import numpy as np 

16import torch 

17from torch import nn 

18 

19from hopwise.model.abstract_recommender import KnowledgeRecommender 

20from hopwise.model.init import xavier_normal_initialization 

21from hopwise.model.loss import EmbLoss 

22from hopwise.utils import InputType 

23 

24 

25class KGCN(KnowledgeRecommender): 

26 r"""KGCN is a knowledge-based recommendation model that captures inter-item relatedness effectively by mining their 

27 associated attributes on the KG. To automatically discover both high-order structure information and semantic 

28 information of the KG, we treat KG as an undirected graph and sample from the neighbors for each entity in the KG 

29 as their receptive field, then combine neighborhood information with bias when calculating the representation of a 

30 given entity. 

31 """ 

32 

33 input_type = InputType.PAIRWISE 

34 

35 def __init__(self, config, dataset): 

36 super().__init__(config, dataset) 

37 

38 # load parameters info 

39 self.embedding_size = config["embedding_size"] 

40 # number of iterations when computing entity representation 

41 self.n_iter = config["n_iter"] 

42 self.aggregator_class = config["aggregator"] # which aggregator to use 

43 self.reg_weight = config["reg_weight"] # weight of l2 regularization 

44 self.neighbor_sample_size = config["neighbor_sample_size"] 

45 

46 # define embedding 

47 self.user_embedding = nn.Embedding(self.n_users, self.embedding_size) 

48 self.entity_embedding = nn.Embedding(self.n_entities, self.embedding_size) 

49 self.relation_embedding = nn.Embedding(self.n_relations + 1, self.embedding_size) 

50 

51 # sample neighbors 

52 kg_graph = dataset.kg_graph(form="coo", value_field="relation_id") 

53 adj_entity, adj_relation = self.construct_adj(kg_graph) 

54 self.adj_entity, self.adj_relation = ( 

55 adj_entity.to(self.device), 

56 adj_relation.to(self.device), 

57 ) 

58 

59 # define function 

60 self.softmax = nn.Softmax(dim=-1) 

61 self.linear_layers = torch.nn.ModuleList() 

62 for i in range(self.n_iter): 

63 self.linear_layers.append( 

64 nn.Linear( 

65 (self.embedding_size if not self.aggregator_class == "concat" else self.embedding_size * 2), 

66 self.embedding_size, 

67 ) 

68 ) 

69 self.ReLU = nn.ReLU() 

70 self.Tanh = nn.Tanh() 

71 

72 self.bce_loss = nn.BCEWithLogitsLoss() 

73 self.l2_loss = EmbLoss() 

74 

75 # parameters initialization 

76 self.apply(xavier_normal_initialization) 

77 self.other_parameter_name = ["adj_entity", "adj_relation"] 

78 

79 def construct_adj(self, kg_graph): 

80 r"""Get neighbors and corresponding relations for each entity in the KG. 

81 

82 Args: 

83 kg_graph(scipy.sparse.coo_matrix): an undirected graph 

84 

85 Returns: 

86 tuple: 

87 - adj_entity(torch.LongTensor): each line stores the sampled neighbor entities for a given entity, 

88 shape: [n_entities, neighbor_sample_size] 

89 - adj_relation(torch.LongTensor): each line stores the corresponding sampled neighbor relations, 

90 shape: [n_entities, neighbor_sample_size] 

91 """ 

92 # self.logger.info('constructing knowledge graph ...') 

93 # treat the KG as an undirected graph 

94 kg_dict = dict() 

95 for triple in zip(kg_graph.row, kg_graph.data, kg_graph.col): 

96 head = triple[0] 

97 relation = triple[1] 

98 tail = triple[2] 

99 if head not in kg_dict: 

100 kg_dict[head] = [] 

101 kg_dict[head].append((tail, relation)) 

102 if tail not in kg_dict: 

103 kg_dict[tail] = [] 

104 kg_dict[tail].append((head, relation)) 

105 

106 # self.logger.info('constructing adjacency matrix ...') 

107 # each line of adj_entity stores the sampled neighbor entities for a given entity 

108 # each line of adj_relation stores the corresponding sampled neighbor relations 

109 entity_num = kg_graph.shape[0] 

110 adj_entity = np.zeros([entity_num, self.neighbor_sample_size], dtype=np.int64) 

111 adj_relation = np.zeros([entity_num, self.neighbor_sample_size], dtype=np.int64) 

112 for entity in range(entity_num): 

113 if entity not in kg_dict.keys(): 

114 adj_entity[entity] = np.array([entity] * self.neighbor_sample_size) 

115 adj_relation[entity] = np.array([0] * self.neighbor_sample_size) 

116 continue 

117 

118 neighbors = kg_dict[entity] 

119 n_neighbors = len(neighbors) 

120 if n_neighbors >= self.neighbor_sample_size: 

121 sampled_indices = np.random.choice( 

122 list(range(n_neighbors)), 

123 size=self.neighbor_sample_size, 

124 replace=False, 

125 ) 

126 else: 

127 sampled_indices = np.random.choice( 

128 list(range(n_neighbors)), 

129 size=self.neighbor_sample_size, 

130 replace=True, 

131 ) 

132 adj_entity[entity] = np.array([neighbors[i][0] for i in sampled_indices]) 

133 adj_relation[entity] = np.array([neighbors[i][1] for i in sampled_indices]) 

134 

135 return torch.from_numpy(adj_entity), torch.from_numpy(adj_relation) 

136 

137 def get_neighbors(self, items): 

138 r"""Get neighbors and corresponding relations for each entity in items from adj_entity and adj_relation. 

139 

140 Args: 

141 items(torch.LongTensor): The input tensor that contains item's id, shape: [batch_size, ] 

142 

143 Returns: 

144 tuple: 

145 - entities(list): Entities is a list of i-iter (i = 0, 1, ..., n_iter) neighbors for the batch of items. 

146 dimensions of entities: {[batch_size, 1], 

147 [batch_size, n_neighbor], 

148 [batch_size, n_neighbor^2], 

149 ..., 

150 [batch_size, n_neighbor^n_iter]} 

151 - relations(list): Relations is a list of i-iter (i = 0, 1, ..., n_iter) corresponding relations for 

152 entities. Relations have the same shape as entities. 

153 """ # noqa: E501 

154 items = torch.unsqueeze(items, dim=1) 

155 entities = [items] 

156 relations = [] 

157 for i in range(self.n_iter): 

158 index = torch.flatten(entities[i]) 

159 neighbor_entities = torch.index_select(self.adj_entity, 0, index).reshape(self.batch_size, -1) 

160 neighbor_relations = torch.index_select(self.adj_relation, 0, index).reshape(self.batch_size, -1) 

161 entities.append(neighbor_entities) 

162 relations.append(neighbor_relations) 

163 return entities, relations 

164 

165 def mix_neighbor_vectors(self, neighbor_vectors, neighbor_relations, user_embeddings): 

166 r"""Mix neighbor vectors on user-specific graph. 

167 

168 Args: 

169 neighbor_vectors(torch.FloatTensor): The embeddings of neighbor entities(items), 

170 shape: [batch_size, -1, neighbor_sample_size, embedding_size] 

171 neighbor_relations(torch.FloatTensor): The embeddings of neighbor relations, 

172 shape: [batch_size, -1, neighbor_sample_size, embedding_size] 

173 user_embeddings(torch.FloatTensor): The embeddings of users, shape: [batch_size, embedding_size] 

174 

175 Returns: 

176 neighbors_aggregated(torch.FloatTensor): The neighbors aggregated embeddings, 

177 shape: [batch_size, -1, embedding_size] 

178 

179 """ 

180 avg = False 

181 if not avg: 

182 user_embeddings = user_embeddings.reshape( 

183 self.batch_size, 1, 1, self.embedding_size 

184 ) # [batch_size, 1, 1, dim] 

185 user_relation_scores = torch.mean( 

186 user_embeddings * neighbor_relations, dim=-1 

187 ) # [batch_size, -1, n_neighbor] 

188 user_relation_scores_normalized = self.softmax(user_relation_scores) # [batch_size, -1, n_neighbor] 

189 

190 user_relation_scores_normalized = torch.unsqueeze( 

191 user_relation_scores_normalized, dim=-1 

192 ) # [batch_size, -1, n_neighbor, 1] 

193 neighbors_aggregated = torch.mean( 

194 user_relation_scores_normalized * neighbor_vectors, dim=2 

195 ) # [batch_size, -1, dim] 

196 else: 

197 neighbors_aggregated = torch.mean(neighbor_vectors, dim=2) # [batch_size, -1, dim] 

198 return neighbors_aggregated 

199 

200 def aggregate(self, user_embeddings, entities, relations): 

201 r"""For each item, aggregate the entity representation and its neighborhood representation into a single vector. 

202 

203 Args: 

204 user_embeddings(torch.FloatTensor): The embeddings of users, shape: [batch_size, embedding_size] 

205 entities(list): entities is a list of i-iter (i = 0, 1, ..., n_iter) neighbors for the batch of items. 

206 dimensions of entities: {[batch_size, 1], 

207 [batch_size, n_neighbor], 

208 [batch_size, n_neighbor^2], 

209 ..., 

210 [batch_size, n_neighbor^n_iter]} 

211 relations(list): relations is a list of i-iter (i = 0, 1, ..., n_iter) corresponding relations for entities. 

212 relations have the same shape as entities. 

213 

214 Returns: 

215 item_embeddings(torch.FloatTensor): The embeddings of items, shape: [batch_size, embedding_size] 

216 

217 """ # noqa: E501 

218 entity_vectors = [self.entity_embedding(i) for i in entities] 

219 relation_vectors = [self.relation_embedding(i) for i in relations] 

220 

221 for i in range(self.n_iter): 

222 entity_vectors_next_iter = [] 

223 for hop in range(self.n_iter - i): 

224 shape = ( 

225 self.batch_size, 

226 -1, 

227 self.neighbor_sample_size, 

228 self.embedding_size, 

229 ) 

230 self_vectors = entity_vectors[hop] 

231 neighbor_vectors = entity_vectors[hop + 1].reshape(shape) 

232 neighbor_relations = relation_vectors[hop].reshape(shape) 

233 

234 neighbors_agg = self.mix_neighbor_vectors( 

235 neighbor_vectors, neighbor_relations, user_embeddings 

236 ) # [batch_size, -1, dim] 

237 

238 if self.aggregator_class == "sum": 

239 output = (self_vectors + neighbors_agg).reshape(-1, self.embedding_size) # [-1, dim] 

240 elif self.aggregator_class == "neighbor": 

241 output = neighbors_agg.reshape(-1, self.embedding_size) # [-1, dim] 

242 elif self.aggregator_class == "concat": 

243 # [batch_size, -1, dim * 2] 

244 output = torch.cat([self_vectors, neighbors_agg], dim=-1) 

245 output = output.reshape(-1, self.embedding_size * 2) # [-1, dim * 2] 

246 else: 

247 raise Exception("Unknown aggregator: " + self.aggregator_class) 

248 

249 output = self.linear_layers[i](output) 

250 # [batch_size, -1, dim] 

251 output = output.reshape(self.batch_size, -1, self.embedding_size) 

252 

253 if i == self.n_iter - 1: 

254 vector = self.Tanh(output) 

255 else: 

256 vector = self.ReLU(output) 

257 

258 entity_vectors_next_iter.append(vector) 

259 entity_vectors = entity_vectors_next_iter 

260 

261 item_embeddings = entity_vectors[0].reshape(self.batch_size, self.embedding_size) 

262 

263 return item_embeddings 

264 

265 def forward(self, user, item): 

266 self.batch_size = item.shape[0] 

267 # [batch_size, dim] 

268 user_e = self.user_embedding(user) 

269 # entities is a list of i-iter (i = 0, 1, ..., n_iter) neighbors for the batch of items. dimensions of entities: # noqa: E501 

270 # {[batch_size, 1], [batch_size, n_neighbor], [batch_size, n_neighbor^2], ..., [batch_size, n_neighbor^n_iter]} 

271 entities, relations = self.get_neighbors(item) 

272 # [batch_size, dim] 

273 item_e = self.aggregate(user_e, entities, relations) 

274 

275 return user_e, item_e 

276 

277 def calculate_loss(self, interaction): 

278 user = interaction[self.USER_ID] 

279 pos_item = interaction[self.ITEM_ID] 

280 neg_item = interaction[self.NEG_ITEM_ID] 

281 

282 user_e, pos_item_e = self.forward(user, pos_item) 

283 user_e, neg_item_e = self.forward(user, neg_item) 

284 

285 pos_item_score = torch.mul(user_e, pos_item_e).sum(dim=1) 

286 neg_item_score = torch.mul(user_e, neg_item_e).sum(dim=1) 

287 

288 predict = torch.cat((pos_item_score, neg_item_score)) 

289 target = torch.zeros(len(user) * 2, dtype=torch.float32).to(self.device) 

290 target[: len(user)] = 1 

291 rec_loss = self.bce_loss(predict, target) 

292 

293 l2_loss = self.l2_loss(user_e, pos_item_e, neg_item_e) 

294 loss = rec_loss + self.reg_weight * l2_loss 

295 

296 return loss 

297 

298 def predict(self, interaction): 

299 user = interaction[self.USER_ID] 

300 item = interaction[self.ITEM_ID] 

301 user_e, item_e = self.forward(user, item) 

302 return torch.mul(user_e, item_e).sum(dim=1) 

303 

304 def full_sort_predict(self, interaction): 

305 user_index = interaction[self.USER_ID] 

306 item_index = torch.tensor(range(self.n_items)).to(self.device) 

307 

308 user = torch.unsqueeze(user_index, dim=1).repeat(1, item_index.shape[0]) 

309 user = torch.flatten(user) 

310 item = torch.unsqueeze(item_index, dim=0).repeat(user_index.shape[0], 1) 

311 item = torch.flatten(item) 

312 

313 user_e, item_e = self.forward(user, item) 

314 score = torch.mul(user_e, item_e).sum(dim=1) 

315 

316 return score.view(-1)