Coverage for hopwise/data/utils.py: 55%

249 statements  

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

1# @Time : 2020/7/21 

2# @Author : Yupeng Hou 

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

4 

5# UPDATE: 

6# @Time : 2021/7/9, 2020/9/17, 2020/8/31, 2021/2/20, 2021/3/1, 2022/7/6 

7# @Author : Yupeng Hou, Yushuo Chen, Kaiyuan Li, Haoran Cheng, Jiawei Guan, Gaowei Zhang 

8# @Email : houyupeng@ruc.edu.cn, chenyushuo@ruc.edu.cn, tsotfsk@outlook.com, chenghaoran29@foxmail.com, guanjw@ruc.edu.cn, zgw15630559577@163.com # noqa: E501 

9 

10"""hopwise.data.utils 

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

12""" 

13 

14# ruff: noqa: F403, F405 

15 

16import importlib 

17import os 

18import pickle 

19import warnings 

20from typing import Literal 

21 

22from hopwise.data.dataloader import * 

23from hopwise.sampler import KGSampler, RepeatableSampler, Sampler 

24from hopwise.utils import KnowledgeEvaluationType, ModelType, ensure_dir, set_color 

25from hopwise.utils.argument_list import dataset_arguments 

26 

27 

28def create_dataset(config): 

29 """Create dataset according to :attr:`config['model']` and :attr:`config['MODEL_TYPE']`. 

30 If :attr:`config['dataset_save_path']` file exists and 

31 its :attr:`config` of dataset is equal to current :attr:`config` of dataset. 

32 It will return the saved dataset in :attr:`config['dataset_save_path']`. 

33 

34 Args: 

35 config (Config): An instance object of Config, used to record parameter information. 

36 

37 Returns: 

38 Dataset: Constructed dataset. 

39 """ 

40 dataset_module = importlib.import_module("hopwise.data.dataset") 

41 

42 # Check if user-item knowledge graph links are available 

43 has_user_item_kg = os.path.isfile( 

44 os.path.join(config["data_path"], f"{config['dataset']}.user_link") 

45 ) and os.path.isfile(os.path.join(config["data_path"], f"{config['dataset']}.item_link")) 

46 

47 # Check for model-specific dataset class 

48 model_dataset_name = config["model"] + "Dataset" 

49 user_item_model_dataset_name = "UserItem" + model_dataset_name 

50 

51 if has_user_item_kg and hasattr(dataset_module, user_item_model_dataset_name): 

52 # Prefer UserItem variant when link files exist 

53 dataset_class = getattr(dataset_module, user_item_model_dataset_name) 

54 elif hasattr(dataset_module, model_dataset_name): 

55 dataset_class = getattr(dataset_module, model_dataset_name) 

56 else: 

57 model_type = config["MODEL_TYPE"] 

58 

59 if has_user_item_kg: 

60 kg_dataset_classname = "UserItemKnowledgeBasedDataset" 

61 path_language_model_dataset_classname = "UserItemKnowledgePathDataset" 

62 else: 

63 kg_dataset_classname = "KnowledgeBasedDataset" 

64 path_language_model_dataset_classname = "KnowledgePathDataset" 

65 

66 type2class = { 

67 ModelType.GENERAL: "Dataset", 

68 ModelType.SEQUENTIAL: "SequentialDataset", 

69 ModelType.CONTEXT: "Dataset", 

70 ModelType.KNOWLEDGE: kg_dataset_classname, 

71 ModelType.TRADITIONAL: "Dataset", 

72 ModelType.DECISIONTREE: "Dataset", 

73 ModelType.PATH_LANGUAGE_MODELING: path_language_model_dataset_classname, 

74 } 

75 dataset_class = getattr(dataset_module, type2class[model_type]) 

76 

77 default_file = os.path.join(config["checkpoint_dir"], f"{config['dataset']}-{dataset_class.__name__}.pth") 

78 file = config["dataset_save_path"] or default_file 

79 if os.path.exists(file): 

80 with open(file, "rb") as f: 

