Coverage for hopwise/data/dataset/kg_dataset.py: 77%

725 statements  

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

1# @Time : 2020/9/3 

2# @Author : Yupeng Hou 

3# @Email : houyupeng@ruc.edu.cn 

4 

5# UPDATE: 

6# @Time : 2020/10/16, 2020/9/15, 2020/10/25, 2022/7/10 

7# @Author : Yupeng Hou, Xingyu Pan, Yushuo Chen, Lanling Xu 

8# @Email : houyupeng@ruc.edu.cn, panxy@ruc.edu.cn, chenyushuo@ruc.edu.cn, xulanling_sherry@163.com 

9 

10# UPDATE: 

11# @Time : 2025 

12# @Author : Giacomo Medda, Alessandro Soccol 

13# @Email : giacomo.medda@unica.it, alessandro.soccol@unica.it 

14 

15"""hopwise.data.kg_dataset 

16 hopwise.data.user_item_kg_dataset 

17########################## 

18""" 

19 

20import copy 

21import os 

22import sys 

23from collections import Counter 

24 

25import numpy as np 

26import pandas as pd 

27import torch 

28from scipy.sparse import coo_matrix 

29 

30from hopwise.data.dataset import Dataset 

31from hopwise.data.interaction import Interaction 

32from hopwise.utils import FeatureSource, FeatureType, KnowledgeEvaluationType, set_color 

33from hopwise.utils.url import decide_download, download_url, extract_zip 

34 

35 

36class KnowledgeBasedDataset(Dataset): 

37 """:class:`KnowledgeBasedDataset` is based on :class:`~hopwise.data.dataset.dataset.Dataset`, 

38 and load ``.kg`` and ``.link`` additionally. 

39 

40 Entities are remapped together with ``item_id`` specially. 

41 All entities are remapped into three consecutive ID sections. 

42 

43 - virtual entities that only exist in interaction data. 

44 - entities that exist both in interaction data and kg triplets. 

45 - entities only exist in kg triplets. 

46 

47 It also provides several interfaces to transfer ``.kg`` features into coo sparse matrix, 

48 csr sparse matrix or :class:`PyG.Data`. 

49 

50 Attributes: 

51 head_entity_field (str): The same as ``config['HEAD_ENTITY_ID_FIELD']``. 

52 

53 tail_entity_field (str): The same as ``config['TAIL_ENTITY_ID_FIELD']``. 

54 

55 relation_field (str): The same as ``config['RELATION_ID_FIELD']``. 

56 

57 entity_field (str): The same as ``config['ENTITY_ID_FIELD']``. 

58 

59 kg_feat (pandas.DataFrame): Internal data structure stores the kg triplets. 

60 It's loaded from file ``.kg``. 

61 

62 item2entity (dict): Dict maps ``item_id`` to ``entity``, 

63 which is loaded from file ``.link``. 

64 

65 entity2item (dict): Dict maps ``entity`` to ``item_id``, 

66 which is loaded from file ``.link``. 

67 

68 Note: 

69 :attr:`entity_field` doesn't exist exactly. It's only a symbol, 

70 representing entity features. 

71 

72 :attr:`ui_relation` is a special relation token, which is used to represent 

73 the interaction relation between users and items. 

74 """ 

75 

76 def __init__(self, config): 

77 super().__init__(config) 

78 

79 def _get_field_from_config(self): 

80 super()._get_field_from_config() 

81 self.head_entity_field = self.config["HEAD_ENTITY_ID_FIELD"] 

82 self.tail_entity_field = self.config["TAIL_ENTITY_ID_FIELD"] 

83 self.relation_field = self.config["RELATION_ID_FIELD"] 

84 self.entity_field = self.config["ENTITY_ID_FIELD"] 

85 self.kg_reverse_r = self.config["kg_reverse_r"] 

86 self.ui_relation = self.config["ui_relation"] 

87 self.entity_kg_num_interval = self.config["entity_kg_num_interval"] 

88 self.relation_kg_num_interval = self.config["relation_kg_num_interval"] 

89 self._check_field("head_entity_field", "tail_entity_field", "relation_field", "entity_field") 

90 self.set_field_property(self.entity_field, FeatureType.TOKEN, FeatureSource.KG, 1) 

91 

92 self.logger.debug(set_color("relation_field", "blue") + f": {self.relation_field}") 

93 self.logger.debug(set_color("entity_field", "blue") + f": {self.entity_field}") 

94 

95 def _data_filtering(self): 

96 super()._data_filtering() 

97 self._filter_kg_by_triple_num() 

98 self._filter_link() 

99 

100 def _filter_kg_by_triple_num(self): 

101 """Filter by number of triples. 

102 

103 The interval of the number of triples can be set, and only entities/relations 

104 whose number of triples is in the specified interval can be retained. 

105 See :doc:`../user_guide/data/data_args` for detail arg setting. 

106 

107 Note: 

108 Lower bound of the interval is also called k-core filtering, which means this method 

109 will filter loops until all the entities and relations has at least k triples. 

110 """ 

111 entity_kg_num_interval = self._parse_intervals_str(self.config["entity_kg_num_interval"]) 

112 relation_kg_num_interval = self._parse_intervals_str(self.config["relation_kg_num_interval"]) 

113 

114 if entity_kg_num_interval is None and relation_kg_num_interval is None: 

115 return 

116 

117 entity_kg_num = Counter() 

118 if entity_kg_num_interval: 

119 head_entity_kg_num = Counter(self.kg_feat[self.head_entity_field].values) 

120 tail_entity_kg_num = Counter(self.kg_feat[self.tail_entity_field].values) 

121 

122 self.head_entity_kg_num = head_entity_kg_num 

123 entity_kg_num = head_entity_kg_num + tail_entity_kg_num 

124 self.entity_kg_num = entity_kg_num 

125 relation_kg_num = Counter(self.kg_feat[self.relation_field].values) if relation_kg_num_interval else Counter() 

126 

127 while True: 

128 ban_entities = self._get_illegal_ids_by_inter_num( 

129 field=f"{self.head_entity_field}-{self.tail_entity_field}", 

130 feat=None, 

131 inter_num=entity_kg_num, 

132 inter_interval=entity_kg_num_interval, 

133 ) 

134 ban_relations = self._get_illegal_ids_by_inter_num( 

135 field=self.relation_field, 

136 feat=None, 

137 inter_num=relation_kg_num, 

138 inter_interval=relation_kg_num_interval, 

139 ) 

140 if len(ban_entities) == 0 and len(ban_relations) == 0: 

141 break 

142 

143 dropped_kg = pd.Series(False, index=self.kg_feat.index) 

144 head_entity_kg = self.kg_feat[self.head_entity_field] 

145 tail_entity_kg = self.kg_feat[self.tail_entity_field] 

146 relation_kg = self.kg_feat[self.relation_field] 

147 dropped_kg |= head_entity_kg.isin(ban_entities) 

148 dropped_kg |= tail_entity_kg.isin(ban_entities) 

149 dropped_kg |= relation_kg.isin(ban_relations) 

150 

151 entity_kg_num -= Counter(head_entity_kg[dropped_kg].values) 

152 entity_kg_num -= Counter(tail_entity_kg[dropped_kg].values) 

153 relation_kg_num -= Counter(relation_kg[dropped_kg].values) 

154 

155 dropped_index = self.kg_feat.index[dropped_kg] 

156 self.logger.debug(f"[{len(dropped_index)}] dropped triples.") 

157 self.kg_feat.drop(dropped_index, inplace=True) 

158 

159 def build(self): 

160 """Processing dataset according to evaluation setting, including Group, Order and Split. 

161 See :class:`~hopwise.config.eval_setting.EvalSetting` for details. 

162 

163 Returns: 

164 list: List of built :class:`Dataset`. 

165 """ 

166 self._change_feat_format() 

167 

168 if self.benchmark_filename_list is not None: 

169 self._drop_unused_col() 

170 cumsum = list(np.cumsum(self.file_size_list)) 

171 datasets = [self.copy(self.inter_feat[start:end]) for start, end in zip([0] + cumsum[:-1], cumsum)] 

172 return datasets 

173 

174 # ordering 

175 ordering_args = self.config["eval_args"]["order"] 

176 if ordering_args == "RO": 

