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
« 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
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
10"""hopwise.data.utils
11########################
12"""
14# ruff: noqa: F403, F405
16import importlib
17import os
18import pickle
19import warnings
20from typing import Literal
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
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']`.
34 Args:
35 config (Config): An instance object of Config, used to record parameter information.
37 Returns:
38 Dataset: Constructed dataset.
39 """
40 dataset_module = importlib.import_module("hopwise.data.dataset")
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"))
47 # Check for model-specific dataset class
48 model_dataset_name = config["model"] + "Dataset"
49 user_item_model_dataset_name = "UserItem" + model_dataset_name
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"]
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"
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])
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
92 dataset = dataset_class(config)
93 if config["save_dataset"]:
94 dataset.save()
95 return dataset
98def _get_dataloader_name(config, dataloaders_folder):
99 path_gen_args = config["path_sample_args"]
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"]
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"
116 if "train_stage" in config and config["MODEL_TYPE"] in [ModelType.PATH_LANGUAGE_MODELING]:
117 filename += f"-{config['train_stage']}"
119 file_path = os.path.join(
120 config["checkpoint_dir"],
121 dataloaders_folder,
122 filename,
123 )
125 return file_path
128def save_split_dataloaders(config, dataloaders):
129 """Save split dataloaders.
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 )
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)]
164 with open(file_path, "wb") as f:
165 pickle.dump(serialization_dataloaders, f)
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.
172 Args:
173 config (Config): An instance object of Config, used to record parameter information.
175 Returns:
176 dataloaders (tuple of AbstractDataLoader or None): The split dataloaders.
177 """
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 )
188 # used if you want to load a specific dataloader
189 dataloaders_save_path = config["dataloaders_save_path"] or default_file
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)
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)
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
247def data_preparation(config, dataset):
248 """Split the dataset by :attr:`config['[valid|test]_eval_args']` and create training, validation and test dataloader.
250 Note:
251 If we can load split dataloaders by :meth:`load_split_dataloaders`, we will not create new split dataloaders.
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.
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()
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 ]
281 kg_sampler = KGSampler(
282 train_kg_dataset,
283 config["train_neg_sample_args"]["distribution"],
284 config["train_neg_sample_args"]["alpha"],
285 )
287 train_inter_sampler, valid_inter_sampler, test_inter_sampler = create_samplers(
288 config, dataset, built_datasets[KnowledgeEvaluationType.REC]
289 )
291 train_data = get_dataloader(config, "train")(
292 config, train_inter_dataset, train_inter_sampler, kg_sampler, shuffle=True
293 )
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 )
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 )
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 )
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 )
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)
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)
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)
347 if config["save_dataloaders"]:
348 save_split_dataloaders(config, dataloaders=(train_data, valid_data, test_data))
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 )
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 = ""
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
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`.
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.
389 Returns:
390 type: The dataloader class that meets the requirements in :attr:`config` and :attr:`phase`.
391 """
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 )
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)
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
429def _get_user_dataloader(config, phase: Literal["train", "valid", "test", "evaluation"]):
430 """Customized function for models that needs only users
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.
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 )
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
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
490def create_samplers(config, dataset, built_datasets):
491 """Create sampler for training, validation and testing.
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.
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
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
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