81 dataset = pickle.load(f) 

82 dataset_args_unchanged = True 

83 for arg in dataset_arguments + ["seed", "repeatable"]: 

84 if config[arg] != dataset.config[arg]: 

85 dataset_args_unchanged = False 

86 break 

87 if dataset_args_unchanged: 

88 logger = getLogger() 

89 logger.info(set_color("Load filtered dataset from", "magenta") + f": [{file}]") 

90 return dataset 

91 

92 dataset = dataset_class(config) 

93 if config["save_dataset"]: 

94 dataset.save() 

95 return dataset 

96 

97 

98def _get_dataloader_name(config, dataloaders_folder): 

99 path_gen_args = config["path_sample_args"] 

100 

101 max_path_per_user = config["MAX_PATHS_PER_USER"] 

102 max_rw_tries_per_iid = config["MAX_RW_TRIES_PER_IID"] 

103 restrict_by_phase = path_gen_args["restrict_by_phase"] 

104 temporal_causality = path_gen_args["temporal_causality"] 

105 strategy = path_gen_args["strategy"] 

106 collaborative_path = path_gen_args["collaborative_path"] 

107 

108 filename = f"{config['dataset']}-for-{config['model']}" 

109 f"-str {strategy}" 

110 f"-mppu {max_path_per_user}" 

111 f"-max_tries {max_rw_tries_per_iid}" 

112 f"-temp {temporal_causality}" 

113 f"-col {collaborative_path}" 

114 f"-restrbyphase {restrict_by_phase}-dataloader.pth" 

115 

116 if "train_stage" in config and config["MODEL_TYPE"] in [ModelType.PATH_LANGUAGE_MODELING]: 

117 filename += f"-{config['train_stage']}" 

118 

119 file_path = os.path.join( 

120 config["checkpoint_dir"], 

121 dataloaders_folder, 

122 filename, 

123 ) 

124 

125 return file_path 

126 

127 

128def save_split_dataloaders(config, dataloaders): 

129 """Save split dataloaders. 

130 

131 Args: 

132 config (Config): An instance object of Config, used to record parameter information. 

133 dataloaders (tuple of AbstractDataLoader): The split dataloaders. 

134 """ 

135 ensure_dir(config["checkpoint_dir"]) 

136 if config["MODEL_TYPE"] == ModelType.PATH_LANGUAGE_MODELING: 

137 dataloaders_folder = f"{config['model']} - {config['dataset']} - dataloaders" 

138 ensure_dir(os.path.join(config["checkpoint_dir"], dataloaders_folder)) 

139 file_path = _get_dataloader_name(config, dataloaders_folder) 

140 else: 

141 file_path = os.path.join( 

142 config["checkpoint_dir"], 

143 f"{config['dataset']}-for-{config['model']}-dataloader.pth", 

144 ) 

145 

146 logger = getLogger() 

147 logger.info(set_color("Saving split dataloaders into", "magenta") + f": [{file_path}]") 

148 serialization_dataloaders = [] 

149 for dataloader in dataloaders: 

150 if isinstance(dataloader, KnowledgeBasedDataLoader): 

151 general_generator_state = dataloader.general_dataloader.generator.get_state() 

152 dataloader.general_dataloader.generator = None 

153 dataloader.general_dataloader.sampler.generator = None 

154 kg_generator_state = dataloader.kg_dataloader.generator.get_state() 

155 dataloader.kg_dataloader.generator = None 

156 dataloader.kg_dataloader.sampler.generator = None 

157 serialization_dataloaders += [(dataloader, general_generator_state, kg_generator_state)] 

158 else: 

159 generator_state = dataloader.generator.get_state() 

160 dataloader.generator = None 

161 dataloader.sampler.generator = None 

162 serialization_dataloaders += [(dataloader, generator_state)] 

163 

164 with open(file_path, "wb") as f: 

165 pickle.dump(serialization_dataloaders, f) 

166 

167 

168def load_split_dataloaders(config): 

