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
« 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
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
10# UPDATE:
11# @Time : 2025
12# @Author : Giacomo Medda, Alessandro Soccol
13# @Email : giacomo.medda@unica.it, alessandro.soccol@unica.it
15"""hopwise.data.kg_dataset
16 hopwise.data.user_item_kg_dataset
17##########################
18"""
20import copy
21import os
22import sys
23from collections import Counter
25import numpy as np
26import pandas as pd
27import torch
28from scipy.sparse import coo_matrix
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
36class KnowledgeBasedDataset(Dataset):
37 """:class:`KnowledgeBasedDataset` is based on :class:`~hopwise.data.dataset.dataset.Dataset`,
38 and load ``.kg`` and ``.link`` additionally.
40 Entities are remapped together with ``item_id`` specially.
41 All entities are remapped into three consecutive ID sections.
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.
47 It also provides several interfaces to transfer ``.kg`` features into coo sparse matrix,
48 csr sparse matrix or :class:`PyG.Data`.
50 Attributes:
51 head_entity_field (str): The same as ``config['HEAD_ENTITY_ID_FIELD']``.
53 tail_entity_field (str): The same as ``config['TAIL_ENTITY_ID_FIELD']``.
55 relation_field (str): The same as ``config['RELATION_ID_FIELD']``.
57 entity_field (str): The same as ``config['ENTITY_ID_FIELD']``.
59 kg_feat (pandas.DataFrame): Internal data structure stores the kg triplets.
60 It's loaded from file ``.kg``.
62 item2entity (dict): Dict maps ``item_id`` to ``entity``,
63 which is loaded from file ``.link``.
65 entity2item (dict): Dict maps ``entity`` to ``item_id``,
66 which is loaded from file ``.link``.
68 Note:
69 :attr:`entity_field` doesn't exist exactly. It's only a symbol,
70 representing entity features.
72 :attr:`ui_relation` is a special relation token, which is used to represent
73 the interaction relation between users and items.
74 """
76 def __init__(self, config):
77 super().__init__(config)
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)
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}")
95 def _data_filtering(self):
96 super()._data_filtering()
97 self._filter_kg_by_triple_num()
98 self._filter_link()
100 def _filter_kg_by_triple_num(self):
101 """Filter by number of triples.
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.
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"])
114 if entity_kg_num_interval is None and relation_kg_num_interval is None:
115 return
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)
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()
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
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)
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)
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)
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.
163 Returns:
164 list: List of built :class:`Dataset`.
165 """
166 self._change_feat_format()
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
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.")
183 # splitting & grouping
184 split_args = self.config["eval_args"]["split"]
185 eval_lp_args = self.config["eval_lp_args"]
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
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.")
206 split_mode = list(split_args.keys())[0]
208 assert len(split_args.keys()) == 1
210 group_by = self.config["eval_args"]["group_by"]
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 )
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 )
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 )
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
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.
268 Args:
269 new_inter_feat (Interaction): The new interaction feature need to be updated.
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
281 def split_by_ratio(self, ratios, data, group_by=None):
282 """Split interaction records by ratios.
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``
289 Returns:
290 list: List of :class:`~Dataset`, whose interaction features has been split.
292 Note:
293 Other than the first one, each part is rounded down.
294 """
296 self.logger.debug(f"split {data['name']} by ratios [{ratios}], group_by=[{group_by}]")
297 data_type = data["name"]
298 data = data["data"]
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)])]
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])
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]
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
324 @property
325 def tail_num(self):
326 """Get the number of different tokens of ``self.tail_entity_field``.
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)
334 def get_tail_feature(self):
335 """Returns:
336 Interaction: tails features
337 """
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
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.
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()
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]
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)
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)
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)
387 # feats are re-indexed for safe index dropping and while loop stop conditions
388 self._reset_index()
390 if remained_inter.all() and remained_kg.all():
391 break
393 def _download(self):
394 super()._download()
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}].")
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)
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)
424 @property
425 def kg_num(self):
426 """Get the number of interaction records.
428 Returns:
429 int: Number of interaction records.
430 """
431 return len(self.kg_feat)
433 @property
434 def sparsity_kg(self):
435 """Get the sparsity of this dataset.
437 Returns:
438 float: Sparsity of this dataset.
439 """
440 return 1 - self.kg_num / (self.entity_num**2)
442 @property
443 def sparsity_kg_rel(self):
444 """Get the sparsity of this dataset.
446 Returns:
447 float: Sparsity of this dataset.
448 """
449 return 1 - self.kg_num / (self.entity_num**2 * self.relation_num)
451 @property
452 def avg_degree_kg_item(self):
453 """Get the average degree of items in the knowledge graph.
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
472 @property
473 def avg_degree_kg(self):
474 """Get the average degree of all entities in the knowledge graph.
476 Returns:
477 float: Average number of triples each entity is involved in.
478 """
479 return 2 * self.kg_num / self.entity_num
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)
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
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
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)
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)
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
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)
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])
542 super()._init_alias()
544 self._rest_fields = np.setdiff1d(self._rest_fields, [self.entity_field], assume_unique=True)
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)
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)
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)
568 new_idx = idmap[old_idx]
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)
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)
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)
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
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]
619 # Internal id gap of a relation and its reverse edge is original relation num
620 reverse_rels = original_rels + original_rel_num
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]
626 # ui_relation may already exist in the relation field when using pre-trained embeddings
627 if original_token == self.ui_relation:
628 continue
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 )
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])
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 )
653 def _remap_ID_all(self):
654 super()._remap_ID_all()
655 self._merge_item_and_entity()
656 self._add_auxiliary_relation()
658 @property
659 def relation_num(self):
660 """Get the number of different tokens of ``self.relation_field``.
662 Returns:
663 int: Number of different tokens of ``self.relation_field``.
664 """
665 return self.num(self.relation_field)
667 @property
668 def entity_num(self):
669 """Get the number of different tokens of entities, including virtual entities.
671 Returns:
672 int: Number of different tokens of entities, including virtual entities.
673 """
674 return self.num(self.entity_field)
676 @property
677 def auxiliary_entity_num(self):
678 """Get the number of different tokens of auxiliary entities (not items).
680 Returns:
681 int: Number of different tokens of auxiliary entities.
682 """
683 return self.entity_num - self.item_num
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()
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()
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()
706 def norm_ckg_adjacency_matrix(self, form="torch.sparse"):
707 """Get the collaborative normalized adjacency matrix of users and items.
709 Construct the square matrix from the training data and normalize it
710 using the laplace matrix.
712 .. math::
713 A_{hat} = D^{-0.5} \times A \times D^{-0.5}
715 Args:
716 form (str, optional): Format of the normalized adjacency matrix. Defaults to ``torch.sparse``.
718 Returns:
719 torch.sparse.FloatTensor: Normalized adjacency matrix.
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.")
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.
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).
737 Returns:
738 tuple: tuple of:
739 - normalized interaction matrix of users and entities (items)
740 - normalized adjacency matrix of the collaborative knowledge graph.
742 """
743 if size is None:
744 size = self.user_num + self.entity_num
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()
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()
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)
759 return norm_matrix, norm_graph
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)
768 def kg_graph(self, form="coo", value_field=None):
769 """Get graph or sparse matrix that describe relations between entities.
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]``.
774 Currently, we support graph in `PyG`_,
775 and two type of sparse matrices, ``coo`` and ``csr``.
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``.
783 Returns:
784 Graph / Sparse matrix of kg triplets.
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.")
803 def _create_ckg_source_target(self, form="numpy"):
804 """Create base collaborative knowledge graph.
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
812 if form == "numpy":
813 hids = self.head_entities + user_num
814 tids = self.tail_entities + user_num
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
824 hids = kg_tensor[self.head_entity_field] + user_num
825 tids = kg_tensor[self.tail_entity_field] + user_num
827 uids = inter_tensor[self.uid_field]
828 iids = inter_tensor[self.iid_field] + user_num
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.")
835 return src, tgt
837 def _create_ckg_sparse_matrix(self, form="coo", show_relation=False):
838 src, tgt = self._create_ckg_source_target(form="numpy")
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
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.")
859 def _create_ckg_graph(self, form="pyg", show_relation=False):
860 src, tgt = self._create_ckg_source_target(form="torch")
862 if show_relation:
863 ui_rel_num = len(self.inter_feat)
865 ui_rel_id = self.field2token_id[self.relation_field][self.ui_relation]
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])
871 if form == "pyg":
872 from torch_geometric.data import Data
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.")
880 def _create_ckg_igraph(self, show_relation=False, directed=True):
881 import igraph as ig
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
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
906 uids = self.inter_feat[self.uid_field].numpy()
907 iids = self.inter_feat[self.iid_field].numpy() + user_num
909 src = np.concatenate([uids, hids])
910 tgt = np.concatenate([iids, tids])
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 )
920 return ig_graph
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.
926 Item ids and entity ids are added by ``user_num`` temporally.
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``.
932 Currently, we support graph in `PyG`_ and `igraph`_,
933 two type of sparse matrices, ``coo`` and ``csr``.
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``.
941 Returns:
942 Graph / Sparse matrix of kg triplets.
944 .. _PyG:
945 https://github.com/rusty1s/pytorch_geometric
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
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.")
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()
971 src = np.concatenate([uids, self.head_entities])
972 tgt = np.concatenate([iids, self.tail_entities])
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])
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"
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()
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()
996 graph_dict[end_type][tgt_id][rel_id].append(src_id)
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()
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()
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)
1013 return graph_dict
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.
1020 Entities are remapped together with ``user_id`` and ``item_id`` specially.
1021 All entities are remapped into three consecutive ID sections.
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.
1027 It also provides several interfaces to transfer ``.kg`` features into coo sparse matrix,
1028 csr sparse matrix or :class:`PyG.Data`.
1030 Attributes:
1031 head_entity_field (str): The same as ``config['HEAD_ENTITY_ID_FIELD']``.
1033 tail_entity_field (str): The same as ``config['TAIL_ENTITY_ID_FIELD']``.
1035 relation_field (str): The same as ``config['RELATION_ID_FIELD']``.
1037 entity_field (str): The same as ``config['ENTITY_ID_FIELD']``.
1039 kg_feat (pandas.DataFrame): Internal data structure stores the kg triplets.
1040 It's loaded from file ``.kg``.
1042 user2entity (dict): Dict maps ``user_id`` to ``entity``,
1043 which is loaded from file ``.user_link``.
1045 entity2user (dict): Dict maps ``entity`` to ``user_id``,
1046 which is loaded from file ``.user_link``.
1048 item2entity (dict): Dict maps ``item_id`` to ``entity``,
1049 which is loaded from file ``.item_link``.
1051 entity2item (dict): Dict maps ``entity`` to ``item_id``,
1052 which is loaded from file ``.item_link``.
1054 Note:
1055 :attr:`entity_field` doesn't exist exactly. It's only a symbol,
1056 representing entity features.
1058 :attr:`ui_relation` is a special relation token, which is used to represent
1059 the interaction relation between users and items.
1060 """
1062 @property
1063 def auxiliary_entity_num(self):
1064 """Get the number of different tokens of auxiliary entities (not users nor items).
1066 Returns:
1067 int: Number of different tokens of auxiliary entities.
1068 """
1069 return self.entity_num - self.user_num - self.item_num
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.
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()
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]
1101 remained_inter = pd.Series(True, index=self.inter_feat.index)
1102 remained_inter &= self.inter_feat[self.iid_field].isin(self.item2entity.keys())
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]
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)
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)
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)
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)
1137 # feats are re-indexed for safe index dropping and while loop stop conditions
1138 self._reset_index()
1140 if remained_inter.all() and remained_kg.all():
1141 break
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 )
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)
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)
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
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
1178 return item2entity, entity2item, user2entity, entity2user
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)
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)
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)
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)
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)
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
1255 def _filter_kg_by_triple_num(self):
1256 """Filter by number of triples.
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.
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"])
1270 if entity_kg_num_interval is None and relation_kg_num_interval is None:
1271 return
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()
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})
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 )
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})
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 )
1302 ban_entities = item_ban_entities | user_ban_entities
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
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)
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)
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)
1329 def _create_ckg_source_target(self, form="numpy"):
1330 """Create base collaborative knowledge graph.
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
1340 uids = self.inter_feat[self.uid_field].numpy()
1341 iids = self.inter_feat[self.iid_field].numpy() + self.user_num
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
1349 hids = kg_tensor[self.head_entity_field]
1350 tids = kg_tensor[self.tail_entity_field]
1352 uids = inter_tensor[self.uid_field]
1353 iids = inter_tensor[self.iid_field] + self.user_num
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.")
1360 return src, tgt
1362 def _create_ckg_sparse_matrix(self, form="coo", show_relation=False):
1363 src, tgt = self._create_ckg_source_target(form="numpy")
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
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.")
1383 def _create_ckg_igraph(self, show_relation=False, directed=True):
1384 import igraph as ig
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
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
1408 uids = self.inter_feat[self.uid_field].numpy()
1409 iids = self.inter_feat[self.iid_field].numpy() + self.user_num
1411 src = np.concatenate([uids, hids])
1412 tgt = np.concatenate([iids, tids])
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 )
1422 return ig_graph