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

1# @Time : 2020/6/28 

2# @Author : Yupeng Hou 

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

4 

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 

9 

10"""hopwise.data.dataset 

11########################## 

12""" 

13 

14import copy 

15import os 

16import pickle 

17import sys 

18from collections import Counter, defaultdict 

19from logging import getLogger 

20 

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 

28 

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 

32 

33 

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. 

39 

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`. 

42 

43 Args: 

44 config (Config): Global configuration object. 

45 

46 Attributes: 

47 dataset_name (str): Name of this dataset. 

48 

49 dataset_path (str): Local file path of this dataset. 

50 

51 field2type (dict): Dict mapping feature name (str) to its type (:class:`~hopwise.utils.enum_type.FeatureType`). 

52 

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

57 

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

62 

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

67 

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. 

72 

73 uid_field (str or None): The same as ``config['USER_ID_FIELD']``. 

74 

75 iid_field (str or None): The same as ``config['ITEM_ID_FIELD']``. 

76 

77 label_field (str or None): The same as ``config['LABEL_FIELD']``. 

78 

79 time_field (str or None): The same as ``config['TIME_FIELD']``. 

80 

81 inter_feat (:class:`Interaction`): Internal data structure stores the interaction features. 

82 It's loaded from file ``.inter``. 

83 

84 user_feat (:class:`Interaction` or None): Internal data structure stores the user features. 

85 It's loaded from file ``.user`` if existed. 

86 

87 item_feat (:class:`Interaction` or None): Internal data structure stores the item features. 

88 It's loaded from file ``.item`` if existed. 

89 

90 feat_name_list (list): A list contains all the features' name (:class:`str`), including additional features. 

91 """ # noqa: E501 

92 

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

99 

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

105 

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

111 

112 def _get_preset(self): 

113 """Initialization useful inside attributes.""" 

114 self.dataset_path = self.config["data_path"] 

115 

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

125 

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

132 

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 ) 

137 

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

140 

141 def _data_processing(self): 

142 """Data preprocessing, including: 

143 

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

153 

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

161 

162 def _data_filtering(self): 

163 """Data filtering 

164 

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 

170 

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

181 

182 def _build_feat_name_list(self): 

183 """Feat list building. 

184 

185 Any feat loaded by Dataset can be found in ``feat_name_list`` 

186 

187 Returns: 

188 built feature name list. 

189 

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 

203 

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) 

208 

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 ) 

218 

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}].") 

223 

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) 

229 

230 basename = os.path.splitext(os.path.basename(path))[0] 

231 rename_atomic_files(self.dataset_path, basename, self.dataset_name) 

232 

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

241 

242 def _load_data(self, token, dataset_path): 

243 """Load features. 

244 

245 Firstly load interaction features, then user/item features optionally, 

246 finally load additional features if ``config['additional_feat_suffix']`` is set. 

247 

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) 

258 

259 def _load_inter_feat(self, token, dataset_path): 

260 """Load interaction features. 

261 

262 If ``config['benchmark_filename']`` is not set, load interaction features from ``.inter``. 

263 

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. 

267 

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

276 

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 

297 

298 def _load_user_or_item_feat(self, token, dataset_path, source, field_name): 

299 """Load user/item features. 

300 

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

306 

307 Returns: 

308 pandas.DataFrame: Loaded feature 

309 

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) 

316 

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

323 

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) 

330 

331 if field in self.field2source: 

332 self.field2source[field] = FeatureSource(source.value + "_id") 

333 return feat 

334 

335 def _load_additional_feat(self, token, dataset_path): 

336 """Load additional features. 

337 

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']``. 

343 

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 

359 

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) 

366 

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. 

370 

371 Args: 

372 source (FeatureSource): source of input file. 

373 

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

387 

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 

392 

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

395 

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 

400 

401 def _load_feat(self, filepath, source): 

402 """Load features according to source into :class:`pandas.DataFrame`. 

403 

404 Set features' properties, e.g. type, source and length. 

405 

406 Args: 

407 filepath (str): path of input file. 

408 source (FeatureSource or str): source of input file. 

409 

410 Returns: 

411 pandas.DataFrame: Loaded feature 

412 

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

419 

420 load_col, unload_col = self._get_load_and_unload_col(source) 

421 if load_col == set(): 

422 return None 

423 

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 

451 

452 if len(columns) == 0: 

453 self.logger.warning(f"No columns has been loaded from [{source}]") 

454 return None 

455 

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 

465 

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 

485 

486 return df 

487 

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

493 

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

498 

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 ) 

508 

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) 

517 

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

530 

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

535 

536 preload_fields = self.config["preload_weight"] 

537 if preload_fields is None: 

538 return 

539 

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 ) 

554 

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 

582 

583 def _fill_nan(self): 

584 """Missing value imputation. 

585 

586 For fields with type :obj:`~hopwise.utils.enum_type.FeatureType.TOKEN`, missing value will be filled by 

587 ``[PAD]``, which indexed as 0. 

588 

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

593 

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) 

605 

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. 

609 

610 .. math:: 

611 x' = \frac{x - x_{min}}{x_{max} - x_{min}} 

612 

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

618 

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 

631 

632 self.logger.debug(set_color("Normalized fields", "blue") + f": {fields}") 

633 

634 for field in fields: 

635 for feat in self.field2feats(field): 

636 

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 

645 

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) 

652 

653 def _discretization(self): 

654 """Discretization if ``config['discretization']`` is set. 

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

656 

657 Note: 

658 Only float-like fields can be discretized. 

659 """ 

660 dis_info = {} 

661 

662 if self.config["discretization"]: 

663 dis_info = self.config["discretization"] 

664 

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] 

674 

675 self.logger.debug(set_color("Normalized fields", "blue") + f": {dis_info.keys()}") 

676 

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

687 

688 for feat in self.field2feats(field): 

689 

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 

698 

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

706 

707 return arr, int(max(arr) + 1) 

708 

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

737 

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) 

756 

757 def _remove_duplication(self): 

758 """Remove duplications in inter_feat. 

759 

760 If :attr:`self.config['rm_dup_inter']` is not ``None``, it will remove duplicated user-item interactions. 

761 

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

770 

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) 

782 

783 def _filter_by_inter_num(self): 

784 """Filter by number of interaction. 

785 

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. 

789 

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 

796 

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

799 

800 if user_inter_num_interval is None and item_inter_num_interval is None: 

801 return 

802 

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

805 

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 ) 

819 

820 if len(ban_users) == 0 and len(ban_items) == 0: 

821 break 

822 

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) 

826 

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) 

830 

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) 

836 

837 user_inter_num -= Counter(user_inter[dropped_inter].values) 

838 item_inter_num -= Counter(item_inter[dropped_inter].values) 

839 

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) 

843 

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] 

846 

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``. 

853 

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 ) 

860 

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

864 

865 ids = {id_ for id_ in inter_num if not self._within_intervals(inter_num[id_], inter_interval)} 

866 

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 

874 

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. 

877 

878 Args: 

879 intervals_str (str): the string of intervals, such as "(0,1];[3,4)". 

880 

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 

886 

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 

895 

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

899 

900 endpoints.append((left_bracket, left_point, right_point, right_bracket)) 

901 return endpoints 

902 

903 def _within_intervals(self, num, intervals): 

904 """Return Ture if the num is in the intervals. 

905 

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 

915 

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

920 

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

924 

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) 

935 

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) 

943 

944 def _del_col(self, feat, field): 

945 """Delete columns 

946 

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] 

965 

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 

970 

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

972 

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) 

976 

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) 

980 

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

982 

983 def _set_label_by_threshold(self): 

984 """Generate 0/1 labels according to value of features. 

985 

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. 

989 

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 

997 

998 self.logger.debug(f"Set label by {threshold}.") 

999 

1000 if len(threshold) != 1: 

1001 raise ValueError("Threshold length should be 1.") 

1002 

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) 

1011 

1012 def _get_remap_list(self, field_list): 

1013 """Transfer set of fields in the same remapping space into remap list. 

1014 

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. 

1018 

1019 Args: 

1020 field_list (numpy.ndarray): List of fields in the same remapping space. 

1021 

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 

1033 

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) 

1039 

1040 for field in self._rest_fields: 

1041 remap_list = self._get_remap_list(np.array([field])) 

1042 self._remap(remap_list) 

1043 

1044 def _concat_remaped_tokens(self, remap_list): 

1045 """Given ``remap_list``, concatenate values in order. 

1046 

1047 Args: 

1048 remap_list (list): See :meth:`_get_remap_list` for detail. 

1049 

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 

1063 

1064 def _remap(self, remap_list): 

1065 """Remap tokens using :meth:`pandas.factorize`. 

1066 

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

1077 

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) 

1087 

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

1094 

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``. 

1098 

1099 Args: 

1100 field (str): field name to get token number. 

1101 

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

1107 

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

1117 

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. 

1122 

1123 Args: 

1124 ftype (FeatureType, optional): Type of features. Defaults to ``None``. 

1125 source (FeatureSource, optional): Source of features. Defaults to ``None``. 

1126 

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 

1139 

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`. 