169 """Load split dataloaders if saved dataloaders exist and 

170 their :attr:`config` of dataset are the same as current :attr:`config` of dataset. 

171 

172 Args: 

173 config (Config): An instance object of Config, used to record parameter information. 

174 

175 Returns: 

176 dataloaders (tuple of AbstractDataLoader or None): The split dataloaders. 

177 """ 

178 

179 if config["MODEL_TYPE"] == ModelType.PATH_LANGUAGE_MODELING: 

180 dataloaders_folder = f"{config['model']} - {config['dataset']} - dataloaders" 

181 default_file = _get_dataloader_name(config, dataloaders_folder) 

182 else: 

183 default_file = os.path.join( 

184 config["checkpoint_dir"], 

185 f"{config['dataset']}-for-{config['model']}-dataloader.pth", 

186 ) 

187 

188 # used if you want to load a specific dataloader 

189 dataloaders_save_path = config["dataloaders_save_path"] or default_file 

190 

191 if not os.path.exists(dataloaders_save_path): 

192 return None 

193 with open(dataloaders_save_path, "rb") as f: 

194 dataloaders = [] 

195 with warnings.catch_warnings(): 

196 warnings.simplefilter(action="ignore", category=FutureWarning) 

197 for dataloader_saved_data in pickle.load(f): 

198 if isinstance(dataloader_saved_data[0], KnowledgeBasedDataLoader): 

199 data_loader, general_generator_state, kg_generator_state = dataloader_saved_data 

200 general_generator = torch.Generator() 

201 general_generator.set_state(general_generator_state) 

202 data_loader.general_dataloader.generator = general_generator 

203 data_loader.general_dataloader.sampler.generator = general_generator 

204 kg_generator = torch.Generator() 

205 kg_generator.set_state(kg_generator_state) 

206 data_loader.kg_dataloader.generator = kg_generator 

207 data_loader.kg_dataloader.sampler.generator = kg_generator 

208 dataloaders.append(data_loader) 

209 else: 

210 data_loader, generator_state = dataloader_saved_data 

211 generator = torch.Generator() 

212 generator.set_state(generator_state) 

213 data_loader.generator = generator 

214 data_loader.sampler.generator = generator 

215 dataloaders.append(data_loader) 

216 

217 eval_lp_args = config["eval_lp_args"] 

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

219 train_data, valid_inter_data, valid_kg_data, test_inter_data, test_kg_data = dataloaders 

220 else: 

221 train_data, valid_data, test_data = dataloaders 

222 for arg in dataset_arguments + ["seed", "repeatable", "eval_args"]: 

223 if isinstance(train_data, KnowledgeBasedDataLoader): 

224 general_config = train_data.general_dataloader.config 

225 kg_config = train_data.kg_dataloader.config 

226 if config[arg] != general_config[arg] and config[arg] != kg_config[arg]: 

227 return None 

228 elif config[arg] != train_data.config[arg]: 

229 return None 

230 train_data.update_config(config) 

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

232 valid_inter_data.update_config(config) 

233 valid_kg_data.update_config(config) 

234 test_inter_data.update_config(config) 

235 test_kg_data.update_config(config) 

236 

237 valid_data = [valid_inter_data, valid_kg_data] 

238 test_data = [test_inter_data, test_kg_data] 

239 else: 

240 valid_data.update_config(config) 

241 test_data.update_config(config) 

242 logger = getLogger() 

243 logger.info(set_color("Load split dataloaders from", "magenta") + f": [{dataloaders_save_path}]") 

244 return train_data, valid_data, test_data 

245 

246 

247def data_preparation(config, dataset): 

248 """Split the dataset by :attr:`config['[valid|test]_eval_args']` and create training, validation and test dataloader. 

249 

250 Note: 

251 If we can load split dataloaders by :meth:`load_split_dataloaders`, we will not create new split dataloaders. 

252 

253 Args: 

254 config (Config): An instance object of Config, used to record parameter information. 

255 dataset (Dataset): An instance object of Dataset, which contains all interaction records. 

256 

257 Returns: 

258 tuple: 

259 - train_data (AbstractDataLoader): The dataloader for training. 

260 - valid_data (AbstractDataLoader): The dataloader for validation. 

261 - test_data (AbstractDataLoader): The dataloader for testing. 

262 """ # noqa: E501 

