Coverage for hopwise/data/dataset/dataset.py: 78%
1009 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/6/28
2# @Author : Yupeng Hou
3# @Email : houyupeng@ruc.edu.cn
5# UPDATE:
6# @Time : 2022/7/8, 2021/12/18 2021/7/14 2021/7/1, 2020/11/10
7# @Author : Zhen Tian, Yupeng Hou, Xingyu Pan, Yushuo Chen, Juyong Jiang
8# @Email : chenyuwuxinn@gmail.com, houyupeng@ruc.edu.cn, xy_pan@foxmail.com, chenyushuo@ruc.edu.cn, csjuyongjiang@gmail.com # noqa: E501
10"""hopwise.data.dataset
11##########################
12"""
14import copy
15import os
16import pickle
17import sys
18from collections import Counter, defaultdict
19from logging import getLogger
21import numpy as np
22import pandas as pd
23import torch
24import torch.distributed as dist
25import torch.nn.utils.rnn as rnn_utils
26import yaml
27from scipy.sparse import coo_matrix, diags, dok_matrix
29from hopwise.data.interaction import Interaction
30from hopwise.utils import FeatureSource, FeatureType, ensure_dir, set_color
31from hopwise.utils.url import decide_download, download_url, extract_zip, makedirs, rename_atomic_files
34class Dataset(torch.utils.data.Dataset):
35 """:class:`Dataset` stores the original dataset in memory.
36 It provides many useful functions for data preprocessing, such as k-core data filtering and missing value
37 imputation. Features are stored as :class:`pandas.DataFrame` inside :class:`~hopwise.data.dataset.dataset.Dataset`.
38 General and Context-aware Models can use this class.
40 By calling method :meth:`~hopwise.data.dataset.dataset.Dataset.build`, it will processing dataset into
41 DataLoaders, according to :class:`~hopwise.config.eval_setting.EvalSetting`.
43 Args:
44 config (Config): Global configuration object.
46 Attributes:
47 dataset_name (str): Name of this dataset.
49 dataset_path (str): Local file path of this dataset.
51 field2type (dict): Dict mapping feature name (str) to its type (:class:`~hopwise.utils.enum_type.FeatureType`).
53 field2source (dict): Dict mapping feature name (str) to its source
54 (:class:`~hopwise.utils.enum_type.FeatureSource`).
55 Specially, if feature is loaded from Arg ``additional_feat_suffix``, its source has type str,
56 which is the suffix of its local file (also the suffix written in Arg ``additional_feat_suffix``).
58 field2id_token (dict): Dict mapping feature name (str) to a :class:`np.ndarray`, which stores the original token
59 of this feature. For example, if ``test`` is token-like feature, ``token_a`` is remapped to 1, ``token_b``
60 is remapped to 2. Then ``field2id_token['test'] = ['[PAD]', 'token_a', 'token_b']``. (Note that 0 is
61 always PADDING for token-like features.)
63 field2token_id (dict): Dict mapping feature name (str) to a dict, which stores the token remap table
64 of this feature. For example, if ``test`` is token-like feature, ``token_a`` is remapped to 1, ``token_b``
65 is remapped to 2. Then ``field2token_id['test'] = {'[PAD]': 0, 'token_a': 1, 'token_b': 2}``.
66 (Note that 0 is always PADDING for token-like features.)
68 field2seqlen (dict): Dict mapping feature name (str) to its sequence length (int).
69 For sequence features, their length can be either set in config,
70 or set to the max sequence length of this feature.
71 For token and float features, their length is 1.
73 uid_field (str or None): The same as ``config['USER_ID_FIELD']``.
75 iid_field (str or None): The same as ``config['ITEM_ID_FIELD']``.
77 label_field (str or None): The same as ``config['LABEL_FIELD']``.
79 time_field (str or None): The same as ``config['TIME_FIELD']``.
81 inter_feat (:class:`Interaction`): Internal data structure stores the interaction features.
82 It's loaded from file ``.inter``.
84 user_feat (:class:`Interaction` or None): Internal data structure stores the user features.
85 It's loaded from file ``.user`` if existed.
87 item_feat (:class:`Interaction` or None): Internal data structure stores the item features.
88 It's loaded from file ``.item`` if existed.
90 feat_name_list (list): A list contains all the features' name (:class:`str`), including additional features.
91 """ # noqa: E501
93 def __init__(self, config):
94 super().__init__()
95 self.config = config
96 self.dataset_name = config["dataset"]
97 self.logger = getLogger()
98 self._from_scratch()
100 def _from_scratch(self):
101 """Load dataset from scratch.
102 Initialize attributes firstly, then load data from atomic files, pre-process the dataset lastly.
103 """
104 self.logger.debug(set_color(f"Loading {self.__class__} from scratch.", "green"))
106 self._get_preset()
107 self._get_field_from_config()
108 self._load_data(self.dataset_name, self.dataset_path)
109 self._init_alias()
110 self._data_processing()
112 def _get_preset(self):
113 """Initialization useful inside attributes."""
114 self.dataset_path = self.config["data_path"]
116 self.field2type = {}
117 self.field2source = {}
118 self.field2id_token = {}
119 self.field2token_id = {}
120 self.field2bucketnum = {}
121 self.field2seqlen = {}
122 self.alias = {}
123 self._preloaded_weight = {}
124 self.benchmark_filename_list = self.config["benchmark_filename"]
126 def _get_field_from_config(self):
127 """Initialization common field names."""
128 self.uid_field = self.config["USER_ID_FIELD"]
129 self.iid_field = self.config["ITEM_ID_FIELD"]
130 self.label_field = self.config["LABEL_FIELD"]
131 self.time_field = self.config["TIME_FIELD"]
133 if (self.uid_field is None) ^ (self.iid_field is None):
134 raise ValueError(
135 "USER_ID_FIELD and ITEM_ID_FIELD need to be set at the same time or not set at the same time."
136 )
138 self.logger.debug(set_color("uid_field", "blue") + f": {self.uid_field}")
139 self.logger.debug(set_color("iid_field", "blue") + f": {self.iid_field}")
141 def _data_processing(self):
142 """Data preprocessing, including:
144 - Data filtering
145 - Remap ID
146 - Missing value imputation
147 - Normalization
148 - Preloading weights initialization
149 """
150 self.feat_name_list = self._build_feat_name_list()
151 if self.benchmark_filename_list is None:
152 self._data_filtering()
154 self._remap_ID_all()
155 self._user_item_feat_preparation()
156 self._fill_nan()
157 self._set_label_by_threshold()
158 self._normalize()
159 self._discretization()
160 self._preload_weight_matrix()
162 def _data_filtering(self):
163 """Data filtering
165 - Filter missing user_id or item_id
166 - Remove duplicated user-item interaction
167 - Value-based data filtering
168 - Remove interaction by user or item
169 - K-core data filtering
171 Note:
172 After filtering, feats(``DataFrame``) has non-continuous index,
173 thus :meth:`~hopwise.data.dataset.dataset.Dataset._reset_index` will reset the index of feats.
174 """
175 self._filter_nan_user_or_item()
176 self._remove_duplication()
177 self._filter_by_field_value()
178 self._filter_inter_by_user_or_item()
179 self._filter_by_inter_num()
180 self._reset_index()
182 def _build_feat_name_list(self):
183 """Feat list building.
185 Any feat loaded by Dataset can be found in ``feat_name_list``
187 Returns:
188 built feature name list.
190 Note:
191 Subclasses can inherit this method to add new feat.
192 """
193 feat_name_list = [
194 feat_name
195 for feat_name in ["inter_feat", "user_feat", "item_feat"]
196 if getattr(self, feat_name, None) is not None
197 ]
198 if self.config["additional_feat_suffix"] is not None:
199 for suf in self.config["additional_feat_suffix"]:
200 if getattr(self, f"{suf}_feat", None) is not None:
201 feat_name_list.append(f"{suf}_feat")
202 return feat_name_list
204 def _get_download_url(self, url_file, allow_none=False):
205 current_path = os.path.dirname(os.path.realpath(__file__))
206 with open(os.path.join(current_path, f"../../properties/dataset/{url_file}.yaml")) as f:
207 dataset2url = yaml.load(f.read(), Loader=self.config.yaml_loader)
209 if self.dataset_name in dataset2url:
210 url = dataset2url[self.dataset_name]
211 return url
212 elif allow_none:
213 return None
214 else:
215 raise ValueError(
216 f"Neither [{self.dataset_path}] exists in the device nor [{self.dataset_name}] a known dataset name."
217 )
219 def _download(self):
220 if self.config["local_rank"] == 0:
221 url = self._get_download_url("url")
222 self.logger.info(f"Prepare to download dataset [{self.dataset_name}] from [{url}].")
224 if decide_download(url):
225 makedirs(self.dataset_path)
226 path = download_url(url, self.dataset_path)
227 extract_zip(path, self.dataset_path)
228 os.unlink(path)
230 basename = os.path.splitext(os.path.basename(path))[0]
231 rename_atomic_files(self.dataset_path, basename, self.dataset_name)
233 self.logger.info("Downloading done.")
234 else:
235 self.logger.info("Stop download.")
236 sys.exit(-1)
237 if dist.is_available() and dist.is_initialized():
238 dist.barrier()
239 elif dist.is_available() and dist.is_initialized():
240 dist.barrier()
242 def _load_data(self, token, dataset_path):
243 """Load features.
245 Firstly load interaction features, then user/item features optionally,
246 finally load additional features if ``config['additional_feat_suffix']`` is set.
248 Args:
249 token (str): dataset name.
250 dataset_path (str): path of dataset dir.
251 """
252 if not os.path.exists(dataset_path):
253 self._download()
254 self._load_inter_feat(token, dataset_path)
255 self.user_feat = self._load_user_or_item_feat(token, dataset_path, FeatureSource.USER, "uid_field")
256 self.item_feat = self._load_user_or_item_feat(token, dataset_path, FeatureSource.ITEM, "iid_field")
257 self._load_additional_feat(token, dataset_path)
259 def _load_inter_feat(self, token, dataset_path):
260 """Load interaction features.
262 If ``config['benchmark_filename']`` is not set, load interaction features from ``.inter``.
264 Otherwise, load interaction features from a file list, named ``dataset_name.xxx.inter``,
265 where ``xxx`` if from ``config['benchmark_filename']``.
266 After loading, ``self.file_size_list`` stores the length of each interaction file.
268 Args:
269 token (str): dataset name.
270 dataset_path (str): path of dataset dir.
271 """
272 if self.benchmark_filename_list is None:
273 inter_feat_path = os.path.join(dataset_path, f"{token}.inter")
274 if not os.path.isfile(inter_feat_path):
275 raise ValueError(f"File {inter_feat_path} not exist.")
277 inter_feat = self._load_feat(inter_feat_path, FeatureSource.INTERACTION)
278 self.logger.debug(f"Interaction feature loaded successfully from [{inter_feat_path}].")
279 self.inter_feat = inter_feat
280 else:
281 sub_inter_lens = []
282 sub_inter_feats = []
283 overall_field2seqlen = defaultdict(int)
284 for filename in self.benchmark_filename_list:
285 file_path = os.path.join(dataset_path, f"{token}.{filename}.inter")
286 if os.path.isfile(file_path):
287 temp = self._load_feat(file_path, FeatureSource.INTERACTION)
288 sub_inter_feats.append(temp)
289 sub_inter_lens.append(len(temp))
290 for field in self.field2seqlen:
291 overall_field2seqlen[field] = max(overall_field2seqlen[field], self.field2seqlen[field])
292 else:
293 raise ValueError(f"File {file_path} not exist.")
294 inter_feat = pd.concat(sub_inter_feats, ignore_index=True)
295 self.inter_feat, self.file_size_list = inter_feat, sub_inter_lens
296 self.field2seqlen = overall_field2seqlen
298 def _load_user_or_item_feat(self, token, dataset_path, source, field_name):
299 """Load user/item features.
301 Args:
302 token (str): dataset name.
303 dataset_path (str): path of dataset dir.
304 source (FeatureSource): source of user/item feature.
305 field_name (str): ``uid_field`` or ``iid_field``
307 Returns:
308 pandas.DataFrame: Loaded feature
310 Note:
311 ``user_id`` and ``item_id`` has source :obj:`~hopwise.utils.enum_type.FeatureSource.USER_ID` and
312 :obj:`~hopwise.utils.enum_type.FeatureSource.ITEM_ID`
313 """
314 feat_path = os.path.join(dataset_path, f"{token}.{source.value}")
315 field = getattr(self, field_name, None)
317 if os.path.isfile(feat_path):
318 feat = self._load_feat(feat_path, source)
319 self.logger.debug(f"[{source.value}] feature loaded successfully from [{feat_path}].")
320 else:
321 feat = None
322 self.logger.debug(f"[{feat_path}] not found, [{source.value}] features are not loaded.")
324 if feat is not None and field is None:
325 raise ValueError(f"{field_name} must be exist if {source.value}_feat exist.")
326 if feat is not None and field not in feat:
327 raise ValueError(f"{field_name} must be loaded if {source.value}_feat is loaded.")
328 if feat is not None:
329 feat.drop_duplicates(subset=[field], keep="first", inplace=True)
331 if field in self.field2source:
332 self.field2source[field] = FeatureSource(source.value + "_id")
333 return feat
335 def _load_additional_feat(self, token, dataset_path):
336 """Load additional features.
338 For those additional features, e.g. pretrained entity embedding, user can set them
339 as ``config['additional_feat_suffix']``, then they will be loaded and stored in
340 :attr:`feat_name_list`. See :doc:`../user_guide/data/data_settings` for details.
341 If ``config['preload_weight']`` and ``config['preload_weight_path']`` are set,
342 those additional features will be loaded from ``config['preload_weight_path']``.
344 Args:
345 token (str): dataset name.
346 dataset_path (str): path of dataset dir.
347 """
348 if self.config["additional_feat_suffix"] is None:
349 return
350 for suf in self.config["additional_feat_suffix"]:
351 if hasattr(self, f"{suf}_feat"):
352 raise ValueError(f"{suf}_feat already exist.")
353 preload_fields = self.config["preload_weight"]
354 load_col = self.config["load_col"]
355 if preload_fields is not None and any(col in load_col[suf] for col in preload_fields):
356 feat_path = self.config["preload_weight_path"] or dataset_path
357 else:
358 feat_path = dataset_path
360 feat_path = os.path.join(feat_path, f"{token}.{suf}")
361 if os.path.isfile(feat_path):
362 feat = self._load_feat(feat_path, suf)
363 else:
364 raise ValueError(f"Additional feature file [{feat_path}] not found.")
365 setattr(self, f"{suf}_feat", feat)
367 def _get_load_and_unload_col(self, source):
368 """Parsing ``config['load_col']`` and ``config['unload_col']`` according to source.
369 See :doc:`../user_guide/config/data_settings` for detail arg setting.
371 Args:
372 source (FeatureSource): source of input file.
374 Returns:
375 tuple: tuple of parsed ``load_col`` and ``unload_col``, details on :doc:`../user_guide/data/data_args`.
376 """
377 if isinstance(source, FeatureSource):
378 source = source.value
379 if self.config["load_col"] is None:
380 load_col = None
381 elif source not in self.config["load_col"]:
382 load_col = set()
383 elif self.config["load_col"][source] == "*":
384 load_col = None
385 else:
386 load_col = set(self.config["load_col"][source])
388 if self.config["unload_col"] is not None and source in self.config["unload_col"]:
389 unload_col = set(self.config["unload_col"][source])
390 else:
391 unload_col = None
393 if load_col and unload_col:
394 raise ValueError(f"load_col [{load_col}] and unload_col [{unload_col}] can not be set the same time.")
396 self.logger.debug(set_color(f"[{source}]: ", "magenta"))
397 self.logger.debug(set_color("\t load_col", "blue") + f": [{load_col}]")
398 self.logger.debug(set_color("\t unload_col", "blue") + f": [{unload_col}]")
399 return load_col, unload_col
401 def _load_feat(self, filepath, source):
402 """Load features according to source into :class:`pandas.DataFrame`.
404 Set features' properties, e.g. type, source and length.
406 Args:
407 filepath (str): path of input file.
408 source (FeatureSource or str): source of input file.
410 Returns:
411 pandas.DataFrame: Loaded feature
413 Note:
414 For sequence features, ``seqlen`` will be loaded, but data in DataFrame will not be cut off.
415 Their length is limited only after calling :meth:`~_dict_to_interaction` or
416 :meth:`~_dataframe_to_interaction`
417 """
418 self.logger.debug(set_color(f"Loading feature from [{filepath}] (source: [{source}]).", "green"))
420 load_col, unload_col = self._get_load_and_unload_col(source)
421 if load_col == set():
422 return None
424 field_separator = self.config["field_separator"]
425 columns = []
426 usecols = []
427 dtype = {}
428 encoding = self.config["encoding"]
429 with open(filepath, encoding=encoding) as f:
430 head = f.readline()[:-1]
431 for field_type in head.split(field_separator):
432 field, ftype = field_type.split(":")
433 try:
434 ftype = FeatureType(ftype)
435 except ValueError:
436 raise ValueError(f"Type {ftype} from field {field} is not supported.")
437 if load_col is not None and field not in load_col:
438 continue
439 if unload_col is not None and field in unload_col:
440 continue
441 if isinstance(source, FeatureSource) or source not in ["link", "user_link", "item_link"]:
442 self.field2source[field] = source
443 self.field2type[field] = ftype
444 if not ftype.value.endswith("seq"):
445 self.field2seqlen[field] = 1
446 if "float" in ftype.value:
447 self.field2bucketnum[field] = 2
448 columns.append(field)
449 usecols.append(field_type)
450 dtype[field_type] = np.float64 if ftype == FeatureType.FLOAT else str
452 if len(columns) == 0:
453 self.logger.warning(f"No columns has been loaded from [{source}]")
454 return None
456 df = pd.read_csv(
457 filepath,
458 delimiter=field_separator,
459 usecols=usecols,
460 dtype=dtype,
461 encoding=encoding,
462 engine="python",
463 )
464 df.columns = columns
466 seq_separator = self.config["seq_separator"]
467 for field in columns:
468 ftype = self.field2type[field]
469 if not ftype.value.endswith("seq"):
470 continue
471 df[field] = df[field].fillna(value="")
472 if ftype == FeatureType.TOKEN_SEQ:
473 df[field] = [np.array(list(filter(None, _.split(seq_separator)))) for _ in df[field].values]
474 elif ftype == FeatureType.FLOAT_SEQ:
475 df[field] = [
476 np.array(list(map(float, filter(None, _.split(seq_separator))))) for _ in df[field].values
477 ]
478 max_seq_len = max(map(len, df[field].values))
479 if self.config["seq_len"] and field in self.config["seq_len"]:
480 seq_len = self.config["seq_len"][field]
481 df[field] = [seq[:seq_len] if len(seq) > seq_len else seq for seq in df[field].values]
482 self.field2seqlen[field] = min(seq_len, max_seq_len)
483 else:
484 self.field2seqlen[field] = max_seq_len
486 return df
488 def _set_alias(self, alias_name, default_value):
489 alias = self.config[f"alias_of_{alias_name}"] or []
490 alias = np.array(list(filter(None, default_value)) + alias)
491 _, idx = np.unique(alias, return_index=True)
492 self.alias[alias_name] = alias[np.sort(idx)]
494 def _init_alias(self):
495 """Set :attr:`alias_of_user_id` and :attr:`alias_of_item_id`. And set :attr:`_rest_fields`."""
496 self._set_alias("user_id", [self.uid_field])
497 self._set_alias("item_id", [self.iid_field])
499 for alias_name_1, alias_1 in self.alias.items():
500 for alias_name_2, alias_2 in self.alias.items():
501 if alias_name_1 != alias_name_2:
502 intersect = np.intersect1d(alias_1, alias_2, assume_unique=True)
503 if len(intersect) > 0:
504 raise ValueError(
505 f"`alias_of_{alias_name_1}` and `alias_of_{alias_name_2}` "
506 f"should not have the same field {list(intersect)}."
507 )
509 self._rest_fields = self.token_like_fields
510 for alias_name, alias in self.alias.items():
511 isin = np.isin(alias, self._rest_fields, assume_unique=True)
512 if isin.all() is False:
513 raise ValueError(
514 f"`alias_of_{alias_name}` should not contain non-token-like field {list(alias[~isin])}."
515 )
516 self._rest_fields = np.setdiff1d(self._rest_fields, alias, assume_unique=True)
518 def _user_item_feat_preparation(self):
519 """Sort :attr:`user_feat` and :attr:`item_feat` by ``user_id`` or ``item_id``.
520 Missing values will be filled later.
521 """
522 if self.user_feat is not None:
523 new_user_df = pd.DataFrame({self.uid_field: np.arange(self.user_num)})
524 self.user_feat = pd.merge(new_user_df, self.user_feat, on=self.uid_field, how="left")
525 self.logger.debug(set_color("ordering user features by user id.", "green"))
526 if self.item_feat is not None:
527 new_item_df = pd.DataFrame({self.iid_field: np.arange(self.item_num)})
528 self.item_feat = pd.merge(new_item_df, self.item_feat, on=self.iid_field, how="left")
529 self.logger.debug(set_color("ordering item features by item id.", "green"))
531 def _preload_weight_matrix(self):
532 """Transfer preload weight features into :class:`numpy.ndarray` with shape ``[id_token_length]``
533 or ``[id_token_length, seqlen]``. See :doc:`../user_guide/data/data_args` for detail arg setting.
534 """
536 preload_fields = self.config["preload_weight"]
537 if preload_fields is None:
538 return
540 self.logger.debug(f"Preload weight matrix for {preload_fields}.")
541 for preload_id_field, preload_value_field in preload_fields.items():
542 if preload_id_field not in self.field2source:
543 raise ValueError(f"Preload id field [{preload_id_field}] not exist.")
544 if preload_value_field not in self.field2source:
545 raise ValueError(f"Preload value field [{preload_value_field}] not exist.")
546 pid_source = self.field2source[preload_id_field]
547 pv_source = self.field2source[preload_value_field]
548 if pid_source != pv_source:
549 raise ValueError(
550 f"Preload id field [{preload_id_field}] is from source [{pid_source}],"
551 f"while preload value field [{preload_value_field}] is from source [{pv_source}], "
552 f"which should be the same."
553 )
555 id_ftype = self.field2type[preload_id_field]
556 value_ftype = self.field2type[preload_value_field]
557 if id_ftype != FeatureType.TOKEN:
558 raise ValueError(f"Preload id field [{preload_id_field}] should be type token, but is [{id_ftype}].")
559 if value_ftype not in {FeatureType.FLOAT, FeatureType.FLOAT_SEQ}:
560 self.logger.warning(
561 f"Field [{preload_value_field}] with type [{value_ftype}] is not `float` or `float_seq`, "
562 f"which will not be handled by preload matrix."
563 )
564 continue
565 token_num = self.num(preload_id_field)
566 feat = self.field2feats(preload_id_field)[0]
567 if value_ftype == FeatureType.FLOAT:
568 matrix = np.zeros(token_num)
569 matrix[feat[preload_id_field]] = feat[preload_value_field]
570 else:
571 max_len = self.field2seqlen[preload_value_field]
572 matrix = np.zeros((token_num, max_len))
573 preload_ids = feat[preload_id_field].values
574 preload_values = feat[preload_value_field].to_list()
575 for pid, prow in zip(preload_ids, preload_values):
576 length = len(prow)
577 if length <= max_len:
578 matrix[pid, :length] = prow
579 else:
580 matrix[pid] = prow[:max_len]
581 self._preloaded_weight[preload_id_field] = matrix
583 def _fill_nan(self):
584 """Missing value imputation.
586 For fields with type :obj:`~hopwise.utils.enum_type.FeatureType.TOKEN`, missing value will be filled by
587 ``[PAD]``, which indexed as 0.
589 For fields with type :obj:`~hopwise.utils.enum_type.FeatureType.FLOAT`, missing value will be filled by
590 the average of original data.
591 """
592 self.logger.debug(set_color("Filling nan", "green"))
594 for feat_name in self.feat_name_list:
595 feat = getattr(self, feat_name)
596 for field in feat:
597 ftype = self.field2type[field]
598 if ftype == FeatureType.TOKEN:
599 feat[field] = feat[field].fillna(value=0)
600 elif ftype == FeatureType.FLOAT:
601 feat[field] = feat[field].fillna(value=feat[field].mean())
602 else:
603 dtype = np.int64 if ftype == FeatureType.TOKEN_SEQ else float
604 feat[field] = feat[field].apply(lambda x: np.array([], dtype=dtype) if isinstance(x, float) else x)
606 def _normalize(self):
607 """Normalization if ``config['normalize_field']`` or ``config['normalize_all']`` is set.
608 See :doc:`../user_guide/data/data_args` for detail arg setting.
610 .. math::
611 x' = \frac{x - x_{min}}{x_{max} - x_{min}}
613 Note:
614 Only float-like fields can be normalized.
615 """
616 if self.config["normalize_field"] is not None and self.config["normalize_all"] is True:
617 raise ValueError("Normalize_field and normalize_all can't be set at the same time.")
619 if self.config["normalize_field"]:
620 fields = self.config["normalize_field"]
621 for field in fields:
622 if field not in self.field2type:
623 raise ValueError(f"Field [{field}] does not exist.")
624 ftype = self.field2type[field]
625 if ftype not in (FeatureType.FLOAT, FeatureType.FLOAT_SEQ):
626 self.logger.warning(f"{field} is not a FLOAT/FLOAT_SEQ feat, which will not be normalized.")
627 elif self.config["normalize_all"]:
628 fields = self.float_like_fields
629 else:
630 return
632 self.logger.debug(set_color("Normalized fields", "blue") + f": {fields}")
634 for field in fields:
635 for feat in self.field2feats(field):
637 def norm(arr):
638 mx, mn = max(arr), min(arr)
639 if mx == mn:
640 self.logger.warning(f"All the same value in [{field}] from [{feat}_feat].")
641 arr = 1.0
642 else:
643 arr = (arr - mn) / (mx - mn)
644 return arr
646 ftype = self.field2type[field]
647 if ftype == FeatureType.FLOAT:
648 feat[field] = norm(feat[field].values)
649 elif ftype == FeatureType.FLOAT_SEQ:
650 split_point = np.cumsum(feat[field].transform(len))[:-1]
651 feat[field] = np.split(norm(feat[field].agg(np.concatenate)), split_point)
653 def _discretization(self):
654 """Discretization if ``config['discretization']`` is set.
655 See :doc:`../user_guide/data/data_args` for detail arg setting.
657 Note:
658 Only float-like fields can be discretized.
659 """
660 dis_info = {}
662 if self.config["discretization"]:
663 dis_info = self.config["discretization"]
665 for field in dis_info.keys():
666 if field not in self.field2type:
667 raise ValueError(f"Field [{field}] does not exist.")
668 if field not in self.config["numerical_features"]:
669 raise ValueError(f"Field [{field}] must be a numerical feature")
670 ftype = self.field2type[field]
671 if ftype not in (FeatureType.FLOAT, FeatureType.FLOAT_SEQ):
672 self.logger.warning(f"{field} is not a FLOAT/FLOAT_SEQ feat, which will not be normalized.")
673 del dis_info[field]
675 self.logger.debug(set_color("Normalized fields", "blue") + f": {dis_info.keys()}")
677 for field in self.config["numerical_features"]:
678 if field in dis_info:
679 info = dis_info[field]
680 method = info["method"]
681 bucket = None
682 if method == "ED":
683 if "bucket" in info:
684 bucket = info["bucket"]
685 else:
686 raise ValueError("The number of buckets must be set when apply equal discretization.")
688 for feat in self.field2feats(field):
690 def disc(arr, method, bucket):
691 if method == "ED": # Equal Distance/Frequency Discretization.
692 lower, upper = min(arr), max(arr) + 1e-9
693 if upper != lower:
694 arr = np.floor((arr - lower) * bucket / (upper - lower) + 1)
695 else:
696 self.logger.warning(f"All the same value in [{field}] from [{feat}_feat].")
697 arr = np.ones_like(arr) * bucket
699 elif method == "LD": # Logarithm Discretization
700 mask = arr > 2 # noqa: PLR2004
701 x = np.floor(np.log(arr * mask + 1e-9) ** 2 + 1)
702 x = np.where(mask, x, arr)
703 _, arr = np.unique(x, return_inverse=True)
704 else:
705 raise ValueError(f"Method [{method}] does not exist.")
707 return arr, int(max(arr) + 1)
709 ftype = self.field2type[field]
710 if ftype == FeatureType.FLOAT:
711 res, self.field2bucketnum[field] = disc(feat[field].values, method, bucket)
712 ret = np.ones_like(res)
713 feat[field] = np.stack([ret, res], axis=-1).tolist()
714 elif ftype == FeatureType.FLOAT_SEQ:
715 split_point = np.cumsum(feat[field].transform(len))[:-1]
716 res, self.field2bucketnum[field] = disc(feat[field].agg(np.concatenate), method, bucket)
717 ret = np.ones_like(res)
718 res, ret = (
719 np.split(res, split_point),
720 np.split(ret, split_point),
721 )
722 feat[field] = list(zip(ret, res))
723 else:
724 for feat in self.field2feats(field):
725 ftype = self.field2type[field]
726 if ftype == FeatureType.FLOAT:
727 feat[field] = np.stack([feat[field], np.ones_like(feat[field])], axis=-1).tolist()
728 else:
729 split_point = np.cumsum(feat[field].transform(len))[:-1]
730 res = ret = feat[field].agg(np.concatenate)
731 res = np.ones_like(ret)
732 res, ret = (
733 np.split(res, split_point),
734 np.split(ret, split_point),
735 )
736 feat[field] = list(zip(ret, res))
738 def _filter_nan_user_or_item(self):
739 """Filter NaN user_id and item_id"""
740 for field, name in zip([self.uid_field, self.iid_field], ["user", "item"]):
741 feat = getattr(self, name + "_feat")
742 if feat is not None:
743 dropped_feat = feat.index[feat[field].isnull()]
744 if len(dropped_feat):
745 self.logger.warning(
746 f"In {name}_feat, line {list(dropped_feat + 2)}, {field} do not exist, so they will be removed." # noqa: E501
747 )
748 feat.drop(feat.index[dropped_feat], inplace=True)
749 if field is not None:
750 dropped_inter = self.inter_feat.index[self.inter_feat[field].isnull()]
751 if len(dropped_inter):
752 self.logger.warning(
753 f"In inter_feat, line {list(dropped_inter + 2)}, {field} do not exist, so they will be removed." # noqa: E501
754 )
755 self.inter_feat.drop(self.inter_feat.index[dropped_inter], inplace=True)
757 def _remove_duplication(self):
758 """Remove duplications in inter_feat.
760 If :attr:`self.config['rm_dup_inter']` is not ``None``, it will remove duplicated user-item interactions.
762 Note:
763 Before removing duplicated user-item interactions, if :attr:`time_field` existed, :attr:`inter_feat`
764 will be sorted by :attr:`time_field` in ascending order.
765 """
766 keep = self.config["rm_dup_inter"]
767 if keep is None:
768 return
769 self._check_field("uid_field", "iid_field")
771 if self.time_field in self.inter_feat:
772 self.inter_feat.sort_values(by=[self.time_field], ascending=True, inplace=True)
773 self.logger.info(
774 f"Records in original dataset have been sorted by value of [{self.time_field}] in ascending order."
775 )
776 else:
777 self.logger.warning(
778 f"Timestamp field has not been loaded or specified, "
779 f"thus strategy [{keep}] of duplication removal may be meaningless."
780 )
781 self.inter_feat.drop_duplicates(subset=[self.uid_field, self.iid_field], keep=keep, inplace=True)
783 def _filter_by_inter_num(self):
784 """Filter by number of interaction.
786 The interval of the number of interactions can be set, and only users/items whose number
787 of interactions is in the specified interval can be retained.
788 See :doc:`../user_guide/data/data_args` for detail arg setting.
790 Note:
791 Lower bound of the interval is also called k-core filtering, which means this method
792 will filter loops until all the users and items has at least k interactions.
793 """
794 if self.uid_field is None or self.iid_field is None:
795 return
797 user_inter_num_interval = self._parse_intervals_str(self.config["user_inter_num_interval"])
798 item_inter_num_interval = self._parse_intervals_str(self.config["item_inter_num_interval"])
800 if user_inter_num_interval is None and item_inter_num_interval is None:
801 return
803 user_inter_num = Counter(self.inter_feat[self.uid_field].values) if user_inter_num_interval else Counter()
804 item_inter_num = Counter(self.inter_feat[self.iid_field].values) if item_inter_num_interval else Counter()
806 while True:
807 ban_users = self._get_illegal_ids_by_inter_num(
808 field=self.uid_field,
809 feat=self.user_feat,
810 inter_num=user_inter_num,
811 inter_interval=user_inter_num_interval,
812 )
813 ban_items = self._get_illegal_ids_by_inter_num(
814 field=self.iid_field,
815 feat=self.item_feat,
816 inter_num=item_inter_num,
817 inter_interval=item_inter_num_interval,
818 )
820 if len(ban_users) == 0 and len(ban_items) == 0:
821 break
823 if self.user_feat is not None:
824 dropped_user = self.user_feat[self.uid_field].isin(ban_users)
825 self.user_feat.drop(self.user_feat.index[dropped_user], inplace=True)
827 if self.item_feat is not None:
828 dropped_item = self.item_feat[self.iid_field].isin(ban_items)
829 self.item_feat.drop(self.item_feat.index[dropped_item], inplace=True)
831 dropped_inter = pd.Series(False, index=self.inter_feat.index)
832 user_inter = self.inter_feat[self.uid_field]
833 item_inter = self.inter_feat[self.iid_field]
834 dropped_inter |= user_inter.isin(ban_users)
835 dropped_inter |= item_inter.isin(ban_items)
837 user_inter_num -= Counter(user_inter[dropped_inter].values)
838 item_inter_num -= Counter(item_inter[dropped_inter].values)
840 dropped_index = self.inter_feat.index[dropped_inter]
841 self.logger.debug(f"[{len(dropped_index)}] dropped interactions.")
842 self.inter_feat.drop(dropped_index, inplace=True)
844 def _get_illegal_ids_by_inter_num(self, field, feat, inter_num, inter_interval=None):
845 """Given inter feat, return illegal ids, whose inter num out of [min_num, max_num]
847 Args:
848 field (str): field name of user_id or item_id.
849 feat (pandas.DataFrame): interaction feature.
850 inter_num (Counter): interaction number counter.
851 inter_interval (list, optional): the allowed interval(s) of the number of interactions.
852 Defaults to ``None``.
854 Returns:
855 set: illegal ids, whose inter num out of inter_intervals.
856 """
857 self.logger.debug(
858 set_color("get_illegal_ids_by_inter_num", "blue") + f": field=[{field}], inter_interval=[{inter_interval}]"
859 )
861 if inter_interval is not None:
862 if len(inter_interval) > 1:
863 self.logger.warning("More than one interval of interaction number are given!")
865 ids = {id_ for id_ in inter_num if not self._within_intervals(inter_num[id_], inter_interval)}
867 if feat is not None:
868 min_num = inter_interval[0][1] if inter_interval else -1
869 for id_ in feat[field].values:
870 if inter_num[id_] < min_num:
871 ids.add(id_)
872 self.logger.debug(f"[{len(ids)}] illegal_ids_by_inter_num, field=[{field}]")
873 return ids
875 def _parse_intervals_str(self, intervals_str):
876 """Given string of intervals, return the list of endpoints tuple, where a tuple corresponds to an interval.
878 Args:
879 intervals_str (str): the string of intervals, such as "(0,1];[3,4)".
881 Returns:
882 list of endpoint tuple, such as [('(', 0, 1.0 , ']'), ('[', 3.0, 4.0 , ')')].
883 """
884 if intervals_str is None:
885 return None
887 endpoints = []
888 for endpoint_pair_str in str(intervals_str).split(";"):
889 endpoint_pair_str = endpoint_pair_str.strip() # noqa: PLW2901
890 left_bracket, right_bracket = endpoint_pair_str[0], endpoint_pair_str[-1]
891 endpoint_pair = endpoint_pair_str[1:-1].split(",")
892 if not (len(endpoint_pair) == 2 and left_bracket in ["(", "["] and right_bracket in [")", "]"]): # noqa: PLR2004
893 self.logger.warning(f"{endpoint_pair_str} is an illegal interval!")
894 continue
896 left_point, right_point = float(endpoint_pair[0]), float(endpoint_pair[1])
897 if left_point > right_point:
898 self.logger.warning(f"{endpoint_pair_str} is an illegal interval!")
900 endpoints.append((left_bracket, left_point, right_point, right_bracket))
901 return endpoints
903 def _within_intervals(self, num, intervals):
904 """Return Ture if the num is in the intervals.
906 Note:
907 return true when the intervals is None.
908 """
909 result = True
910 for i, (left_bracket, left_point, right_point, right_bracket) in enumerate(intervals):
911 temp_result = num >= left_point if left_bracket == "[" else num > left_point
912 temp_result &= num <= right_point if right_bracket == "]" else num < right_point
913 result = temp_result if i == 0 else result | temp_result
914 return result
916 def _filter_by_field_value(self):
917 """Filter features according to its values."""
918 val_intervals = {} if self.config["val_interval"] is None else self.config["val_interval"]
919 self.logger.debug(set_color("drop_by_value", "blue") + f": val={val_intervals}")
921 for field, interval in val_intervals.items():
922 if field not in self.field2type:
923 raise ValueError(f"Field [{field}] not defined in dataset.")
925 if self.field2type[field] in {FeatureType.FLOAT, FeatureType.FLOAT_SEQ}:
926 field_val_interval = self._parse_intervals_str(interval)
927 for feat in self.field2feats(field):
928 feat.drop(
929 feat.index[~self._within_intervals(feat[field].values, field_val_interval)],
930 inplace=True,
931 )
932 else: # token-like field
933 for feat in self.field2feats(field):
934 feat.drop(feat.index[~feat[field].isin(interval)], inplace=True)
936 def _reset_index(self):
937 """Reset index for all feats in :attr:`feat_name_list`."""
938 for feat_name in self.feat_name_list:
939 feat = getattr(self, feat_name)
940 if feat.empty:
941 raise ValueError("Some feat is empty, please check the filtering settings.")
942 feat.reset_index(drop=True, inplace=True)
944 def _del_col(self, feat, field):
945 """Delete columns
947 Args:
948 feat (pandas.DataFrame or Interaction): the feat contains field.
949 field (str): field name to be dropped.
950 """
951 self.logger.debug(f"Delete column [{field}].")
952 if isinstance(feat, Interaction):
953 feat.drop(column=field)
954 else:
955 feat.drop(columns=field, inplace=True)
956 for dct in [
957 self.field2id_token,
958 self.field2token_id,
959 self.field2seqlen,
960 self.field2source,
961 self.field2type,
962 ]:
963 if field in dct:
964 del dct[field]
966 def _filter_inter_by_user_or_item(self):
967 """Remove interaction in inter_feat which user or item is not in user_feat or item_feat."""
968 if self.config["filter_inter_by_user_or_item"] is not True:
969 return
971 remained_inter = pd.Series(True, index=self.inter_feat.index)
973 if self.user_feat is not None:
974 remained_uids = self.user_feat[self.uid_field].values
975 remained_inter &= self.inter_feat[self.uid_field].isin(remained_uids)
977 if self.item_feat is not None:
978 remained_iids = self.item_feat[self.iid_field].values
979 remained_inter &= self.inter_feat[self.iid_field].isin(remained_iids)
981 self.inter_feat.drop(self.inter_feat.index[~remained_inter], inplace=True)
983 def _set_label_by_threshold(self):
984 """Generate 0/1 labels according to value of features.
986 According to ``config['threshold']``, those rows with value lower than threshold will
987 be given negative label, while the other will be given positive label.
988 See :doc:`../user_guide/data/data_args` for detail arg setting.
990 Note:
991 Key of ``config['threshold']`` if a field name.
992 This field will be dropped after label generation.
993 """
994 threshold = self.config["threshold"]
995 if threshold is None:
996 return
998 self.logger.debug(f"Set label by {threshold}.")
1000 if len(threshold) != 1:
1001 raise ValueError("Threshold length should be 1.")
1003 self.set_field_property(self.label_field, FeatureType.FLOAT, FeatureSource.INTERACTION, 1)
1004 for field, value in threshold.items():
1005 if field in self.inter_feat:
1006 self.inter_feat[self.label_field] = (self.inter_feat[field] >= value).astype(int)
1007 else:
1008 raise ValueError(f"Field [{field}] not in inter_feat.")
1009 if field != self.label_field:
1010 self._del_col(self.inter_feat, field)
1012 def _get_remap_list(self, field_list):
1013 """Transfer set of fields in the same remapping space into remap list.
1015 If ``uid_field`` or ``iid_field`` in ``field_set``,
1016 field in :attr:`inter_feat` will be remapped firstly,
1017 then field in :attr:`user_feat` or :attr:`item_feat` will be remapped next, finally others.
1019 Args:
1020 field_list (numpy.ndarray): List of fields in the same remapping space.
1022 Returns:
1023 list: List of tuples (feat, field, ftype) where feat is a pandas.DataFrame,
1024 field is a str, and ftype is a FeatureType. They will be concatenated
1025 in order, and remapped together.
1026 """
1027 remap_list = []
1028 for field in field_list:
1029 ftype = self.field2type[field]
1030 for feat in self.field2feats(field):
1031 remap_list.append((feat, field, ftype))
1032 return remap_list
1034 def _remap_ID_all(self):
1035 """Remap all token-like fields."""
1036 for alias in self.alias.values():
1037 remap_list = self._get_remap_list(alias)
1038 self._remap(remap_list)
1040 for field in self._rest_fields:
1041 remap_list = self._get_remap_list(np.array([field]))
1042 self._remap(remap_list)
1044 def _concat_remaped_tokens(self, remap_list):
1045 """Given ``remap_list``, concatenate values in order.
1047 Args:
1048 remap_list (list): See :meth:`_get_remap_list` for detail.
1050 Returns:
1051 tuple: A tuple of (tokens, split_point) where tokens is the concatenated
1052 array and split_point contains indices to restore the original tokens.
1053 """
1054 tokens = []
1055 for feat, field, ftype in remap_list:
1056 if ftype == FeatureType.TOKEN:
1057 tokens.append(feat[field].values)
1058 elif ftype == FeatureType.TOKEN_SEQ:
1059 tokens.append(feat[field].agg(np.concatenate))
1060 split_point = np.cumsum(list(map(len, tokens)))[:-1]
1061 tokens = np.concatenate(tokens)
1062 return tokens, split_point
1064 def _remap(self, remap_list):
1065 """Remap tokens using :meth:`pandas.factorize`.
1067 Args:
1068 remap_list (list): See :meth:`_get_remap_list` for detail.
1069 """
1070 if len(remap_list) == 0:
1071 return
1072 tokens, split_point = self._concat_remaped_tokens(remap_list)
1073 new_ids_list, mp = pd.factorize(tokens)
1074 new_ids_list = np.split(new_ids_list + 1, split_point)
1075 mp = np.array(["[PAD]"] + list(mp))
1076 token_id = {t: i for i, t in enumerate(mp)}
1078 for (feat, field, ftype), new_ids in zip(remap_list, new_ids_list):
1079 if field not in self.field2id_token:
1080 self.field2id_token[field] = mp
1081 self.field2token_id[field] = token_id
1082 if ftype == FeatureType.TOKEN:
1083 feat[field] = new_ids
1084 elif ftype == FeatureType.TOKEN_SEQ:
1085 split_point = np.cumsum(feat[field].transform(len))[:-1]
1086 feat[field] = np.split(new_ids, split_point)
1088 def _change_feat_format(self):
1089 """Change feat format from :class:`pandas.DataFrame` to :class:`Interaction`."""
1090 for feat_name in self.feat_name_list:
1091 feat = getattr(self, feat_name)
1092 if not isinstance(feat, Interaction):
1093 setattr(self, feat_name, self._dataframe_to_interaction(feat))
1095 def num(self, field):
1096 """Given ``field``, for token-like fields, return the number of different tokens after remapping,
1097 for float-like fields, return ``1``.
1099 Args:
1100 field (str): field name to get token number.
1102 Returns:
1103 int: The number of different tokens (``1`` if ``field`` is a float-like field).
1104 """
1105 if field not in self.field2type:
1106 raise ValueError(f"Field [{field}] not defined in dataset.")
1108 if (
1109 self.field2type[field] in {FeatureType.FLOAT, FeatureType.FLOAT_SEQ}
1110 and field in self.config["numerical_features"]
1111 ):
1112 return self.field2bucketnum[field]
1113 elif self.field2type[field] not in {FeatureType.TOKEN, FeatureType.TOKEN_SEQ}:
1114 return self.field2seqlen[field]
1115 else:
1116 return len(self.field2id_token[field])
1118 def fields(self, ftype=None, source=None):
1119 """Given type and source of features, return all the field name of this type and source.
1120 If ``ftype == None``, the type of returned fields is not restricted.
1121 If ``source == None``, the source of returned fields is not restricted.
1123 Args:
1124 ftype (FeatureType, optional): Type of features. Defaults to ``None``.
1125 source (FeatureSource, optional): Source of features. Defaults to ``None``.
1127 Returns:
1128 list: List of field names.
1129 """
1130 ftype = set(ftype) if ftype is not None else set(FeatureType)
1131 source = set(source) if source is not None else set(FeatureSource)
1132 ret = []
1133 for field in self.field2type:
1134 tp = self.field2type[field]
1135 src = self.field2source[field]
1136 if tp in ftype and src in source:
1137 ret.append(field)
1138 return ret
1140 @property
1141 def float_like_fields(self):
1142 """Get fields of type :obj:`~hopwise.utils.enum_type.FeatureType.FLOAT` and
1143 :obj:`~hopwise.utils.enum_type.FeatureType.FLOAT_SEQ`.
1145 Returns:
1146 list: List of field names.
1147 """
1148 return self.fields(ftype=[FeatureType.FLOAT, FeatureType.FLOAT_SEQ])
1150 @property
1151 def token_like_fields(self):
1152 """Get fields of type :obj:`~hopwise.utils.enum_type.FeatureType.TOKEN` and
1153 :obj:`~hopwise.utils.enum_type.FeatureType.TOKEN_SEQ`.
1155 Returns:
1156 list: List of field names.
1157 """
1158 return self.fields(ftype=[FeatureType.TOKEN, FeatureType.TOKEN_SEQ])
1160 @property
1161 def seq_fields(self):
1162 """Get fields of type :obj:`~hopwise.utils.enum_type.FeatureType.TOKEN_SEQ` and
1163 :obj:`~hopwise.utils.enum_type.FeatureType.FLOAT_SEQ`.
1165 Returns:
1166 list: List of field names.
1167 """
1168 return self.fields(ftype=[FeatureType.FLOAT_SEQ, FeatureType.TOKEN_SEQ])
1170 @property
1171 def non_seq_fields(self):
1172 """Get fields of type :obj:`~hopwise.utils.enum_type.FeatureType.TOKEN` and
1173 :obj:`~hopwise.utils.enum_type.FeatureType.FLOAT`.
1175 Returns:
1176 list: List of field names.
1177 """
1178 return self.fields(ftype=[FeatureType.FLOAT, FeatureType.TOKEN])
1180 def set_field_property(self, field, field_type, field_source, field_seqlen):
1181 """Set a new field's properties.
1183 Args:
1184 field (str): Name of the new field.
1185 field_type (FeatureType): Type of the new field.
1186 field_source (FeatureSource): Source of the new field.
1187 field_seqlen (int): max length of the sequence in ``field``.
1188 ``1`` if ``field``'s type is not sequence-like.
1189 """
1190 self.field2type[field] = field_type
1191 self.field2source[field] = field_source
1192 self.field2seqlen[field] = field_seqlen
1194 def copy_field_property(self, dest_field, source_field):
1195 """Copy properties from ``dest_field`` towards ``source_field``.
1197 Args:
1198 dest_field (str): Destination field.
1199 source_field (str): Source field.
1200 """
1201 self.field2type[dest_field] = self.field2type[source_field]
1202 self.field2source[dest_field] = self.field2source[source_field]
1203 self.field2seqlen[dest_field] = self.field2seqlen[source_field]
1205 def field2feats(self, field):
1206 if field not in self.field2source:
1207 raise ValueError(f"Field [{field}] not defined in dataset.")
1208 if field == self.uid_field:
1209 feats = [self.inter_feat]
1210 if self.user_feat is not None:
1211 feats.append(self.user_feat)
1212 elif field == self.iid_field:
1213 feats = [self.inter_feat]
1214 if self.item_feat is not None:
1215 feats.append(self.item_feat)
1216 else:
1217 source = self.field2source[field]
1218 if not isinstance(source, str):
1219 source = source.value
1220 feats = [getattr(self, f"{source}_feat")]
1221 return feats
1223 def token2id(self, field, tokens):
1224 """Map external tokens to internal ids.
1226 Args:
1227 field (str): Field of external tokens.
1228 tokens (str, list or numpy.ndarray): External tokens.
1230 Returns:
1231 int or numpy.ndarray: The internal ids of external tokens.
1232 """
1233 if isinstance(tokens, str):
1234 if tokens in self.field2token_id[field]:
1235 return self.field2token_id[field][tokens]
1236 else:
1237 raise ValueError(f"token [{tokens}] is not existed in {field}")
1238 elif isinstance(tokens, (list, np.ndarray)):
1239 return np.array([self.token2id(field, token) for token in tokens])
1240 else:
1241 raise TypeError(f"The type of tokens [{tokens}] is not supported")
1243 def id2token(self, field, ids):
1244 """Map internal ids to external tokens.
1246 Args:
1247 field (str): Field of internal ids.
1248 ids (int, list, numpy.ndarray or torch.Tensor): Internal ids.
1250 Returns:
1251 str or numpy.ndarray: The external tokens of internal ids.
1252 """
1253 try:
1254 return self.field2id_token[field][ids]
1255 except IndexError:
1256 if isinstance(ids, list):
1257 raise ValueError(f"[{ids}] is not a one-dimensional list.")
1258 else:
1259 raise ValueError(f"[{ids}] is not a valid ids.")
1261 def counter(self, field):
1262 """Given ``field``, if it is a token field in ``inter_feat``,
1263 return the counter containing the occurrences times in ``inter_feat`` of different tokens,
1264 for other cases, raise ValueError.
1266 Args:
1267 field (str): field name to get token counter.
1269 Returns:
1270 Counter: The counter of different tokens.
1271 """
1272 if field not in self.inter_feat:
1273 raise ValueError(f"Field [{field}] is not defined in ``inter_feat``.")
1274 if self.field2type[field] == FeatureType.TOKEN:
1275 if isinstance(self.inter_feat, pd.DataFrame):
1276 return Counter(self.inter_feat[field].values)
1277 else:
1278 return Counter(self.inter_feat[field].numpy())
1279 else:
1280 raise ValueError(f"Field [{field}] is not a token field.")
1282 @property
1283 def user_counter(self):
1284 """Get the counter containing the occurrences times in ``inter_feat`` of different users.
1286 Returns:
1287 Counter: The counter of different users.
1288 """
1289 self._check_field("uid_field")
1290 return self.counter(self.uid_field)
1292 @property
1293 def item_counter(self):
1294 """Get the counter containing the occurrences times in ``inter_feat`` of different items.
1296 Returns:
1297 Counter: The counter of different items.
1298 """
1299 self._check_field("iid_field")
1300 return self.counter(self.iid_field)
1302 @property
1303 def user_num(self):
1304 """Get the number of different tokens of ``self.uid_field``.
1306 Returns:
1307 int: Number of different tokens of ``self.uid_field``.
1308 """
1309 self._check_field("uid_field")
1310 return self.num(self.uid_field)
1312 @property
1313 def item_num(self):
1314 """Get the number of different tokens of ``self.iid_field``.
1316 Returns:
1317 int: Number of different tokens of ``self.iid_field``.
1318 """
1319 self._check_field("iid_field")
1320 return self.num(self.iid_field)
1322 @property
1323 def inter_num(self):
1324 """Get the number of interaction records.
1326 Returns:
1327 int: Number of interaction records.
1328 """
1329 return len(self.inter_feat)
1331 @property
1332 def avg_actions_of_users(self):
1333 """Get the average number of users' interaction records.
1335 Returns:
1336 numpy.float64: Average number of users' interaction records.
1337 """
1338 if isinstance(self.inter_feat, pd.DataFrame):
1339 return np.mean(self.inter_feat.groupby(self.uid_field).size())
1340 else:
1341 return np.mean(list(Counter(self.inter_feat[self.uid_field].numpy()).values()))
1343 @property
1344 def avg_actions_of_items(self):
1345 """Get the average number of items' interaction records.
1347 Returns:
1348 numpy.float64: Average number of items' interaction records.
1349 """
1350 if isinstance(self.inter_feat, pd.DataFrame):
1351 return np.mean(self.inter_feat.groupby(self.iid_field).size())
1352 else:
1353 return np.mean(list(Counter(self.inter_feat[self.iid_field].numpy()).values()))
1355 @property
1356 def sparsity(self):
1357 """Get the sparsity of this dataset.
1359 Returns:
1360 float: Sparsity of this dataset.
1361 """
1362 return 1 - self.inter_num / self.user_num / self.item_num
1364 def _check_field(self, *field_names):
1365 """Given a name of attribute, check if it's exist.
1367 Args:
1368 *field_names (str): Fields to be checked.
1369 """
1370 for field_name in field_names:
1371 if getattr(self, field_name, None) is None:
1372 raise ValueError(f"{field_name} isn't set.")
1374 def join(self, df):
1375 """Given interaction feature, join user/item feature into it.
1377 Args:
1378 df (Interaction): Interaction feature to be joint.
1380 Returns:
1381 Interaction: Interaction feature after joining operation.
1382 """
1383 if self.user_feat is not None and self.uid_field in df:
1384 df.update(self.user_feat[df[self.uid_field]])
1385 if self.item_feat is not None and self.iid_field in df:
1386 df.update(self.item_feat[df[self.iid_field]])
1387 return df
1389 def __getitem__(self, index, join=True):
1390 df = self.inter_feat[index]
1391 return self.join(df) if join else df
1393 def __len__(self):
1394 return len(self.inter_feat)
1396 def __repr__(self):
1397 return self.__str__()
1399 def __str__(self):
1400 info = [set_color(self.dataset_name, "magenta")]
1401 if self.uid_field:
1402 info.extend(
1403 [
1404 set_color("The number of users", "blue") + f": {self.user_num}",
1405 set_color("Average actions of users", "blue") + f": {self.avg_actions_of_users}",
1406 ]
1407 )
1408 if self.iid_field:
1409 info.extend(
1410 [
1411 set_color("The number of items", "blue") + f": {self.item_num}",
1412 set_color("Average actions of items", "blue") + f": {self.avg_actions_of_items}",
1413 ]
1414 )
1415 info.append(set_color("The number of inters", "blue") + f": {self.inter_num}")
1416 if self.uid_field and self.iid_field:
1417 info.append(set_color("The sparsity of the dataset", "blue") + f": {self.sparsity * 100}%")
1418 info.append(set_color("Remain Fields", "blue") + f": {list(self.field2type)}")
1419 return "\n".join(info)
1421 def copy(self, new_inter_feat):
1422 """Given a new interaction feature, return a new :class:`Dataset` object,
1423 whose interaction feature is updated with ``new_inter_feat``, and all the other attributes the same.
1425 Args:
1426 new_inter_feat (Interaction): The new interaction feature need to be updated.
1428 Returns:
1429 :class:`~Dataset`: the new :class:`~Dataset` object, whose interaction feature has been updated.
1430 """
1431 nxt = copy.copy(self)
1432 nxt.inter_feat = new_inter_feat
1433 return nxt
1435 def _drop_unused_col(self):
1436 """Drop columns which are loaded for data preparation but not used in model."""
1437 unused_col = self.config["unused_col"]
1438 if unused_col is None:
1439 return
1441 for feat_name, unused_fields in unused_col.items():
1442 feat = getattr(self, feat_name + "_feat")
1443 for field in unused_fields:
1444 if field not in feat:
1445 self.logger.warning(
1446 f"Field [{field}] is not in [{feat_name}_feat], which can not be set in `unused_col`."
1447 )
1448 continue
1449 self._del_col(feat, field)
1451 def _grouped_index(self, group_by_list):
1452 index = {}
1453 for i, key in enumerate(group_by_list):
1454 if key not in index:
1455 index[key] = [i]
1456 else:
1457 index[key].append(i)
1458 return index.values()
1460 def _calcu_split_ids(self, tot, ratios):
1461 """Given split ratios, and total number, calculate the number of each part after splitting.
1463 Other than the first one, each part is rounded down.
1465 Args:
1466 tot (int): Total number.
1467 ratios (list): List of split ratios. No need to be normalized.
1469 Returns:
1470 list: Number of each part after splitting.
1471 """
1472 cnt = [int(ratios[i] * tot) for i in range(len(ratios))]
1473 cnt[0] = tot - sum(cnt[1:])
1474 for i in range(1, len(ratios)):
1475 if cnt[0] <= 1:
1476 break
1477 if 0 < ratios[-i] * tot < 1:
1478 cnt[-i] += 1
1479 cnt[0] -= 1
1480 split_ids = np.cumsum(cnt)[:-1]
1481 return list(split_ids)
1483 def split_by_ratio(self, ratios, group_by=None):
1484 """Split interaction records by ratios.
1486 Args:
1487 ratios (list): List of split ratios. No need to be normalized.
1488 group_by (str, optional): Field name that interaction records should grouped by before splitting.
1489 Defaults to ``None``
1491 Returns:
1492 list: List of :class:`~Dataset`, whose interaction features has been split.
1494 Note:
1495 Other than the first one, each part is rounded down.
1496 """
1497 self.logger.debug(f"split by ratios [{ratios}], group_by=[{group_by}]")
1498 tot_ratio = sum(ratios)
1499 ratios = [_ / tot_ratio for _ in ratios]
1501 if group_by is None:
1502 tot_cnt = self.__len__()
1503 split_ids = self._calcu_split_ids(tot=tot_cnt, ratios=ratios)
1504 next_index = [range(start, end) for start, end in zip([0] + split_ids, split_ids + [tot_cnt])]
1505 else:
1506 grouped_inter_feat_index = self._grouped_index(self.inter_feat[group_by].numpy())
1507 next_index = [[] for _ in range(len(ratios))]
1508 for grouped_index in grouped_inter_feat_index:
1509 tot_cnt = len(grouped_index)
1510 split_ids = self._calcu_split_ids(tot=tot_cnt, ratios=ratios)
1511 for index, start, end in zip(next_index, [0] + split_ids, split_ids + [tot_cnt]):
1512 index.extend(grouped_index[start:end])
1514 self._drop_unused_col()
1515 next_df = [self.inter_feat[index] for index in next_index]
1516 next_ds = [self.copy(_) for _ in next_df]
1517 return next_ds
1519 def _split_index_by_leave_one_out(self, grouped_index, leave_one_num):
1520 """Split indexes by strategy leave one out.
1522 Args:
1523 grouped_index (list of list of int): Index to be split.
1524 leave_one_num (int): Number of parts whose length is expected to be ``1``.
1526 Returns:
1527 list: List of index that has been split.
1528 """
1529 next_index = [[] for _ in range(leave_one_num + 1)]
1530 for grp_index in grouped_index:
1531 index = list(grp_index)
1532 tot_cnt = len(index)
1533 legal_leave_one_num = min(leave_one_num, tot_cnt - 1)
1534 pr = tot_cnt - legal_leave_one_num
1535 next_index[0].extend(index[:pr])
1536 for i in range(legal_leave_one_num):
1537 next_index[-legal_leave_one_num + i].append(index[pr])
1538 pr += 1
1539 return next_index
1541 def leave_one_out(self, group_by, leave_one_mode):
1542 """Split interaction records by leave one out strategy.
1544 Args:
1545 group_by (str): Field name that interaction records should grouped by before splitting.
1546 leave_one_mode (str): The way to leave one out. It can only take three values:
1547 'valid_and_test', 'valid_only' and 'test_only'.
1549 Returns:
1550 list: List of :class:`~Dataset`, whose interaction features has been split.
1551 """
1552 self.logger.debug(f"leave one out, group_by=[{group_by}], leave_one_mode=[{leave_one_mode}]")
1553 if group_by is None:
1554 raise ValueError("leave one out strategy require a group field")
1556 grouped_inter_feat_index = self._grouped_index(self.inter_feat[group_by].numpy())
1557 if leave_one_mode == "valid_and_test":
1558 next_index = self._split_index_by_leave_one_out(grouped_inter_feat_index, leave_one_num=2)
1559 elif leave_one_mode == "valid_only":
1560 next_index = self._split_index_by_leave_one_out(grouped_inter_feat_index, leave_one_num=1)
1561 next_index.append([])
1562 elif leave_one_mode == "test_only":
1563 next_index = self._split_index_by_leave_one_out(grouped_inter_feat_index, leave_one_num=1)
1564 next_index = [next_index[0], [], next_index[1]]
1565 else:
1566 raise NotImplementedError(f"The leave_one_mode [{leave_one_mode}] has not been implemented.")
1568 self._drop_unused_col()
1569 next_df = [self.inter_feat[index] for index in next_index]
1570 next_ds = [self.copy(_) for _ in next_df]
1571 return next_ds
1573 def shuffle(self):
1574 """Shuffle the interaction records inplace."""
1575 self.inter_feat.shuffle()
1577 def sort(self, by, ascending=True):
1578 """Sort the interaction records inplace.
1580 Args:
1581 by (str or list of str): Field that as the key in the sorting process.
1582 ascending (bool or list of bool, optional): Results are ascending if ``True``, otherwise descending.
1583 Defaults to ``True``
1584 """
1585 self.inter_feat.sort(by=by, ascending=ascending)
1587 def build(self):
1588 """Processing dataset according to evaluation setting, including Group, Order and Split.
1589 See :class:`~hopwise.config.eval_setting.EvalSetting` for details.
1591 Returns:
1592 list: List of built :class:`Dataset`.
1593 """
1594 self._change_feat_format()
1596 if self.benchmark_filename_list is not None:
1597 self._drop_unused_col()
1598 cumsum = list(np.cumsum(self.file_size_list))
1599 datasets = [self.copy(self.inter_feat[start:end]) for start, end in zip([0] + cumsum[:-1], cumsum)]
1600 return datasets
1602 # ordering
1603 ordering_args = self.config["eval_args"]["order"]
1604 if ordering_args == "RO":
1605 self.shuffle()
1606 elif ordering_args == "TO":
1607 self.sort(by=self.time_field)
1608 else:
1609 raise NotImplementedError(f"The ordering_method [{ordering_args}] has not been implemented.")
1611 # splitting & grouping
1612 split_args = self.config["eval_args"]["split"]
1613 if split_args is None:
1614 raise ValueError("The split_args in eval_args should not be None.")
1615 if not isinstance(split_args, dict):
1616 raise ValueError(f"The split_args [{split_args}] should be a dict.")
1618 split_mode = list(split_args.keys())[0]
1619 assert len(split_args.keys()) == 1
1620 group_by = self.config["eval_args"]["group_by"]
1621 if split_mode == "RS":
1622 if not isinstance(split_args["RS"], list):
1623 raise ValueError(f'The value of "RS" [{split_args}] should be a list.')
1624 if group_by is None or group_by.lower() == "none":
1625 datasets = self.split_by_ratio(split_args["RS"], group_by=None)
1626 elif group_by == "user":
1627 datasets = self.split_by_ratio(split_args["RS"], group_by=self.uid_field)
1628 else:
1629 raise NotImplementedError(f"The grouping method [{group_by}] has not been implemented.")
1630 elif split_mode == "LS":
1631 datasets = self.leave_one_out(group_by=self.uid_field, leave_one_mode=split_args["LS"])
1632 else:
1633 raise NotImplementedError(f"The splitting_method [{split_mode}] has not been implemented.")
1635 return datasets
1637 def save(self):
1638 """Saving this :class:`Dataset` object to :attr:`config['checkpoint_dir']`."""
1639 save_dir = self.config["checkpoint_dir"]
1640 ensure_dir(save_dir)
1641 file = os.path.join(save_dir, f"{self.config['dataset']}-{self.__class__.__name__}.pth")
1642 self.logger.info(set_color("Saving filtered dataset into ", "magenta") + f"[{file}]")
1643 with open(file, "wb") as f:
1644 pickle.dump(self, f)
1646 def get_user_feature(self):
1647 """Returns:
1648 Interaction: user features
1649 """
1650 if self.user_feat is None:
1651 self._check_field("uid_field")
1652 return Interaction({self.uid_field: torch.arange(self.user_num)})
1653 else:
1654 return self.user_feat
1656 def get_item_feature(self):
1657 """Returns:
1658 Interaction: item features
1659 """
1660 if self.item_feat is None:
1661 self._check_field("iid_field")
1662 return Interaction({self.iid_field: torch.arange(self.item_num)})
1663 else:
1664 return self.item_feat
1666 def _create_sparse_matrix(self, df_feat, source_field, target_field, form="coo", value_field=None):
1667 """Get sparse matrix that describe relations between two fields.
1669 Source and target should be token-like fields.
1671 Sparse matrix has shape (``self.num(source_field)``, ``self.num(target_field)``).
1673 For a row of <src, tgt>, ``matrix[src, tgt] = 1`` if ``value_field`` is ``None``,
1674 else ``matrix[src, tgt] = df_feat[value_field][src, tgt]``.
1676 Args:
1677 df_feat (Interaction): Feature where src and tgt exist.
1678 source_field (str): Source field
1679 target_field (str): Target field
1680 form (str, optional): Sparse matrix format. Defaults to ``coo``.
1681 value_field (str, optional): Data of sparse matrix, which should exist in ``df_feat``.
1682 Defaults to ``None``.
1684 Returns:
1685 scipy.sparse: Sparse matrix in form ``coo`` or ``csr``.
1686 """
1687 src = df_feat[source_field]
1688 tgt = df_feat[target_field]
1689 if value_field is None:
1690 data = np.ones(len(df_feat))
1691 else:
1692 if value_field not in df_feat:
1693 raise ValueError(f"Value_field [{value_field}] should be one of `df_feat`'s features.")
1694 data = df_feat[value_field]
1696 mat = coo_matrix((data, (src, tgt)), shape=(self.num(source_field), self.num(target_field)))
1698 if form == "coo":
1699 return mat
1700 elif form == "csr":
1701 return mat.tocsr()
1702 else:
1703 raise NotImplementedError(f"Sparse matrix format [{form}] has not been implemented.")
1705 def _create_graph(self, tensor_feat, source_field, target_field, form="pyg", value_field=None):
1706 """Get graph that describe relations between two fields.
1708 Source and target should be token-like fields.
1710 For an edge of <src, tgt>, ``graph[src, tgt] = 1`` if ``value_field`` is ``None``,
1711 else ``graph[src, tgt] = df_feat[value_field][src, tgt]``.
1713 Currently, we support graph in `PyG`_.
1715 Args:
1716 tensor_feat (Interaction): Feature where src and tgt exist.
1717 source_field (str): Source field
1718 target_field (str): Target field
1719 form (str, optional): Library of graph data structure. Defaults to ``pyg``.
1720 value_field (str, optional): edge attributes of graph, which should exist in ``df_feat``.
1721 Defaults to ``None``.
1723 Returns:
1724 Graph of relations.
1726 .. _PyG:
1727 https://github.com/rusty1s/pytorch_geometric
1728 """
1729 src = tensor_feat[source_field]
1730 tgt = tensor_feat[target_field]
1732 if form == "pyg":
1733 from torch_geometric.data import Data
1735 edge_attr = tensor_feat[value_field] if value_field else None
1736 graph = Data(edge_index=torch.stack([src, tgt]), edge_attr=edge_attr)
1737 return graph
1738 else:
1739 raise NotImplementedError(f"Graph format [{form}] has not been implemented.")
1741 def inter_matrix(self, form="coo", value_field=None):
1742 """Get sparse matrix that describe interactions between user_id and item_id.
1744 Sparse matrix has shape (user_num, item_num).
1746 For a row of <src, tgt>, ``matrix[src, tgt] = 1`` if ``value_field`` is ``None``,
1747 else ``matrix[src, tgt] = self.inter_feat[src, tgt]``.
1749 Args:
1750 form (str, optional): Sparse matrix format. Defaults to ``coo``.
1751 value_field (str, optional): Data of sparse matrix, which should exist in ``df_feat``.
1752 Defaults to ``None``.
1754 Returns:
1755 scipy.sparse: Sparse matrix in form ``coo`` or ``csr``.
1756 """
1757 if not self.uid_field or not self.iid_field:
1758 raise ValueError("dataset does not exist uid/iid, thus can not converted to sparse matrix.")
1759 return self._create_sparse_matrix(self.inter_feat, self.uid_field, self.iid_field, form, value_field)
1761 def _create_norm_adjacency_matrix(self, size=None, symmetric=True):
1762 r"""Get the normalized interaction matrix of users and items.
1764 Construct the square matrix from the training data and normalize it
1765 using the laplace matrix.
1767 Args:
1768 size (int, optional): Size of the normalized interaction matrix. Defaults to ``None``.
1769 If ``None``, the size is set to ``self.user_num + self.item_num``.
1770 symmetric (bool, optional): Whether to use symmetric normalization. Defaults to ``True``.
1771 If ``True``, uses symmetric normalization: ``A_hat = D^{-0.5} * A * D^{-0.5}``.
1772 If ``False``, uses left normalization: ``A_hat = D^{-1} * A``.
1773 Here ``A`` is the adjacency matrix and ``D`` is the diagonal degree matrix.
1775 Returns:
1776 Sparse tensor of the normalized interaction matrix.
1777 """
1778 if size is None:
1779 size = self.user_num + self.item_num
1781 # build adj matrix
1782 A = dok_matrix((size, size), dtype=np.float32)
1783 inter_M = self.inter_matrix(form="coo").astype(np.float32)
1784 inter_M_t = inter_M.transpose()
1785 data_dict = dict(zip(zip(inter_M.row, inter_M.col + self.user_num), [1] * inter_M.nnz))
1786 data_dict.update(
1787 dict(
1788 zip(
1789 zip(inter_M_t.row + self.user_num, inter_M_t.col),
1790 [1] * inter_M_t.nnz,
1791 )
1792 )
1793 )
1794 A._dict.update(data_dict)
1796 # norm adj matrix
1797 sumArr = (A > 0).sum(axis=1)
1798 diag = np.array(sumArr.flatten())[0] + 1e-7 # add epsilon to avoid divide by zero Warning
1799 if symmetric:
1800 diag = np.power(diag, -0.5)
1801 D = diags(diag)
1802 L = D @ A @ D
1803 else:
1804 diag = np.power(diag, -1)
1805 D = diags(diag)
1806 L = D @ A
1808 # convert norm_adj matrix to tensor
1809 L = coo_matrix(L)
1810 row = L.row
1811 col = L.col
1812 i = torch.LongTensor(np.array([row, col]))
1813 data = torch.FloatTensor(L.data)
1815 return torch.sparse.FloatTensor(i, data, torch.Size(L.shape))
1817 def _create_eye_matrix(self):
1818 r"""Construct the identity matrix with the size of item_num + user_num.
1820 Returns:
1821 Sparse tensor of the identity matrix. Shape of (item_num + user_num, item_num + user_num)
1822 """
1823 num = self.item_num + self.user_num
1824 i = torch.LongTensor([range(0, num), range(0, num)])
1825 val = torch.FloatTensor([1] * num)
1826 return torch.sparse.FloatTensor(i, val)
1828 def norm_adjacency_matrix(self, form="torch.sparse"):
1829 """Get the normalized adjacency matrix of users and items.
1831 Construct the square matrix from the training data and normalize it
1832 using the laplace matrix.
1834 .. math::
1835 A_{hat} = D^{-0.5} \times A \times D^{-0.5}
1837 Args:
1838 form (str, optional): Format of the normalized adjacency matrix. Defaults to ``torch.sparse``.
1840 Returns:
1841 torch.sparse.FloatTensor: Normalized adjacency matrix.
1843 Raises:
1844 NotImplementedError: If the format of the normalized adjacency matrix is not implemented.
1845 """
1846 if form == "torch.sparse":
1847 return self._create_norm_adjacency_matrix()
1848 else:
1849 raise NotImplementedError(f"Normalized adjacency matrix format [{form}] has not been implemented.")
1851 def eye_matrix(self, form="torch.sparse"):
1852 """Construct the identity matrix with the size of item_num + user_num.
1854 Args:
1855 form (str, optional): Format of the identity matrix. Defaults
1856 to ``torch.sparse``.
1858 Returns:
1859 torch.sparse.FloatTensor: Identity matrix.
1860 """
1861 if form == "torch.sparse":
1862 return self._create_eye_matrix()
1863 else:
1864 raise NotImplementedError(f"Identity matrix format [{form}] has not been implemented.")
1866 def _history_matrix(self, row, value_field=None, max_history_len=None):
1867 """Get dense matrix describe user/item's history interaction records.
1869 ``history_matrix[i]`` represents ``i``'s history interacted item_id.
1871 ``history_value[i]`` represents ``i``'s history interaction records' values.
1872 ``0`` if ``value_field = None``.
1874 ``history_len[i]`` represents number of ``i``'s history interaction records.
1876 ``0`` is used as padding.
1878 Args:
1879 row (str): ``user`` or ``item``.
1880 value_field (str, optional): Data of matrix, which should exist in ``self.inter_feat``.
1881 Defaults to ``None``.
1882 max_history_len (int): The maximum number of history interaction records.
1883 Defaults to ``None``.
1885 Returns:
1886 tuple:
1887 - History matrix (torch.Tensor): ``history_matrix`` described above.
1888 - History values matrix (torch.Tensor): ``history_value`` described above.
1889 - History length matrix (torch.Tensor): ``history_len`` described above.
1890 """
1891 self._check_field("uid_field", "iid_field")
1893 inter_feat = copy.deepcopy(self.inter_feat)
1894 inter_feat.shuffle()
1895 user_ids, item_ids = (
1896 inter_feat[self.uid_field].numpy(),
1897 inter_feat[self.iid_field].numpy(),
1898 )
1899 if value_field is None:
1900 values = np.ones(len(inter_feat))
1901 else:
1902 if value_field not in inter_feat:
1903 raise ValueError(f"Value_field [{value_field}] should be one of `inter_feat`'s features.")
1904 values = inter_feat[value_field].numpy()
1906 if row == "user":
1907 row_num, max_col_num = self.user_num, self.item_num
1908 row_ids, col_ids = user_ids, item_ids
1909 else:
1910 row_num, max_col_num = self.item_num, self.user_num
1911 row_ids, col_ids = item_ids, user_ids
1913 history_len = np.zeros(row_num, dtype=np.int64)
1914 for row_id in row_ids:
1915 history_len[row_id] += 1
1917 max_inter_num = np.max(history_len)
1918 if max_history_len is not None:
1919 col_num = min(max_history_len, max_inter_num)
1920 else:
1921 col_num = max_inter_num
1923 if col_num > max_col_num * 0.2:
1924 self.logger.warning(
1925 f"Max value of {row}'s history interaction records has reached "
1926 f"{col_num / max_col_num * 100}% of the total."
1927 )
1929 history_matrix = np.zeros((row_num, col_num), dtype=np.int64)
1930 history_value = np.zeros((row_num, col_num))
1931 history_len[:] = 0
1932 for row_id, value, col_id in zip(row_ids, values, col_ids):
1933 if history_len[row_id] >= col_num:
1934 continue
1935 history_matrix[row_id, history_len[row_id]] = col_id
1936 history_value[row_id, history_len[row_id]] = value
1937 history_len[row_id] += 1
1939 return (
1940 torch.LongTensor(history_matrix),
1941 torch.FloatTensor(history_value),
1942 torch.LongTensor(history_len),
1943 )
1945 def history_item_matrix(self, value_field=None, max_history_len=None):
1946 """Get dense matrix describe user's history interaction records.
1948 ``history_matrix[i]`` represents user ``i``'s history interacted item_id.
1950 ``history_value[i]`` represents user ``i``'s history interaction records' values,
1951 ``0`` if ``value_field = None``.
1953 ``history_len[i]`` represents number of user ``i``'s history interaction records.
1955 ``0`` is used as padding.
1957 Args:
1958 value_field (str, optional): Data of matrix, which should exist in ``self.inter_feat``.
1959 Defaults to ``None``.
1961 max_history_len (int): The maximum number of user's history interaction records.
1962 Defaults to ``None``.
1964 Returns:
1965 tuple:
1966 - History matrix (torch.Tensor): ``history_matrix`` described above.
1967 - History values matrix (torch.Tensor): ``history_value`` described above.
1968 - History length matrix (torch.Tensor): ``history_len`` described above.
1969 """
1970 return self._history_matrix(row="user", value_field=value_field, max_history_len=max_history_len)
1972 def history_user_matrix(self, value_field=None, max_history_len=None):
1973 """Get dense matrix describe item's history interaction records.
1975 ``history_matrix[i]`` represents item ``i``'s history interacted user_id.
1977 ``history_value[i]`` represents item ``i``'s history interaction records' values,
1978 ``0`` if ``value_field = None``.
1980 ``history_len[i]`` represents number of item ``i``'s history interaction records.
1982 ``0`` is used as padding.
1984 Args:
1985 value_field (str, optional): Data of matrix, which should exist in ``self.inter_feat``.
1986 Defaults to ``None``.
1988 max_history_len (int): The maximum number of item's history interaction records.
1989 Defaults to ``None``.
1991 Returns:
1992 tuple:
1993 - History matrix (torch.Tensor): ``history_matrix`` described above.
1994 - History values matrix (torch.Tensor): ``history_value`` described above.
1995 - History length matrix (torch.Tensor): ``history_len`` described above.
1996 """
1997 return self._history_matrix(row="item", value_field=value_field, max_history_len=max_history_len)
1999 def _get_used_ids(self, source_field, target_field):
2000 """Get used ids from the interaction features.
2002 Args:
2003 source_field (str, optional): Source field name.
2004 target_field (str, optional): Target field name.
2006 Returns:
2007 numpy.ndarray: A numpy array of sets, where each set contains the item ids
2008 that a user has interacted with.
2009 """
2010 if isinstance(self.inter_feat, pd.DataFrame):
2011 source_values = self.inter_feat[source_field].values
2012 target_values = self.inter_feat[target_field].values
2013 else:
2014 source_values = self.inter_feat[source_field].numpy()
2015 target_values = self.inter_feat[target_field].numpy()
2017 cur = np.array([set() for _ in range(self.user_num)])
2018 for source, target in zip(source_values, target_values):
2019 cur[source].add(target)
2020 return cur
2022 def get_user_used_ids(self):
2023 """Get used item ids of each user from the interaction features.
2025 Returns:
2026 numpy.ndarray: A numpy array of sets, where each set contains the item ids
2027 that a user has interacted with.
2028 """
2029 self._check_field("uid_field", "iid_field")
2031 return self._get_used_ids(self.uid_field, self.iid_field)
2033 def get_item_used_ids(self):
2034 """Get used user ids of each item from the interaction features.
2036 Returns:
2037 numpy.ndarray: A numpy array of sets, where each set contains the user ids
2038 that have interacted with an item.
2039 """
2040 self._check_field("uid_field", "iid_field")
2042 return self._get_used_ids(self.iid_field, self.uid_field)
2044 def get_preload_weight(self, field):
2045 """Get preloaded weight matrix, whose rows are sorted by token ids.
2047 ``0`` is used as padding.
2049 Args:
2050 field (str): preloaded feature field name.
2052 Returns:
2053 numpy.ndarray: preloaded weight matrix. See :doc:`../user_guide/config/data_settings` for details.
2054 """
2055 if field not in self._preloaded_weight:
2056 raise ValueError(f"Field [{field}] not in preload_weight")
2057 return self._preloaded_weight[field]
2059 def _dataframe_to_interaction(self, data):
2060 """Convert :class:`pandas.DataFrame` to :class:`~hopwise.data.interaction.Interaction`.
2062 Args:
2063 data (pandas.DataFrame): data to be converted.
2065 Returns:
2066 :class:`~hopwise.data.interaction.Interaction`: Converted data.
2067 """
2068 new_data = {}
2069 for k in data:
2070 value = data[k].values
2071 ftype = self.field2type[k]
2072 if ftype == FeatureType.TOKEN:
2073 new_data[k] = torch.LongTensor(value)
2074 elif ftype == FeatureType.FLOAT:
2075 if k in self.config["numerical_features"]:
2076 new_data[k] = torch.FloatTensor(value.tolist())
2077 else:
2078 new_data[k] = torch.FloatTensor(value)
2079 elif ftype == FeatureType.TOKEN_SEQ:
2080 seq_data = [torch.LongTensor(d[: self.field2seqlen[k]]) for d in value]
2081 new_data[k] = rnn_utils.pad_sequence(seq_data, batch_first=True)
2082 elif ftype == FeatureType.FLOAT_SEQ:
2083 if k in self.config["numerical_features"]:
2084 base = [torch.FloatTensor(d[0][: self.field2seqlen[k]]) for d in value]
2085 base = rnn_utils.pad_sequence(base, batch_first=True)
2086 index = [torch.FloatTensor(d[1][: self.field2seqlen[k]]) for d in value]
2087 index = rnn_utils.pad_sequence(index, batch_first=True)
2088 new_data[k] = torch.stack([base, index], dim=-1)
2089 else:
2090 seq_data = [torch.FloatTensor(d[: self.field2seqlen[k]]) for d in value]
2091 new_data[k] = rnn_utils.pad_sequence(seq_data, batch_first=True)
2092 return Interaction(new_data)