1144 

1145 Returns: 

1146 list: List of field names. 

1147 """ 

1148 return self.fields(ftype=[FeatureType.FLOAT, FeatureType.FLOAT_SEQ]) 

1149 

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`. 

1154 

1155 Returns: 

1156 list: List of field names. 

1157 """ 

1158 return self.fields(ftype=[FeatureType.TOKEN, FeatureType.TOKEN_SEQ]) 

1159 

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`. 

1164 

1165 Returns: 

1166 list: List of field names. 

1167 """ 

1168 return self.fields(ftype=[FeatureType.FLOAT_SEQ, FeatureType.TOKEN_SEQ]) 

1169 

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`. 

1174 

1175 Returns: 

1176 list: List of field names. 

1177 """ 

1178 return self.fields(ftype=[FeatureType.FLOAT, FeatureType.TOKEN]) 

1179 

1180 def set_field_property(self, field, field_type, field_source, field_seqlen): 

1181 """Set a new field's properties. 

1182 

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 

1193 

1194 def copy_field_property(self, dest_field, source_field): 

1195 """Copy properties from ``dest_field`` towards ``source_field``. 

1196 

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] 

1204 

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 

1222 

1223 def token2id(self, field, tokens): 

1224 """Map external tokens to internal ids. 

1225 