263 dataloaders = load_split_dataloaders(config) 

264 if dataloaders is not None: 

265 train_data, valid_data, test_data = dataloaders 

266 dataset._change_feat_format() 

267 else: 

268 model_type = config["MODEL_TYPE"] 

269 model_input_type = config["MODEL_INPUT_TYPE"] 

270 # model = config["model"] 

271 built_datasets = dataset.build() 

272 

273 if model_type in [ModelType.KNOWLEDGE] and model_input_type not in [InputType.USERWISE]: 

274 if isinstance(built_datasets, dict): 

275 # then the kg has been split 

276 train_kg_dataset, valid_kg_dataset, test_kg_dataset = built_datasets[KnowledgeEvaluationType.LP] 

277 train_inter_dataset, valid_inter_dataset, test_inter_dataset = built_datasets[ 

278 KnowledgeEvaluationType.REC 

279 ] 

280 

281 kg_sampler = KGSampler( 

282 train_kg_dataset, 

283 config["train_neg_sample_args"]["distribution"], 

284 config["train_neg_sample_args"]["alpha"], 

285 ) 

286 

287 train_inter_sampler, valid_inter_sampler, test_inter_sampler = create_samplers( 

288 config, dataset, built_datasets[KnowledgeEvaluationType.REC] 

289 ) 

290 

291 train_data = get_dataloader(config, "train")( 

292 config, train_inter_dataset, train_inter_sampler, kg_sampler, shuffle=True 

293 ) 

294 

295 valid_kg_sampler = KGSampler(valid_kg_dataset, distribution=None) 

296 valid_kg_data = get_dataloader(config, "valid", task=KnowledgeEvaluationType.LP)( 

297 config, valid_kg_dataset, valid_kg_sampler, shuffle=False 

298 ) 

299 valid_inter_data = get_dataloader(config, "valid", task=KnowledgeEvaluationType.REC)( 

300 config, valid_inter_dataset, valid_inter_sampler, shuffle=False 

301 ) 

302 

303 test_kg_sampler = KGSampler(test_kg_dataset, distribution=None) 

304 test_kg_data = get_dataloader(config, "valid", task=KnowledgeEvaluationType.LP)( 

305 config, test_kg_dataset, test_kg_sampler, shuffle=False 

306 ) 

307 test_inter_data = get_dataloader(config, "test", task=KnowledgeEvaluationType.REC)( 

308 config, test_inter_dataset, test_inter_sampler, shuffle=False 

309 ) 

310 

311 if config["save_dataloaders"]: 

312 save_split_dataloaders( 

313 config, 

314 dataloaders=(train_data, valid_inter_data, valid_kg_data, test_inter_data, test_kg_data), 

315 ) 

316 

317 valid_data = [valid_inter_data, valid_kg_data] 

318 test_data = [test_inter_data, test_kg_data] 

319 else: 

320 # then we want to use the whole kg 

321 kg_sampler = KGSampler( 

322 dataset, 

323 config["train_neg_sample_args"]["distribution"], 

324 config["train_neg_sample_args"]["alpha"], 

325 ) 

326 

327 train_dataset, valid_dataset, test_dataset = built_datasets 

328 train_sampler, valid_sampler, test_sampler = create_samplers(config, dataset, built_datasets) 

329 train_data = get_dataloader(config, "train")( 

330 config, train_dataset, train_sampler, kg_sampler, shuffle=True 

331 ) 

332 valid_data = get_dataloader(config, "valid")(config, valid_dataset, valid_sampler, shuffle=False) 

333 test_data = get_dataloader(config, "test")(config, test_dataset, test_sampler, shuffle=False) 

334 

335 if config["save_dataloaders"]: 

336 save_split_dataloaders(config, dataloaders=(train_data, valid_data, test_data)) 

337 else: 