177 self.shuffle() 

178 elif ordering_args == "TO": 

179 self.sort(by=self.time_field) 

180 else: 

181 raise NotImplementedError("The ordering_method [{ordering_args}] has not been implemented.") 

182 

183 # splitting & grouping 

184 split_args = self.config["eval_args"]["split"] 

185 eval_lp_args = self.config["eval_lp_args"] 

186 

187 if eval_lp_args is not None and eval_lp_args["knowledge_split"] is not None: 

188 knowledge_split_args = eval_lp_args["knowledge_split"] 

189 print("Splitting the knowledge graph") 

190 if not isinstance(knowledge_split_args, dict): 

191 raise ValueError(f"The knowledge_split_args [{knowledge_split_args}] should be a dict.") 

192 else: 

193 knowledge_split_mode = list(knowledge_split_args.keys())[0] 

194 assert len(knowledge_split_args.keys()) == 1 

195 knowledge_group_by = eval_lp_args["knowledge_group_by"] 

196 else: 

197 knowledge_split_mode = None 

198 knowledge_group_by = None 

199 

200 # split_args is for interaction data 

201 if split_args is None: 

202 raise ValueError("The split_args in eval_args should not be None.") 

203 if not isinstance(split_args, dict): 

204 raise ValueError(f"The split_args [{split_args}] should be a dict.") 

205 

206 split_mode = list(split_args.keys())[0] 

207 

208 assert len(split_args.keys()) == 1 

209 

210 group_by = self.config["eval_args"]["group_by"] 

211 

212 datasets = dict() 

213 if knowledge_split_mode == "RS": 

214 # Manage knowledge graph split 

215 if not isinstance(knowledge_split_args["RS"], list): 

216 raise ValueError( 

217 f'The value of "RS" in knowledge_split_args [{knowledge_split_args}] should be a list.' 

218 ) 

219 

220 if knowledge_group_by is not None: 

221 if knowledge_group_by.lower() == "head": 

222 knowledge_group_by = self.head_entity_field 

223 elif knowledge_group_by.lower() == "tail": 

224 knowledge_group_by = self.tail_entity_field 

225 elif knowledge_group_by.lower() == "relation": 

226 knowledge_group_by = self.relation_field 

227 else: 

228 raise NotImplementedError( 

229 f"The knowledge grouping method [{knowledge_group_by}] has not been implemented." 

230 ) 

231 

232 datasets[KnowledgeEvaluationType.LP] = self.split_by_ratio( 

233 knowledge_split_args["RS"], 

234 data={"data": self.kg_feat, "name": KnowledgeEvaluationType.LP}, 

235 group_by=knowledge_group_by, 

236 ) 

237 

238 if split_mode == "RS": 

239 # Manage interaction split 

240 if not isinstance(split_args["RS"], list): 

241 raise ValueError(f'The value of "RS" in split_args [{split_args}] should be a list.') 

242 if group_by is None: 

243 datasets[KnowledgeEvaluationType.REC] = self.split_by_ratio( 

244 split_args["RS"], 

245 data={"data": self.inter_feat, "name": KnowledgeEvaluationType.REC}, 

246 group_by=None, 

247 ) 

248 elif group_by.lower() == "user": 

249 datasets[KnowledgeEvaluationType.REC] = self.split_by_ratio( 

250 split_args["RS"], 

251 data={"data": self.inter_feat, "name": KnowledgeEvaluationType.REC}, 

252 group_by=self.uid_field, 

253 ) 

254 else: 

255 raise NotImplementedError(f"The grouping method [{group_by}] has not been implemented.") 

256 elif split_mode == "LS": 

257 datasets[KnowledgeEvaluationType.REC] = self.leave_one_out( 

258 group_by=self.uid_field, leave_one_mode=split_args["LS"] 

259 ) 

260 else: 

261 raise NotImplementedError(f"The splitting_method [{split_mode}] has not been implemented.") 

262 return datasets[KnowledgeEvaluationType.REC] if KnowledgeEvaluationType.LP not in datasets else datasets 

263 

264 def copy(self, new_inter_feat, data_type=KnowledgeEvaluationType.REC): 

265 """Given a new interaction feature, return a new :class:`Dataset` object, 

266 whose interaction feature is updated with ``new_inter_feat``, and all the other attributes the same. 

267 

268 Args: 

269 new_inter_feat (Interaction): The new interaction feature need to be updated. 

270 

271 Returns: 

272 :class:`~Dataset`: the new :class:`~Dataset` object, whose interaction feature has been updated. 

273 """ 

274 nxt = copy.copy(self) 

275 if data_type == KnowledgeEvaluationType.REC: 

276 nxt.inter_feat = new_inter_feat 

277 else: 

278 nxt.kg_feat = new_inter_feat 

279 return nxt 

280 

281 def split_by_ratio(self, ratios, data, group_by=None): 

282 """Split interaction records by ratios. 

283 

284 Args: 

285 ratios (list): List of split ratios. No need to be normalized. 

286 group_by (str, optional): Field name that interaction records should grouped by before splitting. 

287 Defaults to ``None`` 

288 

289 Returns: 

290 list: List of :class:`~Dataset`, whose interaction features has been split. 

291 

292 Note: 

293 Other than the first one, each part is rounded down. 

294 """ 

295 

296 self.logger.debug(f"split {data['name']} by ratios [{ratios}], group_by=[{group_by}]") 

297 data_type = data["name"] 

298 data = data["data"] 

299 

300 tot_ratio = sum(ratios) 

301 ratios = [_ / tot_ratio for _ in ratios] 

302 if group_by is None: 

303 split_ids = self._calcu_split_ids(tot=len(data), ratios=ratios) 

304 next_index = [range(start, end) for start, end in zip([0] + split_ids, split_ids + [len(data)])] 

305 

306 else: 

307 grouped_data_feat_index = self._grouped_index(data[group_by].numpy()) 

308 next_index = [[] for _ in range(len(ratios))] 

309 for grouped_index in grouped_data_feat_index: 

310 tot_cnt = len(grouped_index) 

311 split_ids = self._calcu_split_ids(tot=tot_cnt, ratios=ratios) 

312 for index, start, end in zip(next_index, [0] + split_ids, split_ids + [tot_cnt]): 

313 index.extend(grouped_index[start:end]) 

314 

315 self._drop_unused_col() 

316 next_df = [data[index] for index in next_index] 

317 next_ds = [self.copy(split, data_type) for split in next_df] 

318 

319 if data_type == KnowledgeEvaluationType.LP: 

320 # self.kg_feat now have only train data, to prevent data leakage 

321 self.kg_feat = next_df[0] 

322 return next_ds 

323 

324 @property 

325 def tail_num(self): 

326 """Get the number of different tokens of ``self.tail_entity_field``. 

327 

328 Returns: 

329 int: Number of different tokens of ``self.tail_entity_field``. 

330 """ 

331 self._check_field("tail_entity_field") 

332 return self.num(self.tail_entity_field) 

333 

334 def get_tail_feature(self): 

335 """Returns: 

336 Interaction: tails features 

337 """ 

338 

339 if self.tail_feat is None: 

340 self._check_field("tail_entity_field") 

341 return Interaction({self.tail_entity_field: torch.arange(self.tail_num)}) 

342 else: 

343 return self.tail_feat 

344 

345 def _filter_link(self): 

346 """Filter rows of :attr:`item2entity` and :attr:`entity2item`, 

347 whose ``entity_id`` doesn't occur in kg triplets and 

348 ``item_id`` doesn't occur in interaction records. 

349 

350 Dropped items are propagated to :attr:`inter_feat`, :attr:`kg_feat` and :attr:`item_feat`. 

351 """ 

352 while True: 

353 # loop is needed because dropping triples can remove an entity from the kg, 

354 # which in turn can make a still linked item illegal 

355 item_tokens = self._get_rec_item_token() 

356 ent_tokens = self._get_entity_token() 

357 

358 illegal_item = set() 

359 illegal_ent = set() 

360 for item in self.item2entity: 

361 ent = self.item2entity[item] 

362 if item not in item_tokens or ent not in ent_tokens: 

363 illegal_item.add(item) 

364 illegal_ent.add(ent) 

365 for item in illegal_item: 

366 del self.item2entity[item] 