1226 Args: 

1227 field (str): Field of external tokens. 

1228 tokens (str, list or numpy.ndarray): External tokens. 

1229 

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

1242 

1243 def id2token(self, field, ids): 

1244 """Map internal ids to external tokens. 

1245 

1246 Args: 

1247 field (str): Field of internal ids. 

1248 ids (int, list, numpy.ndarray or torch.Tensor): Internal ids. 

1249 

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

1260 

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. 

1265 

1266 Args: 

1267 field (str): field name to get token counter. 

1268 

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

1281 

1282 @property 

1283 def user_counter(self): 

1284 """Get the counter containing the occurrences times in ``inter_feat`` of different users. 

1285 

1286 Returns: 

1287 Counter: The counter of different users. 

1288 """ 

1289 self._check_field("uid_field") 

1290 return self.counter(self.uid_field) 

1291 

1292 @property 

1293 def item_counter(self): 

1294 """Get the counter containing the occurrences times in ``inter_feat`` of different items. 

1295 

1296 Returns: 

1297 Counter: The counter of different items. 

1298 """ 

1299 self._check_field("iid_field") 

1300 return self.counter(self.iid_field) 

1301 

1302 @property 

1303 def user_num(self): 

1304 """Get the number of different tokens of ``self.uid_field``. 

1305 

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) 

1311 

1312 @property 

1313 def item_num(self): 

1314 """Get the number of different tokens of ``self.iid_field``. 

1315 

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) 

1321 

1322 @property 

1323 def inter_num(self): 

1324 """Get the number of interaction records. 

1325 

1326 Returns: 

1327 int: Number of interaction records. 

1328 """ 

1329 return len(self.inter_feat) 

1330 

1331 @property 

1332 def avg_actions_of_users(self): 

1333 """Get the average number of users' interaction records. 

1334 

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

1342 

1343 @property 

1344 def avg_actions_of_items(self): 

1345 """Get the average number of items' interaction records. 

1346 

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

1354 

1355 @property 

1356 def sparsity(self): 

1357 """Get the sparsity of this dataset. 

1358 

1359 Returns: 

1360 float: Sparsity of this dataset. 

1361 """ 

1362 return 1 - self.inter_num / self.user_num / self.item_num 

1363 

1364 def _check_field(self, *field_names): 

1365 """Given a name of attribute, check if it's exist. 