338 train_dataset, valid_dataset, test_dataset = built_datasets 

339 train_sampler, valid_sampler, test_sampler = create_samplers(config, dataset, built_datasets) 

340 

341 train_data = get_dataloader(config, "train")( 

342 config, train_dataset, train_sampler, shuffle=config["shuffle"] 

343 ) 

344 valid_data = get_dataloader(config, "valid")(config, valid_dataset, valid_sampler, shuffle=False) 

345 test_data = get_dataloader(config, "test")(config, test_dataset, test_sampler, shuffle=False) 

346 

347 if config["save_dataloaders"]: 

348 save_split_dataloaders(config, dataloaders=(train_data, valid_data, test_data)) 

349 

350 logger = getLogger() 

351 logger.info( 

352 set_color("[Training]: ", "magenta") 

353 + set_color("train_batch_size", "cyan") 

354 + " = " 

355 + set_color(f"[{config['train_batch_size']}]", "yellow") 

356 + set_color(" train_neg_sample_args", "cyan") 

357 + ": " 

358 + set_color(f"[{config['train_neg_sample_args']}]", "yellow") 

359 ) 

360 

361 if config["eval_lp_args"] is not None and config["eval_lp_args"]["knowledge_split"] is not None: 

362 eval_lp_args_info = ( 

363 set_color(" eval_lp_args", "cyan") + ": " + set_color(f"[{config['eval_lp_args']}]", "yellow") 

364 ) 

365 else: 

366 eval_lp_args_info = "" 

367 

368 logger.info( 

369 set_color("[Evaluation]: ", "magenta") 

370 + set_color("eval_batch_size", "cyan") 

371 + " = " 

372 + set_color(f"[{config['eval_batch_size']}]", "yellow") 

373 + set_color(" eval_args", "cyan") 

374 + ": " 

375 + set_color(f"[{config['eval_args']}]", "yellow") 

376 + eval_lp_args_info 

377 ) 

378 return train_data, valid_data, test_data 

379 

380 

381def get_dataloader(config, phase: Literal["train", "valid", "test", "evaluation"], task=KnowledgeEvaluationType.REC): # noqa: PLR0911 

382 """Return a dataloader class according to :attr:`config` and :attr:`phase`. 

383 

384 Args: 

385 config (Config): An instance object of Config, used to record parameter information. 

386 phase (str): The stage of dataloader. It can only take 4 values: 'train', 'valid', 'test' or 'evaluation'. 

387 Notes: 'evaluation' has been deprecated, please use 'valid' or 'test' instead. 

388 

389 Returns: 

390 type: The dataloader class that meets the requirements in :attr:`config` and :attr:`phase`. 

391 """ 

392 

393 if phase not in ["train", "valid", "test", "evaluation"]: 

394 raise ValueError("`phase` can only be 'train', 'valid', 'test' or 'evaluation'.") 

395 if phase == "evaluation": 

396 phase = "test" 

397 warnings.warn( 

398 "'evaluation' has been deprecated, please use 'valid' or 'test' instead.", 

399 DeprecationWarning, 

400 ) 

401 

402 if config["MODEL_INPUT_TYPE"] == InputType.USERWISE: 

403 if config["model"] in ["TPRec"]: 

404 if config["train_stage"] in ["policy"]: 

405 return _get_user_dataloader(config, phase) 

406 else: 

407 return _get_user_dataloader(config, phase) 

408 

409 model_type = config["MODEL_TYPE"] 

410 if phase == "train": 

411 # Return Dataloader based on the modeltype 

412 if model_type == ModelType.KNOWLEDGE: 

413 return KnowledgeBasedDataLoader 

414 else: 

415 return TrainDataLoader 

416 else: 

417 eval_mode = config["eval_args"]["mode"][phase] 

418 if eval_mode == "full": 

419 if model_type == ModelType.PATH_LANGUAGE_MODELING: 

420 return KnowledgePathEvalDataLoader 

421 else: 

422 if task is KnowledgeEvaluationType.LP: 

423 return FullSortLPEvalDataLoader 