367 for ent in illegal_ent: 

368 del self.entity2item[ent] 

369 

370 remained_inter = pd.Series(True, index=self.inter_feat.index) 

371 remained_inter &= self.inter_feat[self.iid_field].isin(self.item2entity.keys()) 

372 self.inter_feat.drop(self.inter_feat.index[~remained_inter], inplace=True) 

373 

374 # dropped items are propagated to the kg, otherwise their entities would still be 

375 # remapped as plain kg entities, even though the items do not exist anymore 

376 remained_kg = pd.Series(True, index=self.kg_feat.index) 

377 remained_kg &= ~self.kg_feat[self.head_entity_field].isin(illegal_ent) 

378 remained_kg &= ~self.kg_feat[self.tail_entity_field].isin(illegal_ent) 

379 self.kg_feat.drop(self.kg_feat.index[~remained_kg], inplace=True) 

380 

381 # if dropped items are not propagated to item_feat, item_num is larger and 

382 # the entity field2id_token includes mappings of items missing from inter_feat 

383 if self.item_feat is not None: 

384 remained_item = self.item_feat[self.iid_field].isin(self.item2entity.keys()) 

385 self.item_feat.drop(self.item_feat.index[~remained_item], inplace=True) 

386 

387 # feats are re-indexed for safe index dropping and while loop stop conditions 

388 self._reset_index() 

389 

390 if remained_inter.all() and remained_kg.all(): 

391 break 

392 

393 def _download(self): 

394 super()._download() 

395 

396 url = self._get_download_url("kg_url", allow_none=True) 

397 if url is None: 

398 return 

399 self.logger.info(f"Prepare to download linked knowledge graph from [{url}].") 

400 

401 if decide_download(url): 

402 # No need to create dir, as `super()._download()` has created one. 

403 path = download_url(url, self.dataset_path) 

404 extract_zip(path, self.dataset_path) 

405 os.unlink(path) 

406 self.logger.info( 

407 f"\nLinked KG for [{self.dataset_name}] requires additional conversion " 

408 f"to atomic files (.kg and .link).\n" 

409 f"Please refer to https://github.com/RUCAIBox/RecSysDatasets/tree/master/conversion_tools#knowledge-aware-datasets " # noqa: E501 

410 f"for detailed instructions.\n" 

411 f"You can run hopwise after the conversion, see you soon." 

412 ) 

413 sys.exit(0) 

414 else: 

415 self.logger.info("Stop download.") 

416 sys.exit(-1) 

417 

418 def _load_data(self, token, dataset_path): 

419 super()._load_data(token, dataset_path) 

420 self.kg_feat = self._load_kg(self.dataset_name, self.dataset_path) 

421 self.tail_feat = None 

422 self.item2entity, self.entity2item = self._load_link(self.dataset_name, self.dataset_path) 

423 

424 @property 

425 def kg_num(self): 

426 """Get the number of interaction records. 

427 

428 Returns: 

429 int: Number of interaction records. 

430 """ 

431 return len(self.kg_feat) 

432 

433 @property 

434 def sparsity_kg(self): 

435 """Get the sparsity of this dataset. 

436 

437 Returns: 

438 float: Sparsity of this dataset. 

439 """ 

440 return 1 - self.kg_num / (self.entity_num**2) 

441 

442 @property 

443 def sparsity_kg_rel(self): 

444 """Get the sparsity of this dataset. 

445 

446 Returns: 

447 float: Sparsity of this dataset. 

448 """ 

449 return 1 - self.kg_num / (self.entity_num**2 * self.relation_num) 

450 

451 @property 

452 def avg_degree_kg_item(self): 

453 """Get the average degree of items in the knowledge graph. 

454 

455 Returns: 

456 float: Average number of KG triples each item is involved in. 

457 """ # assumes a DataFrame or dict with head, relation, tail 

458 if isinstance(self.kg_feat, pd.DataFrame): 

459 head_counts = self.kg_feat[self.head_entity_field].value_counts() 

460 tail_counts = self.kg_feat[self.tail_entity_field].value_counts() 

461 total_counts = head_counts.add(tail_counts, fill_value=0) 

462 item_degrees = total_counts[total_counts.index.astype(str).isin(self.item2entity.keys())] 

463 return item_degrees.mean() if not item_degrees.empty else 0.0 

464 else: 

465 # fallback if not using pandas 

466 head = self.kg_feat[self.head_entity_field].numpy() 

467 tail = self.kg_feat[self.tail_entity_field].numpy() 

468 counter = Counter(head) + Counter(tail) 

469 item_degrees = [counter[pid] for pid in self.item2entity.keys()] 

470 return np.mean(item_degrees) if item_degrees else 0.0 

471 

472 @property 

473 def avg_degree_kg(self): 

474 """Get the average degree of all entities in the knowledge graph. 

475 

476 Returns: 

477 float: Average number of triples each entity is involved in. 

478 """ 

479 return 2 * self.kg_num / self.entity_num 

480 

481 def __str__(self): 

482 info = [ 

483 super().__str__(), 

484 set_color("The number of entities","green") + f": {self.entity_num}", 

485 set_color("The number of relations","green")+ f": {self.relation_num}", 

486 set_color("The number of triples","green")+ f": {self.kg_num}", 

487 set_color("The number of items that have been linked to KG", "green") + f": {len(self.item2entity)}", 

488 set_color("The number of items that have not been linked to KG", 

489 "green") + f": {self.item_num - len(self.item2entity)}", 

490 set_color("The sparsity of the KG","green") + f": {self.sparsity_kg_rel}", 

491 set_color("The sparsity of the KG (relation-aware)","green") + f": {self.sparsity_kg}", 

492 set_color("The average degree of entities in the KG","green") + f": {self.avg_degree_kg}", 

493 set_color("The average degree of items in the KG","green") + f": {self.avg_degree_kg_item}", 

494 ] # yapf: disable 

495 return "\n".join(info) 

496 

497 def _build_feat_name_list(self): 

498 feat_name_list = super()._build_feat_name_list() 

499 if self.kg_feat is not None: 

500 feat_name_list.append("kg_feat") 

501 return feat_name_list 

502 

503 def _load_kg(self, token, dataset_path): 

504 self.logger.debug(set_color(f"Loading kg from [{dataset_path}].", "green")) 

505 kg_path = os.path.join(dataset_path, f"{token}.kg") 

506 if not os.path.isfile(kg_path): 

507 raise ValueError(f"[{token}.kg] not found in [{dataset_path}].") 

508 df = self._load_feat(kg_path, FeatureSource.KG) 

509 self._check_kg(df) 

510 return df 

511 

512 def _check_kg(self, kg): 

513 kg_warn_message = "kg data requires field [{}]" 

514 assert self.head_entity_field in kg, kg_warn_message.format(self.head_entity_field) 

515 assert self.tail_entity_field in kg, kg_warn_message.format(self.tail_entity_field) 

516 assert self.relation_field in kg, kg_warn_message.format(self.relation_field) 

517 

518 def _load_link(self, token, dataset_path): 

519 self.logger.debug(set_color(f"Loading link from [{dataset_path}].", "green")) 

520 link_path = os.path.join(dataset_path, f"{token}.link") 

521 if not os.path.isfile(link_path): 

522 raise ValueError(f"[{token}.link] not found in [{dataset_path}].") 

523 df = self._load_feat(link_path, "link") 

524 self._check_link(df) 

525 

526 item2entity, entity2item = {}, {} 

527 for item_id, entity_id in zip(df[self.iid_field].values, df[self.entity_field].values): 

528 item2entity[item_id] = entity_id 

529 entity2item[entity_id] = item_id 

530 return item2entity, entity2item 

531 

532 def _check_link(self, link): 

533 link_warn_message = "link data requires field [{}]" 

534 assert self.entity_field in link, link_warn_message.format(self.entity_field) 

535 assert self.iid_field in link, link_warn_message.format(self.iid_field) 

536 

537 def _init_alias(self): 

538 """Add :attr:`alias_of_entity_id`, :attr:`alias_of_relation_id` and update :attr:`_rest_fields`.""" 

539 self._set_alias("entity_id", [self.head_entity_field, self.tail_entity_field]) 

540 self._set_alias("relation_id", [self.relation_field]) 

541 