1366 

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

1373 

1374 def join(self, df): 

1375 """Given interaction feature, join user/item feature into it. 

1376 

1377 Args: 

1378 df (Interaction): Interaction feature to be joint. 

1379 

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 

1388 

1389 def __getitem__(self, index, join=True): 

1390 df = self.inter_feat[index] 

1391 return self.join(df) if join else df 

1392 

1393 def __len__(self): 

1394 return len(self.inter_feat) 

1395 

1396 def __repr__(self): 

1397 return self.__str__() 

1398 

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) 

1420 

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. 

1424 

1425 Args: 

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

1427 

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 

1434 

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 

1440 

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) 

1450 

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

1459 

1460 def _calcu_split_ids(self, tot, ratios): 

1461 """Given split ratios, and total number, calculate the number of each part after splitting. 

1462 

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

1464 

1465 Args: 

1466 tot (int): Total number. 

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

1468 

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) 

1482 

1483 def split_by_ratio(self, ratios, group_by=None): 

1484 """Split interaction records by ratios. 

1485 

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

1490 

1491 Returns: 

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

1493 

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] 

1500 

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

1513 

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 

1518 

1519 def _split_index_by_leave_one_out(self, grouped_index, leave_one_num): 

1520 """Split indexes by strategy leave one out. 

1521 

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``. 

1525 

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 

1540 

1541 def leave_one_out(self, group_by, leave_one_mode): 

1542 """Split interaction records by leave one out strategy. 

1543 

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'. 

1548 

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

1555 

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

1567 

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 

1572 

1573 def shuffle(self): 

1574 """Shuffle the interaction records inplace.""" 

1575 self.inter_feat.shuffle() 

1576 

1577 def sort(self, by, ascending=True): 

1578 """Sort the interaction records inplace. 

1579 

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) 

1586 

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. 

1590 

1591 Returns: 

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

1593 """ 

1594 self._change_feat_format() 

1595 

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 

1601 

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

1610 

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

1617 

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

1634 

1635 return datasets 

1636 

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) 

1645 

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 

1655 

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 

1665 

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. 

1668 

1669 Source and target should be token-like fields. 

1670 

1671 Sparse matrix has shape (``self.num(source_field)``, ``self.num(target_field)``). 

1672 

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]``. 

1675 

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``. 

1683 

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] 

1695 

1696 mat = coo_matrix((data, (src, tgt)), shape=(self.num(source_field), self.num(target_field))) 

1697 

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

1704 

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. 

1707 

1708 Source and target should be token-like fields. 

1709 

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]``. 

1712 

1713 Currently, we support graph in `PyG`_. 

1714 

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``. 

1722 

1723 Returns: 

1724 Graph of relations. 

1725 

1726 .. _PyG: 

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

1728 """ 

1729 src = tensor_feat[source_field] 

1730 tgt = tensor_feat[target_field] 

1731 

1732 if form == "pyg": 

1733 from torch_geometric.data import Data 

1734 

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

1740 

1741 def inter_matrix(self, form="coo", value_field=None): 

1742 """Get sparse matrix that describe interactions between user_id and item_id. 

1743 

1744 Sparse matrix has shape (user_num, item_num). 

1745 

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]``. 

1748 

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``. 

1753 

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) 

1760 

1761 def _create_norm_adjacency_matrix(self, size=None, symmetric=True): 

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

1763 

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

1765 using the laplace matrix. 

1766 

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. 

1774 

1775 Returns: 

1776 Sparse tensor of the normalized interaction matrix. 

1777 """ 

1778 if size is None: 

1779 size = self.user_num + self.item_num 

1780 

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) 

1795 

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 

1807 

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) 

1814 

1815 return torch.sparse.FloatTensor(i, data, torch.Size(L.shape)) 

1816 

1817 def _create_eye_matrix(self): 

1818 r"""Construct the identity matrix with the size of item_num + user_num. 

1819 

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) 

1827 

1828 def norm_adjacency_matrix(self, form="torch.sparse"): 

1829 """Get the normalized adjacency matrix of users and items. 

1830 

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

1832 using the laplace matrix. 

1833 

1834 .. math:: 

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

1836 

1837 Args: 

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

1839 

1840 Returns: 