424 return FullSortRecEvalDataLoader 

425 else: 

426 return NegSampleEvalDataLoader 

427 

428 

429def _get_user_dataloader(config, phase: Literal["train", "valid", "test", "evaluation"]): 

430 """Customized function for models that needs only users 

431 

432 Args: 

433 config (Config): An instance object of Config, used to record parameter information. 

434 phase (str): The stage of dataloader. It can only take 4 values: 'train', 'valid', 'test' or 'evaluation'. 

435 Notes: 'evaluation' has been deprecated, please use 'valid' or 'test' instead. 

436 

437 Returns: 

438 type: The dataloader class that meets the requirements in :attr:`config` and :attr:`phase`. 

439 """ 

440 if phase not in ["train", "valid", "test", "evaluation"]: 

441 raise ValueError("`phase` can only be 'train', 'valid', 'test' or 'evaluation'.") 

442 if phase == "evaluation": 

443 phase = "test" 

444 warnings.warn( 

445 "'evaluation' has been deprecated, please use 'valid' or 'test' instead.", 

446 DeprecationWarning, 

447 ) 

448 

449 if phase == "train": 

450 return UserDataLoader 

451 else: 

452 eval_mode = config["eval_args"]["mode"][phase] 

453 if eval_mode == "full": 

454 return FullSortRecEvalDataLoader 

455 else: 

456 return NegSampleEvalDataLoader 

457 

458 

459def _create_sampler( 

460 dataset, 

461 built_datasets, 

462 distribution: str, 

463 repeatable: bool, 

464 alpha: float = 1.0, 

465 base_sampler=None, 

466): 

467 phases = ["train", "valid", "test"] 

468 sampler = None 

469 if distribution != "none": 

470 if base_sampler is not None: 

471 base_sampler.set_distribution(distribution) 

472 return base_sampler 

473 if not repeatable: 

474 sampler = Sampler( 

475 phases, 

476 built_datasets, 

477 distribution, 

478 alpha, 

479 ) 

480 else: 

481 sampler = RepeatableSampler( 

482 phases, 

483 dataset, 

484 distribution, 

485 alpha, 

486 ) 

487 return sampler 

488 

489 

490def create_samplers(config, dataset, built_datasets): 

491 """Create sampler for training, validation and testing. 

492 

493 Args: 

494 config (Config): An instance object of Config, used to record parameter information. 

495 dataset (Dataset): An instance object of Dataset, which contains all interaction records. 

496 built_datasets (list of Dataset): A list of split Dataset, which contains dataset for 

497 training, validation and testing. 

498 

499 Returns: 

500 tuple: 

501 - train_sampler (AbstractSampler): The sampler for training. 

502 - valid_sampler (AbstractSampler): The sampler for validation. 

503 - test_sampler (AbstractSampler): The sampler for testing. 

504 """ 

505 train_neg_sample_args = config["train_neg_sample_args"] 

506 valid_neg_sample_args = config["valid_neg_sample_args"] 

507 test_neg_sample_args = config["test_neg_sample_args"] 

508 repeatable = config["repeatable"] 

509 base_sampler = _create_sampler( 

510 dataset, 

511 built_datasets, 

512 train_neg_sample_args["distribution"], 

513 repeatable, 

514 train_neg_sample_args["alpha"], 

515 ) 

516 train_sampler = base_sampler.set_phase("train") if base_sampler else None 

517 

518 valid_sampler = _create_sampler( 

519 dataset, 

520 built_datasets, 

521 valid_neg_sample_args["distribution"], 

522 repeatable, 

523 base_sampler=base_sampler, 

524 ) 

525 valid_sampler = valid_sampler.set_phase("valid") if valid_sampler else None 

526 

527 test_sampler = _create_sampler( 

528 dataset, 

529 built_datasets, 

530 test_neg_sample_args["distribution"], 

531 repeatable, 

532 base_sampler=base_sampler, 

533 ) 

534 test_sampler = test_sampler.set_phase("test") if test_sampler else None 

535 return train_sampler, valid_sampler, test_sampler