542 super()._init_alias() 

543 

544 self._rest_fields = np.setdiff1d(self._rest_fields, [self.entity_field], assume_unique=True) 

545 

546 def _get_rec_item_token(self): 

547 """Get set of entity tokens from fields in ``rec`` level.""" 

548 remap_list = self._get_remap_list(self.alias["item_id"]) 

549 tokens, _ = self._concat_remaped_tokens(remap_list) 

550 return set(tokens) 

551 

552 def _get_entity_token(self): 

553 """Get set of entity tokens from fields in ``ent`` level.""" 

554 remap_list = self._get_remap_list(self.alias["entity_id"]) 

555 tokens, _ = self._concat_remaped_tokens(remap_list) 

556 return set(tokens) 

557 

558 def _reset_ent_remapID(self, field, idmap, id2token, token2id): 

559 self.field2id_token[field] = id2token 

560 self.field2token_id[field] = token2id 

561 for feat in self.field2feats(field): 

562 ftype = self.field2type[field] 

563 if ftype == FeatureType.TOKEN: 

564 old_idx = feat[field].values 

565 else: 

566 old_idx = feat[field].agg(np.concatenate) 

567 

568 new_idx = idmap[old_idx] 

569 

570 if ftype == FeatureType.TOKEN: 

571 feat[field] = new_idx 

572 else: 

573 split_point = np.cumsum(feat[field].transform(len))[:-1] 

574 feat[field] = np.split(new_idx, split_point) 

575 

576 def _merge_item_and_entity(self): 

577 """Merge item-id and entity-id into the same id-space.""" 

578 item_token = self.field2id_token[self.iid_field] 

579 entity_token = self.field2id_token[self.head_entity_field] 

580 item_num = len(item_token) 

581 link_num = len(self.item2entity) 

582 entity_num = len(entity_token) 

583 

584 # reset item id 

585 item_priority = np.array([token in self.item2entity for token in item_token]) 

586 item_order = np.argsort(item_priority, kind="stable") 

587 item_id_map = np.zeros_like(item_order) 

588 item_id_map[item_order] = np.arange(item_num) 

589 new_item_id2token = item_token[item_order] 

590 new_item_token2id = {t: i for i, t in enumerate(new_item_id2token)} 

591 for field in self.alias["item_id"]: 

592 self._reset_ent_remapID(field, item_id_map, new_item_id2token, new_item_token2id) 

593 

594 # reset entity id 

595 entity_priority = np.array([token != "[PAD]" and token not in self.entity2item for token in entity_token]) 

596 entity_order = np.argsort(entity_priority, kind="stable") 

597 entity_id_map = np.zeros_like(entity_order) 

598 for i in entity_order[1 : link_num + 1]: 

599 entity_id_map[i] = new_item_token2id[self.entity2item[entity_token[i]]] 

600 entity_id_map[entity_order[link_num + 1 :]] = np.arange(item_num, item_num + entity_num - link_num - 1) 

601 new_entity_id2token = np.concatenate([new_item_id2token, entity_token[entity_order[link_num + 1 :]]]) 

602 for i in range(item_num - link_num, item_num): 

603 new_entity_id2token[i] = self.item2entity[new_entity_id2token[i]] 

604 new_entity_token2id = {t: i for i, t in enumerate(new_entity_id2token)} 

605 for field in self.alias["entity_id"]: 

606 self._reset_ent_remapID(field, entity_id_map, new_entity_id2token, new_entity_token2id) 

607 self.field2id_token[self.entity_field] = new_entity_id2token 

608 self.field2token_id[self.entity_field] = new_entity_token2id 

609 

610 def _add_auxiliary_relation(self): 

611 """Add auxiliary relations in ``self.relation_field``.""" 

612 if self.kg_reverse_r: 

613 # '0' is used for padding, so the number needs to be reduced by one 

614 original_rel_num = len(self.field2id_token[self.relation_field]) - 1 

615 original_hids = self.kg_feat[self.head_entity_field] 

616 original_tids = self.kg_feat[self.tail_entity_field] 

617 original_rels = self.kg_feat[self.relation_field] 

618 

619 # Internal id gap of a relation and its reverse edge is original relation num 

620 reverse_rels = original_rels + original_rel_num 

621 

622 # Add mapping for internal and external ID of relations 

623 for i in range(1, original_rel_num + 1): 

624 original_token = self.field2id_token[self.relation_field][i] 

625 

626 # ui_relation may already exist in the relation field when using pre-trained embeddings 

627 if original_token == self.ui_relation: 

628 continue 

629 

630 reverse_token = original_token + "_r" 

631 self.field2token_id[self.relation_field][reverse_token] = i + original_rel_num 

632 self.field2id_token[self.relation_field] = np.append( 

633 self.field2id_token[self.relation_field], reverse_token 

634 ) 

635 

636 # Update knowledge graph triples with reverse relations 

637 reverse_kg_data = { 

638 self.head_entity_field: original_tids, 

639 self.relation_field: reverse_rels, 

640 self.tail_entity_field: original_hids, 

641 } 

642 reverse_kg_feat = pd.DataFrame(reverse_kg_data) 

643 self.kg_feat = pd.concat([self.kg_feat, reverse_kg_feat]) 

644 

645 # Add UI-relation pairs in the relation field 

646 if self.ui_relation not in self.field2token_id[self.relation_field]: 

647 kg_rel_num = len(self.field2id_token[self.relation_field]) 

648 self.field2token_id[self.relation_field][self.ui_relation] = kg_rel_num 

649 self.field2id_token[self.relation_field] = np.append( 

650 self.field2id_token[self.relation_field], self.ui_relation 

651 ) 

652 

653 def _remap_ID_all(self): 

654 super()._remap_ID_all() 

655 self._merge_item_and_entity() 

656 self._add_auxiliary_relation() 

657 

658 @property 

659 def relation_num(self): 

660 """Get the number of different tokens of ``self.relation_field``. 

661 

662 Returns: 

663 int: Number of different tokens of ``self.relation_field``. 

664 """ 

665 return self.num(self.relation_field) 

666 

667 @property 

668 def entity_num(self): 

669 """Get the number of different tokens of entities, including virtual entities. 

670 

671 Returns: 

672 int: Number of different tokens of entities, including virtual entities. 

673 """ 

674 return self.num(self.entity_field) 

675 

676 @property 

677 def auxiliary_entity_num(self): 

678 """Get the number of different tokens of auxiliary entities (not items). 

679 

680 Returns: 

681 int: Number of different tokens of auxiliary entities. 

682 """ 

683 return self.entity_num - self.item_num 

684 

685 @property 

686 def head_entities(self): 

687 """Returns: 

688 numpy.ndarray: List of head entities of kg triplets. 

689 """ 

690 return self.kg_feat[self.head_entity_field].numpy() 

691 

692 @property 

693 def tail_entities(self): 

694 """Returns: 

695 numpy.ndarray: List of tail entities of kg triplets. 

696 """ 

697 return self.kg_feat[self.tail_entity_field].numpy() 

698 

699 @property 

700 def relations(self): 

701 """Returns: 

702 numpy.ndarray: List of relations of kg triplets. 

703 """ 

704 return self.kg_feat[self.relation_field].numpy() 

705 

706 def norm_ckg_adjacency_matrix(self, form="torch.sparse"): 

707 """Get the collaborative normalized adjacency matrix of users and items. 

708 

709 Construct the square matrix from the training data and normalize it 

710 using the laplace matrix. 

711 

712 .. math:: 

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

714 

715 Args: 

716 form (str, optional): Format of the normalized adjacency matrix. Defaults to ``torch.sparse``. 

717 

718 Returns: 

719 torch.sparse.FloatTensor: Normalized adjacency matrix. 

720 

721 Raises: 

722 NotImplementedError: If the format of the normalized adjacency matrix is not implemented. 

723 """ 

724 if form == "torch.sparse": 

725 return self._create_norm_ckg_adjacency_matrix() 

726 else: 

727 raise NotImplementedError(f"Normalized adjacency matrix format [{form}] has not been implemented.") 

728 

729 def _create_norm_ckg_adjacency_matrix(self, size=None, symmetric=True): 