1841 torch.sparse.FloatTensor: Normalized adjacency matrix. 

1842 

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

1850 

1851 def eye_matrix(self, form="torch.sparse"): 

1852 """Construct the identity matrix with the size of item_num + user_num. 

1853 

1854 Args: 

1855 form (str, optional): Format of the identity matrix. Defaults 

1856 to ``torch.sparse``. 

1857 

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

1865 

1866 def _history_matrix(self, row, value_field=None, max_history_len=None): 

1867 """Get dense matrix describe user/item's history interaction records. 

1868 

1869 ``history_matrix[i]`` represents ``i``'s history interacted item_id. 

1870 

1871 ``history_value[i]`` represents ``i``'s history interaction records' values. 

1872 ``0`` if ``value_field = None``. 

1873 

1874 ``history_len[i]`` represents number of ``i``'s history interaction records. 

1875 

1876 ``0`` is used as padding. 

1877 

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``. 

1884 

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

1892 

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

1905 

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 

1912 

1913 history_len = np.zeros(row_num, dtype=np.int64) 

1914 for row_id in row_ids: 

1915 history_len[row_id] += 1 

1916 

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 

1922 

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 ) 

1928 

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 

1938 

1939 return ( 

1940 torch.LongTensor(history_matrix), 

1941 torch.FloatTensor(history_value), 

1942 torch.LongTensor(history_len), 

1943 ) 

1944 

1945 def history_item_matrix(self, value_field=None, max_history_len=None): 

1946 """Get dense matrix describe user's history interaction records. 

1947 

1948 ``history_matrix[i]`` represents user ``i``'s history interacted item_id. 

1949 

1950 ``history_value[i]`` represents user ``i``'s history interaction records' values, 

1951 ``0`` if ``value_field = None``. 

1952 

1953 ``history_len[i]`` represents number of user ``i``'s history interaction records. 

1954 

1955 ``0`` is used as padding. 

1956 

1957 Args: 

1958 value_field (str, optional): Data of matrix, which should exist in ``self.inter_feat``. 

1959 Defaults to ``None``. 

1960 

1961 max_history_len (int): The maximum number of user's history interaction records. 

1962 Defaults to ``None``. 

1963 

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) 

1971 

1972 def history_user_matrix(self, value_field=None, max_history_len=None): 

1973 """Get dense matrix describe item's history interaction records. 

1974 

1975 ``history_matrix[i]`` represents item ``i``'s history interacted user_id. 

1976 

1977 ``history_value[i]`` represents item ``i``'s history interaction records' values, 

1978 ``0`` if ``value_field = None``. 

1979 

1980 ``history_len[i]`` represents number of item ``i``'s history interaction records. 

1981 

1982 ``0`` is used as padding. 

1983 

1984 Args: 

1985 value_field (str, optional): Data of matrix, which should exist in ``self.inter_feat``. 

1986 Defaults to ``None``. 

1987 

1988 max_history_len (int): The maximum number of item's history interaction records. 

1989 Defaults to ``None``. 

1990 

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) 

1998 

1999 def _get_used_ids(self, source_field, target_field): 

2000 """Get used ids from the interaction features. 

2001 

2002 Args: 

2003 source_field (str, optional): Source field name. 

2004 target_field (str, optional): Target field name. 

2005 

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

2016 

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 

2021 

2022 def get_user_used_ids(self): 

2023 """Get used item ids of each user from the interaction features. 

2024 

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

2030 

2031 return self._get_used_ids(self.uid_field, self.iid_field) 

2032 

2033 def get_item_used_ids(self): 

2034 """Get used user ids of each item from the interaction features. 

2035 

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

2041 

2042 return self._get_used_ids(self.iid_field, self.uid_field) 

2043 

2044 def get_preload_weight(self, field): 

2045 """Get preloaded weight matrix, whose rows are sorted by token ids. 

2046 

2047 ``0`` is used as padding. 

2048 

2049 Args: 

2050 field (str): preloaded feature field name. 

2051 

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] 

2058 

2059 def _dataframe_to_interaction(self, data): 

2060 """Convert :class:`pandas.DataFrame` to :class:`~hopwise.data.interaction.Interaction`. 

2061 

2062 Args: 

2063 data (pandas.DataFrame): data to be converted. 

2064 

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)