730 """Get the normalized interaction matrix of users and entities (items) and 

731 the normalized adjacency matrix of the collaborative knowledge graph. 

732 

733 Uses :func:`~hopwise.data.dataset.dataset.Dataset._create_norm_adjacency_matrix` 

734 to get the normalized adjacency matrix of the collaborative knowledge graph 

735 and then extract the normalized interaction matrix of users and entities (items). 

736 

737 Returns: 

738 tuple: tuple of: 

739 - normalized interaction matrix of users and entities (items) 

740 - normalized adjacency matrix of the collaborative knowledge graph. 

741 

742 """ 

743 if size is None: 

744 size = self.user_num + self.entity_num 

745 

746 norm_graph = self._create_norm_adjacency_matrix(size=size, symmetric=symmetric) 

747 if not norm_graph.is_coalesced(): 

748 norm_graph = norm_graph.coalesce() 

749 

750 row, col = norm_graph.indices().cpu().numpy() 

751 values = norm_graph.values().cpu().numpy() 

752 mat = coo_matrix((values, (row, col)), shape=tuple(norm_graph.shape)) 

753 norm_matrix = mat.tocsr()[: self.user_num, self.user_num :].tocoo() 

754 

755 indices = torch.LongTensor(np.array([norm_matrix.row, norm_matrix.col])) 

756 data = torch.FloatTensor(norm_matrix.data) 

757 norm_matrix = torch.sparse.FloatTensor(indices, data, norm_matrix.shape) 

758 

759 return norm_matrix, norm_graph 

760 

761 @property 

762 def entities(self): 

763 """Returns: 

764 numpy.ndarray: List of entity id, including virtual entities. 

765 """ 

766 return np.arange(self.entity_num) 

767 

768 def kg_graph(self, form="coo", value_field=None): 

769 """Get graph or sparse matrix that describe relations between entities. 

770 

771 For an edge of <src, tgt>, ``graph[src, tgt] = 1`` if ``value_field`` is ``None``, 

772 else ``graph[src, tgt] = self.kg_feat[value_field][src, tgt]``. 

773 

774 Currently, we support graph in `PyG`_, 

775 and two type of sparse matrices, ``coo`` and ``csr``. 

776 

777 Args: 

778 form (str, optional): Format of sparse matrix, or library of graph data structure. 

779 Defaults to ``coo``. 

780 value_field (str, optional): edge attributes of graph, or data of sparse matrix, 

781 Defaults to ``None``. 

782 

783 Returns: 

784 Graph / Sparse matrix of kg triplets. 

785 

786 .. _PyG: 

787 https://github.com/rusty1s/pytorch_geometric 

788 """ 

789 args = [ 

790 self.kg_feat, 

791 self.head_entity_field, 

792 self.tail_entity_field, 

793 form, 

794 value_field, 

795 ] 

796 if form in ["coo", "csr"]: 

797 return self._create_sparse_matrix(*args) 

798 elif form in ["pyg"]: 

799 return self._create_graph(*args) 

800 else: 

801 raise NotImplementedError("kg graph format [{}] has not been implemented.") 

802 

803 def _create_ckg_source_target(self, form="numpy"): 

804 """Create base collaborative knowledge graph. 

805 

806 Args: 

807 form (str, optional): The format of the returned graph source and target. 

808 Defaults to ``numpy``. 

809 """ 

810 user_num = self.user_num 

811 

812 if form == "numpy": 

813 hids = self.head_entities + user_num 

814 tids = self.tail_entities + user_num 

815 

816 uids = self.inter_feat[self.uid_field].numpy() 

817 iids = self.inter_feat[self.iid_field].numpy() + user_num 

818 src = np.concatenate([uids, iids, hids]) 

819 tgt = np.concatenate([iids, uids, tids]) 

820 elif form == "torch": 

821 kg_tensor = self.kg_feat 

822 inter_tensor = self.inter_feat 

823 

824 hids = kg_tensor[self.head_entity_field] + user_num 

825 tids = kg_tensor[self.tail_entity_field] + user_num 

826 

827 uids = inter_tensor[self.uid_field] 

828 iids = inter_tensor[self.iid_field] + user_num 

829 

830 src = torch.cat([uids, iids, hids]) 

831 tgt = torch.cat([iids, uids, tids]) 

832 else: 

833 raise NotImplementedError(f"form [{form}] has not been implemented.") 

834 

835 return src, tgt 

836 

837 def _create_ckg_sparse_matrix(self, form="coo", show_relation=False): 

838 src, tgt = self._create_ckg_source_target(form="numpy") 

839 

840 ui_rel_num = self.inter_num 

841 ui_rel_id = self.relation_num - 1 

842 assert self.field2id_token[self.relation_field][ui_rel_id] == self.ui_relation 

843 

844 if not show_relation: 

845 data = np.ones(len(src)) 

846 else: 

847 kg_rel = self.kg_feat[self.relation_field].numpy() 

848 ui_rel = np.full(2 * ui_rel_num, ui_rel_id, dtype=kg_rel.dtype) 

849 data = np.concatenate([ui_rel, kg_rel]) 

850 node_num = self.entity_num + self.user_num 

851 mat = coo_matrix((data, (src, tgt)), shape=(node_num, node_num)) 

852 if form == "coo": 

853 return mat 

854 elif form == "csr": 

855 return mat.tocsr() 

856 else: 

857 raise NotImplementedError(f"Sparse matrix format [{form}] has not been implemented.") 

858 

859 def _create_ckg_graph(self, form="pyg", show_relation=False): 

860 src, tgt = self._create_ckg_source_target(form="torch") 

861 

862 if show_relation: 

863 ui_rel_num = len(self.inter_feat) 

864 

865 ui_rel_id = self.field2token_id[self.relation_field][self.ui_relation] 

866 

867 kg_rel = self.kg_feat[self.relation_field] 

868 ui_rel = torch.full((2 * ui_rel_num,), ui_rel_id, dtype=kg_rel.dtype) 

869 edge = torch.cat([ui_rel, kg_rel]) 

870 

871 if form == "pyg": 

872 from torch_geometric.data import Data 

873 

874 edge_attr = edge if show_relation else None 

875 graph = Data(edge_index=torch.stack([src, tgt]), edge_attr=edge_attr) 

876 return graph 

877 else: 

878 raise NotImplementedError(f"Graph format [{form}] has not been implemented.") 

879 

880 def _create_ckg_igraph(self, show_relation=False, directed=True): 

881 import igraph as ig 

882 

883 vertex_type_attrs = np.concatenate( 

884 [ 

885 [self.uid_field] * self.user_num, 

886 [self.iid_field] * self.item_num, 

887 [self.entity_field] * (self.entity_num - self.item_num), 

888 ], 

889 axis=0, 

890 ) 

891 if show_relation: 

892 n_ui_relations = self.inter_num * 2 if directed else self.inter_num 

893 edge_type_attrs = np.concatenate( 

894 [[self.ui_relation] * n_ui_relations, self.field2id_token[self.relation_field][self.relations]], axis=0 

895 ) 

896 else: 

897 edge_type_attrs = None 

898 

899 if directed: 

900 src, tgt = self._create_ckg_source_target(form="numpy") 

901 else: 

902 user_num = self.user_num 

903 hids = self.head_entities + user_num 

904 tids = self.tail_entities + user_num 

905 

906 uids = self.inter_feat[self.uid_field].numpy() 

907 iids = self.inter_feat[self.iid_field].numpy() + user_num 

908 

909 src = np.concatenate([uids, hids]) 

910 tgt = np.concatenate([iids, tids]) 

911 

912 tuple_graph = list(zip(src, tgt)) 

913 ig_graph = ig.Graph( 

914 edges=tuple_graph, 

915 vertex_attrs={"type": vertex_type_attrs}, 

916 edge_attrs={"type": edge_type_attrs} if show_relation else None, 

917 directed=directed, 

918 ) 

919 

920 return ig_graph 

921 

922 def ckg_graph(self, form="coo", value_field=None): 

923 """Get graph or sparse matrix that describe relations of CKG, 

924 which combines interactions and kg triplets into the same graph. 

925 

926 Item ids and entity ids are added by ``user_num`` temporally. 

927 

928 For an edge of <src, tgt>, ``graph[src, tgt] = 1`` if ``value_field`` is ``None``, 

929 else ``graph[src, tgt] = self.kg_feat[self.relation_field][src, tgt]`` 

930 or ``graph[src, tgt] = self.ui_relation``. 

931 

932 Currently, we support graph in `PyG`_ and `igraph`_, 

933 two type of sparse matrices, ``coo`` and ``csr``. 

934 

935 Args: 

936 form (str, optional): Format of sparse matrix, or library of graph data structure. 

937 Defaults to ``coo``. 

938 value_field (str, optional): ``self.relation_field`` or ``None``, 

939 Defaults to ``None``. 

940 

941 Returns: 

942 Graph / Sparse matrix of kg triplets. 

943 

944 .. _PyG: 

945 https://github.com/rusty1s/pytorch_geometric 

946 

947 .. _igraph: 

948 https://python.igraph.org/en/stable/ 

949 """ 

950 if value_field is not None and value_field != self.relation_field: 

951 raise ValueError(f"Value_field [{value_field}] can only be [{self.relation_field}] in ckg_graph.") 

952 show_relation = value_field is not None 

953 

954 if form in ["coo", "csr"]: 

955 return self._create_ckg_sparse_matrix(form, show_relation) 

956 elif form in ["pyg"]: 

957 return self._create_ckg_graph(form, show_relation) 

958 elif form == "igraph": 

959 return self._create_ckg_igraph(show_relation) 

960 else: 

961 raise NotImplementedError("ckg graph format [{}] has not been implemented.") 

962 

963 def ckg_dict_graph(self, ui_bidirectional=True): 

964 """Get a dictionary representation of the collaborative knowledge graph. 

965 Returns: 

966 dict: Dictionary representation of the collaborative knowledge graph. 

967 """ 

968 uids = self.inter_feat[self.uid_field].numpy() 

969 iids = self.inter_feat[self.iid_field].numpy() 

970 

971 src = np.concatenate([uids, self.head_entities]) 

972 tgt = np.concatenate([iids, self.tail_entities]) 

973 

974 ui_relation_id = self.field2token_id[self.relation_field][self.ui_relation] 

975 rels = np.concatenate([np.full(self.inter_num, ui_relation_id), self.relations]) 

976 

977 graph_dict = {"user": {}, "entity": {}} 

978 for idx, (src_id, rel_id, tgt_id) in enumerate(zip(src, rels, tgt)): 

979 if rel_id == ui_relation_id: 

980 src_type = "user" 

981 end_type = "entity" 

982 

983 if src_id not in graph_dict[src_type]: 

984 graph_dict[src_type][src_id] = dict() 

985 if rel_id not in graph_dict[src_type][src_id]: 

986 graph_dict[src_type][src_id][rel_id] = list() 

987 

988 # UI interaction case 

989 graph_dict[src_type][src_id][rel_id].append(tgt_id) 

990 if ui_bidirectional: 

991 if tgt_id not in graph_dict[end_type]: 

992 graph_dict[end_type][tgt_id] = dict() 

993 if rel_id not in graph_dict[end_type][tgt_id]: 

994 graph_dict[end_type][tgt_id][rel_id] = list() 

995 

996 graph_dict[end_type][tgt_id][rel_id].append(src_id) 

997 

998 else: 

999 if src_id not in graph_dict["entity"]: 

1000 graph_dict["entity"][src_id] = dict() 

1001 if rel_id not in graph_dict["entity"][src_id]: 

1002 graph_dict["entity"][src_id][rel_id] = list() 

1003 

1004 if tgt_id not in graph_dict["entity"]: 

1005 graph_dict["entity"][tgt_id] = dict() 

1006 if rel_id not in graph_dict["entity"][tgt_id]: 

1007 graph_dict["entity"][tgt_id][rel_id] = list() 

1008 

1009 # KG case 

1010 graph_dict["entity"][src_id][rel_id].append(tgt_id) 

1011 graph_dict["entity"][tgt_id][rel_id].append(src_id) 

1012 

1013 return graph_dict 

1014 

1015 

1016class UserItemKnowledgeBasedDataset(KnowledgeBasedDataset): 

1017 """:class:`UserItemKnowledgeBasedDataset` is based on :class:`~hopwise.data.dataset.dataset.KnowledgeBasedDataset`, 

1018 and load ``.kg`` and ``.user_link`` and ``.item_link`` additionally. 

1019 

1020 Entities are remapped together with ``user_id`` and ``item_id`` specially. 

1021 All entities are remapped into three consecutive ID sections. 

1022 

1023 - virtual entities that only exist in interaction data. 

1024 - entities that exist both in interaction data and kg triplets. 

1025 - entities only exist in kg triplets. 

1026 

1027 It also provides several interfaces to transfer ``.kg`` features into coo sparse matrix, 

1028 csr sparse matrix or :class:`PyG.Data`. 

1029 

1030 Attributes: 

1031 head_entity_field (str): The same as ``config['HEAD_ENTITY_ID_FIELD']``. 

1032 

1033 tail_entity_field (str): The same as ``config['TAIL_ENTITY_ID_FIELD']``. 

1034 

1035 relation_field (str): The same as ``config['RELATION_ID_FIELD']``. 

1036 

1037 entity_field (str): The same as ``config['ENTITY_ID_FIELD']``. 

1038 

1039 kg_feat (pandas.DataFrame): Internal data structure stores the kg triplets. 

1040 It's loaded from file ``.kg``. 

1041 

1042 user2entity (dict): Dict maps ``user_id`` to ``entity``, 

1043 which is loaded from file ``.user_link``. 

1044 

1045 entity2user (dict): Dict maps ``entity`` to ``user_id``, 

1046 which is loaded from file ``.user_link``. 

1047 

1048 item2entity (dict): Dict maps ``item_id`` to ``entity``, 

1049 which is loaded from file ``.item_link``. 

1050 

1051 entity2item (dict): Dict maps ``entity`` to ``item_id``, 

1052 which is loaded from file ``.item_link``. 

1053 

1054 Note: 

1055 :attr:`entity_field` doesn't exist exactly. It's only a symbol, 

1056 representing entity features. 

1057 

1058 :attr:`ui_relation` is a special relation token, which is used to represent 

1059 the interaction relation between users and items. 

1060 """ 

1061 

1062 @property 

1063 def auxiliary_entity_num(self): 

1064 """Get the number of different tokens of auxiliary entities (not users nor items). 

1065 

1066 Returns: 

1067 int: Number of different tokens of auxiliary entities. 

1068 """ 

1069 return self.entity_num - self.user_num - self.item_num 

1070 

1071 def _filter_link(self): 

1072 """Filter rows of :attr:`item2entity` and :attr:`entity2item`, 

1073 whose ``entity_id`` doesn't occur in kg triplets and 

1074 ``item_id`` doesn't occur in interaction records. 

1075 Extended to also filter rows of :attr:`user2entity` and :attr:`entity2user`, 

1076 whose ``entity_id`` doesn't occur in kg triplets and 

1077 ``user_id`` doesn't occur in interaction records. 

1078 

1079 Dropped users and items are propagated to :attr:`inter_feat`, :attr:`kg_feat`, 

1080 :attr:`item_feat` and :attr:`user_feat`. 

1081 """ 

1082 while True: 

1083 # loop is needed in case dropped index lead to drop of user/item 

1084 # causing incompatibility between link mappings and field2id_token 

1085 item_tokens = self._get_rec_token("item_id") 

1086 user_tokens = self._get_rec_token("user_id") 

1087 ent_tokens = self._get_entity_token() 

1088 

1089 illegal_item = set() 

1090 illegal_item_ent = set() 

1091 for item in self.item2entity: 

1092 ent = self.item2entity[item] 

1093 if item not in item_tokens or ent not in ent_tokens: 

1094 illegal_item.add(item) 

1095 illegal_item_ent.add(ent) 

1096 for item in illegal_item: 

1097 del self.item2entity[item] 

1098 for ent in illegal_item_ent: 

1099 del self.entity2item[ent] 

1100 

1101 remained_inter = pd.Series(True, index=self.inter_feat.index) 

1102 remained_inter &= self.inter_feat[self.iid_field].isin(self.item2entity.keys()) 

1103 

1104 illegal_user = set() 

1105 illegal_user_ent = set() 

1106 for user in self.user2entity: 

1107 ent = self.user2entity[user] 

1108 if user not in user_tokens or ent not in ent_tokens: 

1109 illegal_user.add(user) 

1110 illegal_user_ent.add(ent) 

1111 for user in illegal_user: 

1112 del self.user2entity[user] 

1113 for ent in illegal_user_ent: 

1114 del self.entity2user[ent] 

1115 

1116 remained_inter &= self.inter_feat[self.uid_field].isin(self.user2entity.keys()) 

1117 self.inter_feat.drop(self.inter_feat.index[~remained_inter], inplace=True) 

1118 

1119 # dropped users and items are propagated to the kg, otherwise their entities would still 

1120 # be remapped as plain kg entities, even though they do not exist anymore 

1121 illegal_ent = illegal_item_ent | illegal_user_ent 

1122 remained_kg = pd.Series(True, index=self.kg_feat.index) 

1123 remained_kg &= ~self.kg_feat[self.head_entity_field].isin(illegal_ent) 

1124 remained_kg &= ~self.kg_feat[self.tail_entity_field].isin(illegal_ent) 

1125 self.kg_feat.drop(self.kg_feat.index[~remained_kg], inplace=True) 

1126 

1127 # if dropped users/items are not propagated to user_feat/item_feat, user_num and item_num 

1128 # are larger and the entity field2id_token includes mappings missing from inter_feat 

1129 if self.item_feat is not None: 

1130 remained_item = self.item_feat[self.iid_field].isin(self.item2entity.keys()) 

1131 self.item_feat.drop(self.item_feat.index[~remained_item], inplace=True) 

1132 

1133 if self.user_feat is not None: 

1134 remained_user = self.user_feat[self.uid_field].isin(self.user2entity.keys()) 

1135 self.user_feat.drop(self.user_feat.index[~remained_user], inplace=True) 

1136 

1137 # feats are re-indexed for safe index dropping and while loop stop conditions 

1138 self._reset_index() 

1139 

1140 if remained_inter.all() and remained_kg.all(): 

1141 break 

1142 

1143 def _load_data(self, token, dataset_path): 

1144 super(KnowledgeBasedDataset, self)._load_data(token, dataset_path) 

1145 self.kg_feat = self._load_kg(self.dataset_name, self.dataset_path) 

1146 self.tail_feat = None 

1147 self.item2entity, self.entity2item, self.user2entity, self.entity2user = self._load_link( 

1148 self.dataset_name, self.dataset_path 

1149 ) 

1150 

1151 def __str__(self): 

1152 info = [ 

1153 super().__str__(), 

1154 set_color("The number of users that have been linked to KG", "green") + f": {len(self.user2entity)}", 

1155 ] 

1156 return "\n".join(info) 

1157 

1158 def _load_link(self, token, dataset_path): 

1159 self.logger.debug(set_color(f"Loading link from [{dataset_path}].", "green")) 

1160 item_link_path = os.path.join(dataset_path, f"{token}.item_link") 

1161 user_link_path = os.path.join(dataset_path, f"{token}.user_link") 

1162 if not os.path.isfile(item_link_path) and not os.path.isfile(user_link_path): 

1163 raise ValueError(f"[{token}.item_link] and [{token}.user_link] not found in [{dataset_path}].") 

1164 item_df = self._load_feat(item_link_path, "item_link") 

1165 user_df = self._load_feat(user_link_path, "user_link") 

1166 self._check_link(item_df, user_df) 

1167 

1168 item2entity, entity2item = {}, {} 

1169 for item_id, entity_id in zip(item_df[self.iid_field].values, item_df[self.entity_field].values): 

1170 item2entity[item_id] = entity_id 

1171 entity2item[entity_id] = item_id 

1172 

1173 user2entity, entity2user = {}, {} 

1174 for user_id, entity_id in zip(user_df[self.uid_field].values, user_df[self.entity_field].values): 

1175 user2entity[user_id] = entity_id 

1176 entity2user[entity_id] = user_id 

1177 

1178 return item2entity, entity2item, user2entity, entity2user 

1179 

1180 def _check_link(self, item_link, user_link): 

1181 link_warn_message = "link data requires field [{}]" 

1182 assert self.entity_field in item_link, link_warn_message.format(self.entity_field) 

1183 assert self.iid_field in item_link, link_warn_message.format(self.iid_field) 

1184 assert self.entity_field in user_link, link_warn_message.format(self.entity_field) 

1185 assert self.uid_field in user_link, link_warn_message.format(self.uid_field) 

1186 

1187 def _get_rec_token(self, field): 

1188 """Get set of entity tokens from fields in ``rec`` level.""" 

1189 remap_list = self._get_remap_list(self.alias[field]) 

1190 tokens, _ = self._concat_remaped_tokens(remap_list) 

1191 return set(tokens) 

1192 

1193 def _merge_item_and_entity(self): 

1194 """Merge item-id and entity-id into the same id-space.""" 

1195 item_token = self.field2id_token[self.iid_field] 

1196 user_token = self.field2id_token[self.uid_field] 

1197 entity_token = self.field2id_token[self.head_entity_field] 

1198 item_num = len(item_token) 

1199 user_num = len(user_token) 

1200 item_link_num = len(self.item2entity) 

1201 user_link_num = len(self.user2entity) 

1202 entity_num = len(entity_token) 

1203 

1204 # reset user id 

1205 user_priority = np.array([token in self.user2entity for token in user_token]) 

1206 user_order = np.argsort(user_priority, kind="stable") 

1207 user_id_map = np.zeros_like(user_order) 

1208 user_id_map[user_order] = np.arange(user_num) 

1209 new_user_id2token = user_token[user_order] 

1210 new_user_token2id = {t: i for i, t in enumerate(new_user_id2token)} 

1211 for field in self.alias["user_id"]: 

1212 self._reset_ent_remapID(field, user_id_map, new_user_id2token, new_user_token2id) 

1213 

1214 # reset item id 

1215 item_priority = np.array([token in self.item2entity for token in item_token]) 

1216 item_order = np.argsort(item_priority, kind="stable") 

1217 item_id_map = np.zeros_like(item_order) 

1218 item_id_map[item_order] = np.arange(item_num) 

1219 new_item_id2token = item_token[item_order] 

1220 new_item_token2id = {t: i for i, t in enumerate(new_item_id2token)} 

1221 for field in self.alias["item_id"]: 

1222 self._reset_ent_remapID(field, item_id_map, new_item_id2token, new_item_token2id) 

1223 

1224 # reset entity id 

1225 entity_priority = np.array( 

1226 [ # these values will be used to set the order in which the entities are remapped 

1227 # 0 for padding and user, 1 for item, 2 for other entities 

1228 0 if token == "[PAD]" or token in self.entity2user else (1 if token in self.entity2item else 2) 

1229 for token in entity_token 

1230 ] 

1231 ) 

1232 entity_order = np.argsort(entity_priority, kind="stable") 

1233 entity_id_map = np.zeros_like(entity_order) 

1234 for i in entity_order[1 : user_link_num + 1]: 

1235 entity_id_map[i] = new_user_token2id[self.entity2user[entity_token[i]]] 

1236 new_item_entity_token2id = {t: i + self.user_num for i, t in enumerate(new_item_id2token)} 

1237 for i in entity_order[user_link_num + 1 : user_link_num + item_link_num + 1]: 

1238 entity_id_map[i] = new_item_entity_token2id[self.entity2item[entity_token[i]]] 

1239 entity_id_map[entity_order[user_link_num + item_link_num + 1 :]] = np.arange( 

1240 user_num + item_num, user_num + item_num + entity_num - user_link_num - item_link_num - 1 

1241 ) 

1242 new_entity_id2token = np.concatenate( 

1243 [new_user_id2token, new_item_id2token, entity_token[entity_order[user_link_num + item_link_num + 1 :]]] 

1244 ) 

1245 for i in range(user_num - user_link_num, user_num): 

1246 new_entity_id2token[i] = self.user2entity[new_entity_id2token[i]] 

1247 for i in range(user_num + item_num - item_link_num, user_num + item_num): 

1248 new_entity_id2token[i] = self.item2entity[new_entity_id2token[i]] 

1249 new_entity_token2id = {t: i for i, t in enumerate(new_entity_id2token)} 

1250 for field in self.alias["entity_id"]: 

1251 self._reset_ent_remapID(field, entity_id_map, new_entity_id2token, new_entity_token2id) 

1252 self.field2id_token[self.entity_field] = new_entity_id2token 

1253 self.field2token_id[self.entity_field] = new_entity_token2id 

1254 

1255 def _filter_kg_by_triple_num(self): 

1256 """Filter by number of triples. 

1257 

1258 The interval of the number of triples can be set, and only entities/relations 

1259 whose number of triples is in the specified interval can be retained. 

1260 See :doc:`../user_guide/data/data_args` for detail arg setting. 

1261 

1262 Note: 

1263 Lower bound of the interval is also called k-core filtering, which means this method 

1264 will filter loops until all the entities and relations has at least k triples. 

1265 """ 

1266 entity_kg_num_interval = self._parse_intervals_str(self.config["entity_kg_num_interval"]) 

1267 relation_kg_num_interval = self._parse_intervals_str(self.config["relation_kg_num_interval"]) 

1268 user_entity_kg_num_interval = self._parse_intervals_str(self.config["user_entity_kg_num_interval"]) 

1269 

1270 if entity_kg_num_interval is None and relation_kg_num_interval is None: 

1271 return 

1272 

1273 entity_kg_num = Counter() 

1274 if entity_kg_num_interval is not None or user_entity_kg_num_interval is not None: 

1275 head_entity_kg_num = Counter(self.kg_feat[self.head_entity_field].values) 

1276 tail_entity_kg_num = Counter(self.kg_feat[self.tail_entity_field].values) 

1277 entity_kg_num = head_entity_kg_num + tail_entity_kg_num 

1278 relation_kg_num = Counter(self.kg_feat[self.relation_field].values) if relation_kg_num_interval else Counter() 

1279 

1280 while True: 

1281 item_entity_kg_num = Counter({k: v for k, v in entity_kg_num.items() if k not in self.entity2user}) 

1282 

1283 item_ban_entities = self._get_illegal_ids_by_inter_num( 

1284 field=f"{self.head_entity_field}-{self.tail_entity_field}", 

1285 feat=None, 

1286 inter_num=item_entity_kg_num, 

1287 inter_interval=entity_kg_num_interval, 

1288 ) 

1289 

1290 if user_entity_kg_num_interval is None: 

1291 ban_entities = item_ban_entities 

1292 else: 

1293 user_entity_kg_num = Counter({k: v for k, v in entity_kg_num.items() if k in self.entity2user}) 

1294 

1295 user_ban_entities = self._get_illegal_ids_by_inter_num( 

1296 field=f"{self.head_entity_field}-{self.tail_entity_field}", 

1297 feat=None, 

1298 inter_num=user_entity_kg_num, 

1299 inter_interval=user_entity_kg_num_interval, 

1300 ) 

1301 

1302 ban_entities = item_ban_entities | user_ban_entities 

1303 

1304 ban_relations = self._get_illegal_ids_by_inter_num( 

1305 field=self.relation_field, 

1306 feat=None, 

1307 inter_num=relation_kg_num, 

1308 inter_interval=relation_kg_num_interval, 

1309 ) 

1310 if len(ban_entities) == 0 and len(ban_relations) == 0: 

1311 break 

1312 

1313 dropped_kg = pd.Series(False, index=self.kg_feat.index) 

1314 head_entity_kg = self.kg_feat[self.head_entity_field] 

1315 tail_entity_kg = self.kg_feat[self.tail_entity_field] 

1316 relation_kg = self.kg_feat[self.relation_field] 

1317 dropped_kg |= head_entity_kg.isin(ban_entities) 

1318 dropped_kg |= tail_entity_kg.isin(ban_entities) 

1319 dropped_kg |= relation_kg.isin(ban_relations) 

1320 

1321 entity_kg_num -= Counter(head_entity_kg[dropped_kg].values) 

1322 entity_kg_num -= Counter(tail_entity_kg[dropped_kg].values) 

1323 relation_kg_num -= Counter(relation_kg[dropped_kg].values) 

1324 

1325 dropped_index = self.kg_feat.index[dropped_kg] 

1326 self.logger.debug(f"[{len(dropped_index)}] dropped triples.") 

1327 self.kg_feat.drop(dropped_index, inplace=True) 

1328 

1329 def _create_ckg_source_target(self, form="numpy"): 

1330 """Create base collaborative knowledge graph. 

1331 

1332 Args: 

1333 form (str, optional): The format of the returned graph source and target. 

1334 Defaults to ``numpy``. 

1335 """ 

1336 if form == "numpy": 

1337 hids = self.head_entities 

1338 tids = self.tail_entities 

1339 

1340 uids = self.inter_feat[self.uid_field].numpy() 

1341 iids = self.inter_feat[self.iid_field].numpy() + self.user_num 

1342 

1343 src = np.concatenate([uids, iids, hids]) 

1344 tgt = np.concatenate([iids, uids, tids]) 

1345 elif form == "torch": 

1346 kg_tensor = self.kg_feat 

1347 inter_tensor = self.inter_feat 

1348 

1349 hids = kg_tensor[self.head_entity_field] 

1350 tids = kg_tensor[self.tail_entity_field] 

1351 

1352 uids = inter_tensor[self.uid_field] 

1353 iids = inter_tensor[self.iid_field] + self.user_num 

1354 

1355 src = torch.cat([uids, iids, hids]) 

1356 tgt = torch.cat([iids, uids, tids]) 

1357 else: 

1358 raise NotImplementedError(f"form [{form}] has not been implemented.") 

1359 

1360 return src, tgt 

1361 

1362 def _create_ckg_sparse_matrix(self, form="coo", show_relation=False): 

1363 src, tgt = self._create_ckg_source_target(form="numpy") 

1364 

1365 ui_rel_num = self.inter_num 

1366 ui_rel_id = self.relation_num - 1 

1367 assert self.field2id_token[self.relation_field][ui_rel_id] == self.ui_relation 

1368 

1369 if not show_relation: 

1370 data = np.ones(len(src)) 

1371 else: 

1372 kg_rel = self.kg_feat[self.relation_field].numpy() 

1373 ui_rel = np.full(2 * ui_rel_num, ui_rel_id, dtype=kg_rel.dtype) 

1374 data = np.concatenate([ui_rel, kg_rel]) 

1375 mat = coo_matrix((data, (src, tgt)), shape=(self.entity_num, self.entity_num)) 

1376 if form == "coo": 

1377 return mat 

1378 elif form == "csr": 

1379 return mat.tocsr() 

1380 else: 

1381 raise NotImplementedError(f"Sparse matrix format [{form}] has not been implemented.") 

1382 

1383 def _create_ckg_igraph(self, show_relation=False, directed=True): 

1384 import igraph as ig 

1385 

1386 vertex_type_attrs = np.concatenate( 

1387 [ 

1388 [self.uid_field] * self.user_num, 

1389 [self.iid_field] * self.item_num, 

1390 [self.entity_field] * (self.auxiliary_entity_num), 

1391 ], 

1392 axis=0, 

1393 ) 

1394 if show_relation: 

1395 n_ui_relations = self.inter_num * 2 if directed else self.inter_num 

1396 edge_type_attrs = np.concatenate( 

1397 [[self.ui_relation] * n_ui_relations, self.field2id_token[self.relation_field][self.relations]], axis=0 

1398 ) 

1399 else: 

1400 edge_type_attrs = None 

1401 

1402 if directed: 

1403 src, tgt = self._create_ckg_source_target(form="numpy") 

1404 else: 

1405 hids = self.head_entities 

1406 tids = self.tail_entities 

1407 

1408 uids = self.inter_feat[self.uid_field].numpy() 

1409 iids = self.inter_feat[self.iid_field].numpy() + self.user_num 

1410 

1411 src = np.concatenate([uids, hids]) 

1412 tgt = np.concatenate([iids, tids]) 

1413 

1414 tuple_graph = list(zip(src, tgt)) 

1415 ig_graph = ig.Graph( 

1416 edges=tuple_graph, 

1417 vertex_attrs={"type": vertex_type_attrs}, 

1418 edge_attrs={"type": edge_type_attrs} if show_relation else None, 

1419 directed=directed, 

1420 ) 

1421 

1422 return ig_graph