Coverage for hopwise/config/configurator.py: 76%
403 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-09-30 13:25 +0000
1# @Time : 2020/6/28
2# @Author : Zihan Lin
3# @Email : linzihan.super@foxmail.com
5# UPDATE
6# @Time : 2020/10/04, 2021/3/2, 2021/2/17, 2021/6/30, 2022/7/6
7# @Author : Shanlei Mu, Yupeng Hou, Jiawei Guan, Xingyu Pan, Gaowei Zhang
8# @Email : slmu@ruc.edu.cn, houyupeng@ruc.edu.cn, Guanjw@ruc.edu.cn, xy_pan@foxmail.com, zgw15630559577@163.com
10# @Time : 2025
11# @Author : Giacomo Medda, Alessandro Soccol
12# @Email : giacomo.medda@unica.it, alessandro.soccol@unica.it
15"""hopwise.config.configurator
16################################
17"""
19import os
20import re
21import sys
22import warnings
23from logging import getLogger
24from typing import Literal
26import yaml
28from hopwise.evaluator import metric_types, smaller_metrics
29from hopwise.utils import (
30 Enum,
31 EvaluatorType,
32 InputType,
33 ModelType,
34 dataset_arguments,
35 deep_dict_update,
36 evaluation_arguments,
37 general_arguments,
38 get_model,
39 set_color,
40 training_arguments,
41)
44class Config:
45 """Configurator module that load the defined parameters.
47 Configurator module will first load the default parameters from the fixed properties in hopwise and then
48 load parameters from the external input.
50 External input supports three kind of forms: config file, command line and parameter dictionaries.
52 - config file: It's a file that record the parameters to be modified or added. It should be in ``yaml`` format,
53 e.g. a config file is 'example.yaml', the content is:
55 learning_rate: 0.001
57 train_batch_size: 2048
59 - command line: It should be in the format as '---learning_rate=0.001'
61 - parameter dictionaries: It should be a dict, where the key is parameter name and the value is parameter value,
62 e.g. config_dict = {'learning_rate': 0.001}
64 Configuration module allows the above three kind of external input format to be used together,
65 the priority order is as following:
67 command line > parameter dictionaries > config file
69 e.g. If we set learning_rate=0.01 in config file, learning_rate=0.02 in command line,
70 learning_rate=0.03 in parameter dictionaries.
72 Finally the learning_rate is equal to 0.02.
73 """
75 NESTED_KEY_SEPARATOR = "."
77 def __init__(self, model=None, dataset=None, config_file_list=None, config_dict=None):
78 """Args:
79 model (str/AbstractRecommender): the model name or the model class, default is None, if it is None, config
80 will search the parameter 'model' from the external input as the model name or model class.
81 dataset (str): the dataset name, default is None, if it is None, config will search the parameter 'dataset'
82 from the external input as the dataset name.
83 config_file_list (list of str): the external config file, it allows multiple config files, default is None.
84 config_dict (dict): the external parameter dictionaries, default is None.
85 """
86 self.compatibility_settings()
87 self._init_parameters_category()
88 self.yaml_loader = self._build_yaml_loader()
89 self.file_config_dict = self._load_config_files(config_file_list)
90 self.variable_config_dict = self._load_variable_config_dict(config_dict)
91 self.cmd_config_dict = self._load_cmd_line()
92 self._merge_external_config_dict()
94 self.model, self.model_class, self.dataset = self._get_model_and_dataset(model, dataset)
95 self._load_internal_config_dict(self.model, self.model_class, self.dataset)
96 self.final_config_dict = self._get_final_config_dict()
97 self._set_default_parameters()
98 self._init_device()
99 self._set_env_behavior()
100 self._set_torch_dtype()
101 self._set_train_neg_sample_args()
102 self._set_eval_neg_sample_args("valid")
103 self._set_eval_neg_sample_args("test")
105 def _init_parameters_category(self):
106 self.parameters = dict()
107 self.parameters["General"] = general_arguments
108 self.parameters["Training"] = training_arguments
109 self.parameters["Evaluation"] = evaluation_arguments
110 self.parameters["Dataset"] = dataset_arguments
112 def _build_yaml_loader(self):
113 loader = yaml.FullLoader
114 loader.add_implicit_resolver(
115 "tag:yaml.org,2002:float",
116 re.compile(
117 """^(?:
118 [-+]?(?:[0-9][0-9_]*)\\.[0-9_]*(?:[eE][-+]?[0-9]+)?
119 |[-+]?(?:[0-9][0-9_]*)(?:[eE][-+]?[0-9]+)
120 |\\.[0-9_]+(?:[eE][-+][0-9]+)?
121 |[-+]?[0-9][0-9_]*(?::[0-5]?[0-9])+\\.[0-9_]*
122 |[-+]?\\.(?:inf|Inf|INF)
123 |\\.(?:nan|NaN|NAN))$""",
124 re.X,
125 ),
126 list("-+0123456789."),
127 )
128 return loader
130 def _convert_config_dict(self, config_dict):
131 r"""This function convert the str parameters to their original type."""
132 config_keys = list(config_dict.keys())
133 for key in config_keys:
134 param = config_dict[key]
136 if not isinstance(param, str):
137 if isinstance(param, dict):
138 if self.NESTED_KEY_SEPARATOR in key:
139 raise SyntaxError(
140 f"If '{self.NESTED_KEY_SEPARATOR}' is used in the key, "
141 f"the value should be str, but got [{param}] of type [{type(param)}] instead."
142 )
143 config_dict[key] = self._convert_config_dict(param)
144 continue
145 try:
146 value = eval(param)
147 if value is not None and not isinstance(value, (str, int, float, list, tuple, dict, bool, Enum)):
148 value = param
149 except (NameError, SyntaxError, TypeError):
150 if isinstance(param, str):
151 if param.lower() == "true":
152 value = True
153 elif param.lower() == "false":
154 value = False
155 else:
156 value = param
157 else:
158 value = param
159 if self.NESTED_KEY_SEPARATOR in key:
160 nested_cmd_dict = param
161 for nested_key in reversed(key.split(self.NESTED_KEY_SEPARATOR)):
162 nested_cmd_dict = {nested_key: nested_cmd_dict}
163 deep_dict_update(config_dict, nested_cmd_dict)
164 del config_dict[key]
165 else:
166 config_dict[key] = value
167 return config_dict
169 def _load_config_files(self, file_list):
170 file_config_dict = dict()
171 if file_list:
172 for file in file_list:
173 with open(file, encoding="utf-8") as f:
174 deep_dict_update(file_config_dict, yaml.load(f.read(), Loader=self.yaml_loader))
175 return file_config_dict
177 def _load_variable_config_dict(self, config_dict):
178 # HyperTuning may set the parameters such as mlp_hidden_size in NeuMF in the format of ['[]', '[]']
179 # then config_dict will receive a str '[]', but indeed it's a list []
180 # temporarily use _convert_config_dict to solve this problem
181 return self._convert_config_dict(config_dict) if config_dict else dict()
183 def _load_cmd_line(self):
184 r"""Read parameters from command line and convert it to str."""
185 cmd_config_dict = dict()
186 unrecognized_args = []
187 if "ipykernel_launcher" not in sys.argv[0]:
188 for arg in sys.argv[1:]:
189 if not arg.startswith("--") or len(arg[2:].split("=")) != 2: # noqa: PLR2004
190 unrecognized_args.append(arg)
191 continue
192 cmd_arg_name, cmd_arg_value = arg[2:].split("=")
193 if cmd_arg_name in cmd_config_dict and cmd_arg_value != cmd_config_dict[cmd_arg_name]:
194 raise SyntaxError("There are duplicate command arg '%s' with different value." % arg)
195 elif self.NESTED_KEY_SEPARATOR in cmd_arg_name:
196 nested_cmd_dict = self._convert_config_dict({cmd_arg_name: cmd_arg_value})
197 deep_dict_update(cmd_config_dict, nested_cmd_dict)
198 else:
199 cmd_config_dict[cmd_arg_name] = cmd_arg_value
200 cmd_config_dict = self._convert_config_dict(cmd_config_dict)
201 return cmd_config_dict
203 def _merge_external_config_dict(self):
204 external_config_dict = dict()
205 deep_dict_update(external_config_dict, self.file_config_dict)
206 deep_dict_update(external_config_dict, self.variable_config_dict)
207 deep_dict_update(external_config_dict, self.cmd_config_dict)
208 self.external_config_dict = external_config_dict
210 def _get_model_and_dataset(self, model, dataset):
211 if model is None:
212 try:
213 model = self.external_config_dict["model"]
214 except KeyError:
215 raise KeyError(
216 "model need to be specified in at least one of the these ways: "
217 "[model variable, config file, config dict, command line] "
218 )
219 if not isinstance(model, str):
220 final_model_class = model
221 final_model = model.__name__
222 else:
223 final_model = model
224 final_model_class = get_model(final_model)
226 if dataset is None:
227 try:
228 final_dataset = self.external_config_dict["dataset"]
229 except KeyError:
230 raise KeyError(
231 "dataset need to be specified in at least one of the these ways: "
232 "[dataset variable, config file, config dict, command line] "
233 )
234 else:
235 final_dataset = dataset
237 return final_model, final_model_class, final_dataset
239 def _update_internal_config_dict(self, file):
240 with open(file, encoding="utf-8") as f:
241 config_dict = yaml.load(f.read(), Loader=self.yaml_loader)
242 if config_dict is not None:
243 deep_dict_update(self.internal_config_dict, config_dict)
244 return config_dict
246 def _load_internal_config_dict(self, model, model_class, dataset):
247 current_path = os.path.dirname(os.path.realpath(__file__))
248 overall_init_file = os.path.join(current_path, "../properties/overall.yaml")
249 model_init_file = os.path.join(current_path, "../properties/model/" + model + ".yaml")
250 sample_init_file = os.path.join(current_path, "../properties/dataset/sample.yaml")
251 dataset_init_file = os.path.join(current_path, "../properties/dataset/" + dataset + ".yaml")
253 quick_start_config_path = os.path.join(current_path, "../properties/quick_start_config/")
254 context_aware_init = os.path.join(quick_start_config_path, "context-aware.yaml")
255 context_aware_on_ml_100k_init = os.path.join(quick_start_config_path, "context-aware_ml-100k.yaml")
256 DIN_init = os.path.join(quick_start_config_path, "sequential_DIN.yaml")
257 DIN_on_ml_100k_init = os.path.join(quick_start_config_path, "sequential_DIN_on_ml-100k.yaml")
258 sequential_init = os.path.join(quick_start_config_path, "sequential.yaml")
259 special_sequential_on_ml_100k_init = os.path.join(
260 quick_start_config_path, "special_sequential_on_ml-100k.yaml"
261 )
262 sequential_embedding_model_init = os.path.join(quick_start_config_path, "sequential_embedding_model.yaml")
263 knowledge_base_init = os.path.join(quick_start_config_path, "knowledge_base.yaml")
264 knowledge_base_on_ml_100k_init = os.path.join(quick_start_config_path, "knowledge_base_on_ml-100k.yaml")
265 knowledge_path_base_init = os.path.join(quick_start_config_path, "knowledge_path_base.yaml")
267 self.internal_config_dict = dict()
268 for file in [
269 overall_init_file,
270 sample_init_file,
271 ]:
272 if os.path.isfile(file):
273 config_dict = self._update_internal_config_dict(file)
274 if file == dataset_init_file:
275 self.parameters["Dataset"] += [
276 key for key in config_dict.keys() if key not in self.parameters["Dataset"]
277 ]
279 self.internal_config_dict["MODEL_TYPE"] = model_class.type
280 if self.internal_config_dict["MODEL_TYPE"] == ModelType.GENERAL:
281 pass
282 elif self.internal_config_dict["MODEL_TYPE"] in {
283 ModelType.CONTEXT,
284 ModelType.DECISIONTREE,
285 }:
286 self._update_internal_config_dict(context_aware_init)
287 if dataset == "ml-100k":
288 self._update_internal_config_dict(context_aware_on_ml_100k_init)
289 elif self.internal_config_dict["MODEL_TYPE"] == ModelType.SEQUENTIAL:
290 if model in ["DIN", "DIEN"]:
291 self._update_internal_config_dict(DIN_init)
292 if dataset == "ml-100k":
293 self._update_internal_config_dict(DIN_on_ml_100k_init)
294 else:
295 self._update_internal_config_dict(sequential_init)
296 if model in ["GRU4RecKG", "KSR"]:
297 self._update_internal_config_dict(sequential_embedding_model_init)
298 if dataset == "ml-100k" and model in [
299 "GRU4RecF",
300 "SASRecF",
301 "FDSA",
302 "S3Rec",
303 ]:
304 self._update_internal_config_dict(special_sequential_on_ml_100k_init)
306 elif self.internal_config_dict["MODEL_TYPE"] == ModelType.KNOWLEDGE:
307 self._update_internal_config_dict(knowledge_base_init)
308 if dataset == "ml-100k":
309 self._update_internal_config_dict(knowledge_base_on_ml_100k_init)
310 elif self.internal_config_dict["MODEL_TYPE"] == ModelType.PATH_LANGUAGE_MODELING:
311 self._update_internal_config_dict(knowledge_path_base_init)
312 if dataset == "ml-100k":
313 self._update_internal_config_dict(knowledge_base_on_ml_100k_init)
315 # avoids parallelism issues with HuggingFace Tokenizers
316 os.environ["TOKENIZERS_PARALLELISM"] = "false"
318 for file in [
319 dataset_init_file,
320 model_init_file,
321 ]:
322 if os.path.isfile(file):
323 config_dict = self._update_internal_config_dict(file)
324 if file == dataset_init_file:
325 self.parameters["Dataset"] += [
326 key for key in config_dict.keys() if key not in self.parameters["Dataset"]
327 ]
329 def _get_final_config_dict(self):
330 final_config_dict = dict()
331 deep_dict_update(final_config_dict, self.internal_config_dict)
332 deep_dict_update(final_config_dict, self.external_config_dict)
333 return final_config_dict
335 def _set_default_parameters(self):
336 self.final_config_dict["dataset"] = self.dataset
337 self.final_config_dict["model"] = self.model
338 if self.dataset == "ml-100k" and self.external_config_dict.get("data_path", None) is None:
339 current_path = os.path.dirname(os.path.realpath(__file__))
340 self.final_config_dict["data_path"] = os.path.join(current_path, "../dataset_example/" + self.dataset)
341 else:
342 self.final_config_dict["data_path"] = os.path.join(self.final_config_dict["data_path"], self.dataset)
344 if hasattr(self.model_class, "input_type"):
345 self.final_config_dict["MODEL_INPUT_TYPE"] = self.model_class.input_type
346 elif "loss_type" in self.final_config_dict:
347 if self.final_config_dict["loss_type"] in ["CE"]:
348 if (
349 self.final_config_dict["MODEL_TYPE"] == ModelType.SEQUENTIAL
350 and self.final_config_dict.get("train_neg_sample_args") is not None
351 ):
352 raise ValueError(
353 f"train_neg_sample_args [{self.final_config_dict['train_neg_sample_args']}] should be None "
354 f"when the loss_type is CE."
355 )
356 self.final_config_dict["MODEL_INPUT_TYPE"] = InputType.POINTWISE
357 elif self.final_config_dict["loss_type"] in ["BPR"]:
358 self.final_config_dict["MODEL_INPUT_TYPE"] = InputType.PAIRWISE
359 else:
360 raise ValueError("Either Model has attr 'input_type',or arg 'loss_type' should exist in config.")
362 # handle special cases for models that needs pretrain in one way
363 if self.final_config_dict["model"] in ["TPRec"]:
364 train_stage = self.final_config_dict.get("train_stage")
366 if train_stage is not None and train_stage in ["pretrain"]:
367 self.final_config_dict["MODEL_INPUT_TYPE"] = InputType.PAIRWISE
369 metrics = self.final_config_dict["metrics"]
370 if isinstance(metrics, str):
371 self.final_config_dict["metrics"] = [metrics]
373 eval_type = set()
374 for metric in self.final_config_dict["metrics"]:
375 if metric.lower() in metric_types:
376 eval_type.add(metric_types[metric.lower()])
377 else:
378 raise NotImplementedError(f"There is no metric named '{metric}'")
379 if len(eval_type) > 1:
380 raise RuntimeError("Ranking metrics and value metrics can not be used at the same time.")
381 self.final_config_dict["eval_type"] = eval_type.pop()
383 if self.final_config_dict["MODEL_TYPE"] == ModelType.SEQUENTIAL and not self.final_config_dict["repeatable"]:
384 raise ValueError(
385 "Sequential models currently only support repeatable recommendation, "
386 "please set `repeatable` as `True`."
387 )
389 valid_metric = self.final_config_dict["valid_metric"].split("@")[0]
390 self.final_config_dict["valid_metric_bigger"] = False if valid_metric.lower() in smaller_metrics else True
392 topk = self.final_config_dict["topk"]
393 if isinstance(topk, (int, list)):
394 if isinstance(topk, int):
395 topk = [topk]
396 for k in topk:
397 if k <= 0:
398 raise ValueError(f"topk must be a positive integer or a list of positive integers, but get `{k}`")
399 self.final_config_dict["topk"] = topk
400 else:
401 raise TypeError(f"The topk [{topk}] must be a integer, list")
403 # Knowledge Graph
404 eval_lp_args = self.final_config_dict.get("eval_lp_args")
405 if (
406 self.final_config_dict["MODEL_TYPE"] == ModelType.KNOWLEDGE
407 and eval_lp_args is not None
408 and eval_lp_args["knowledge_split"] is not None
409 ):
410 metrics_kg = self.final_config_dict["metrics_lp"]
411 if isinstance(metrics_kg, str):
412 self.final_config_dict["metrics_lp"] = [metrics_kg]
414 eval_type_kg = set()
415 for metric in self.final_config_dict["metrics_lp"]:
416 if metric.lower() in metric_types:
417 eval_type_kg.add(metric_types[metric.lower()])
418 else:
419 raise NotImplementedError(f"There is no metric named '{metric}'")
420 if len(eval_type_kg) > 1:
421 raise RuntimeError("Ranking metrics and value metrics can not be used at the same time.")
423 self.final_config_dict["eval_type"] = eval_type_kg.pop()
425 if "valid_metrics_kg" in self.final_config_dict:
426 valid_metric_kg = self.final_config_dict["valid_metric_kg"].split("@")[0]
427 self.final_config_dict["valid_metric_bigger_kg"] = (
428 False if valid_metric_kg.lower() in smaller_metrics else True
429 )
431 topk_kg = self.final_config_dict["topk_kg"]
432 if isinstance(topk_kg, (int, list)):
433 if isinstance(topk_kg, int):
434 topk_kg = [topk_kg]
435 for k in topk_kg:
436 if k <= 0:
437 raise ValueError(
438 f"topk_kg must be a positive integer or a list of positive integers, but get `{k}`"
439 )
440 self.final_config_dict["topk_kg"] = topk_kg
441 else:
442 raise TypeError(f"The topk_kg [{topk_kg}] must be a integer, list")
444 if "additional_feat_suffix" in self.final_config_dict:
445 ad_suf = self.final_config_dict["additional_feat_suffix"]
446 if isinstance(ad_suf, str):
447 self.final_config_dict["additional_feat_suffix"] = [ad_suf]
449 # train_neg_sample_args checking
450 default_train_neg_sample_args = {
451 "distribution": "uniform",
452 "sample_num": 1,
453 "alpha": 1.0,
454 "dynamic": False,
455 "candidate_num": 0,
456 }
458 if (
459 self.final_config_dict.get("neg_sampling") is not None
460 or self.final_config_dict.get("training_neg_sample_num") is not None
461 ):
462 logger = getLogger()
463 logger.warning(
464 "Warning: Parameter 'neg_sampling' or 'training_neg_sample_num' has been deprecated in the new version, " # noqa: E501
465 "please use 'train_neg_sample_args' instead and check the API documentation for proper usage."
466 )
468 if self.final_config_dict.get("train_neg_sample_args") is not None:
469 if not isinstance(self.final_config_dict["train_neg_sample_args"], dict):
470 raise ValueError(
471 f"train_neg_sample_args:[{self.final_config_dict['train_neg_sample_args']}] should be a dict."
472 )
473 for op_args, op_args_values in default_train_neg_sample_args.items():
474 if op_args not in self.final_config_dict["train_neg_sample_args"]:
475 self.final_config_dict["train_neg_sample_args"][op_args] = op_args_values
477 # eval_args checking
478 default_eval_args = {
479 "split": {"RS": [0.8, 0.1, 0.1]},
480 "order": "RO",
481 "group_by": "user",
482 "mode": {"valid": "full", "test": "full"},
483 }
484 if not isinstance(self.final_config_dict["eval_args"], dict):
485 raise ValueError(f"eval_args:[{self.final_config_dict['eval_args']}] should be a dict.")
487 deep_dict_update(default_eval_args, self.final_config_dict["eval_args"])
489 mode = default_eval_args["mode"]
490 # backward compatible
491 if isinstance(mode, str):
492 default_eval_args["mode"] = {"valid": mode, "test": mode}
494 # in case there is only one key in `mode`, e.g., mode: {'valid': 'uni100'} or mode: {'test': 'full'}
495 if isinstance(mode, dict):
496 default_mode = mode.get("valid", mode.get("test", "full"))
497 default_eval_args["mode"] = {
498 "valid": mode.get("valid", default_mode),
499 "test": mode.get("test", default_mode),
500 }
502 self.final_config_dict["eval_args"] = default_eval_args
503 if (
504 self.final_config_dict["eval_type"] == EvaluatorType.VALUE
505 and "full" in self.final_config_dict["eval_args"]["mode"].values()
506 ):
507 raise NotImplementedError("Full sort evaluation do not match value-based metrics!")
509 if (
510 self.final_config_dict["MODEL_TYPE"] == ModelType.PATH_LANGUAGE_MODELING
511 and self.final_config_dict.get("context_length") is None
512 ):
513 if self.final_config_dict.get("path_hop_length") is None:
514 raise ValueError("Path language modeling requires path_hop_length to be specified.")
516 # 2 * path_hop_length + 1(U) + BOS + EOS
517 self.final_config_dict["context_length"] = (self.final_config_dict["path_hop_length"] * 2) + 3
519 default_path_sample_args = {
520 "temporal_causality": False,
521 "collaborative_path": True,
522 "strategy": "constrained-rw",
523 "path_token_separator": " ",
524 "restrict_by_phase": False,
525 "MAX_CONSECUTIVE_INVALID": 10,
526 "MAX_RW_TRIES_PER_IID": 1,
527 "MAX_RW_PATHS_PER_HOP": 1,
528 }
529 if not isinstance(self.final_config_dict.get("path_sample_args", {}), dict):
530 raise ValueError(f"path_sample_args:[{self.final_config_dict['path_sample_args']}] should be a dict.")
532 deep_dict_update(default_path_sample_args, self.final_config_dict.get("path_sample_args", {}))
533 if default_path_sample_args["temporal_causality"] and not default_path_sample_args["restrict_by_phase"]:
534 default_path_sample_args["restrict_by_phase"] = True
535 logger.warning("Since temporal_causality is True so restrict_by_phase has been automatically set to True.")
537 self.final_config_dict["path_sample_args"] = default_path_sample_args
539 def _init_device(self):
540 if isinstance(self.final_config_dict["gpu_id"], tuple):
541 self.final_config_dict["gpu_id"] = ",".join(map(str, list(self.final_config_dict["gpu_id"])))
542 else:
543 self.final_config_dict["gpu_id"] = str(self.final_config_dict["gpu_id"])
544 gpu_id = self.final_config_dict["gpu_id"]
545 os.environ["CUDA_VISIBLE_DEVICES"] = gpu_id
546 import torch
548 if "local_rank" not in self.final_config_dict:
549 self.final_config_dict["single_spec"] = True
550 self.final_config_dict["local_rank"] = 0
551 self.final_config_dict["device"] = (
552 torch.device("cpu") if len(gpu_id) == 0 or not torch.cuda.is_available() else torch.device("cuda")
553 )
554 else:
555 assert len(gpu_id.split(",")) >= self.final_config_dict["nproc"]
556 torch.distributed.init_process_group(
557 backend="nccl",
558 rank=self.final_config_dict["local_rank"] + self.final_config_dict["offset"],
559 world_size=self.final_config_dict["world_size"],
560 init_method="tcp://" + self.final_config_dict["ip"] + ":" + str(self.final_config_dict["port"]),
561 )
562 self.final_config_dict["device"] = torch.device("cuda", self.final_config_dict["local_rank"])
563 self.final_config_dict["single_spec"] = False
564 torch.cuda.set_device(self.final_config_dict["local_rank"])
565 if self.final_config_dict["local_rank"] != 0:
566 self.final_config_dict["state"] = "error"
567 self.final_config_dict["show_progress"] = False
568 self.final_config_dict["verbose"] = False
570 def _set_train_neg_sample_args(self):
571 train_neg_sample_args = self.final_config_dict.get("train_neg_sample_args")
572 if train_neg_sample_args is None or train_neg_sample_args == "None":
573 self.final_config_dict["train_neg_sample_args"] = {
574 "distribution": "none",
575 "sample_num": "none",
576 "alpha": "none",
577 "dynamic": False,
578 "candidate_num": 0,
579 }
580 else:
581 if not isinstance(train_neg_sample_args, dict):
582 raise ValueError(f"train_neg_sample_args:[{train_neg_sample_args}] should be a dict.")
584 distribution = train_neg_sample_args["distribution"]
585 if distribution is None or distribution == "None":
586 self.final_config_dict["train_neg_sample_args"] = {
587 "distribution": "none",
588 "sample_num": "none",
589 "alpha": "none",
590 "dynamic": False,
591 "candidate_num": 0,
592 }
593 elif distribution not in ["uniform", "popularity"]:
594 raise ValueError(
595 f"The distribution [{distribution}] of train_neg_sample_args should in ['uniform', 'popularity']"
596 )
598 def _set_eval_neg_sample_args(self, phase: Literal["valid", "test"]):
599 eval_mode = self.final_config_dict["eval_args"]["mode"][phase]
600 if not isinstance(eval_mode, str):
601 raise ValueError(f"mode [{eval_mode}] in eval_args should be a str.")
602 if eval_mode == "labeled":
603 eval_neg_sample_args = {"distribution": "none", "sample_num": "none"}
604 elif eval_mode == "full":
605 eval_neg_sample_args = {"distribution": "uniform", "sample_num": "none"}
606 elif eval_mode[0:3] == "uni":
607 sample_num = int(eval_mode[3:])
608 eval_neg_sample_args = {"distribution": "uniform", "sample_num": sample_num}
609 elif eval_mode[0:3] == "pop":
610 sample_num = int(eval_mode[3:])
611 eval_neg_sample_args = {
612 "distribution": "popularity",
613 "sample_num": sample_num,
614 }
615 else:
616 raise ValueError(f"the mode [{eval_mode}] in eval_args is not supported.")
617 self.final_config_dict[f"{phase}_neg_sample_args"] = eval_neg_sample_args
619 def _set_torch_dtype(self):
620 """
621 Convert a string dtype to a torch dtype.
622 """
623 import torch
625 weight_precision = self.final_config_dict.get("weight_precision", "float32")
627 if weight_precision == "float32":
628 weight_precision = torch.float32
629 elif weight_precision == "float16":
630 weight_precision = torch.float16
631 elif weight_precision == "bfloat16":
632 weight_precision = torch.bfloat16
633 else:
634 raise ValueError(f"Unsupported weight_precision: {weight_precision}")
636 self.final_config_dict["weight_precision"] = weight_precision
638 def _set_env_behavior(self):
639 """
640 Set behavior of utilities or similar libraries based on environment variables.
641 """
642 # Updates global progress bar API based on config
643 import hopwise.utils.logger as logger_module
645 progress_bar = logger_module.ProgressBar(self.final_config_dict.get("progress_bar_rich", True))
646 setattr(logger_module, "_progress_bar", progress_bar)
648 def __setitem__(self, key, value):
649 if not isinstance(key, str):
650 raise TypeError("index must be a str.")
651 self.final_config_dict[key] = value
653 def __getattr__(self, item):
654 if "final_config_dict" not in self.__dict__:
655 raise AttributeError("'Config' object has no attribute 'final_config_dict'")
656 if item in self.final_config_dict:
657 return self.final_config_dict[item]
658 raise AttributeError(f"'Config' object has no attribute '{item}'")
660 def __getitem__(self, item):
661 return self.final_config_dict.get(item)
663 def __contains__(self, key):
664 if not isinstance(key, str):
665 raise TypeError("index must be a str.")
666 return key in self.final_config_dict
668 def __str__(self):
669 args_info = "\n"
670 for category in self.parameters:
671 args_info += set_color(category + " Hyper Parameters:\n", "magenta")
672 args_info += "\n".join(
673 [
674 (set_color("{}", "cyan") + " =" + set_color(" {}", "yellow")).format(arg, value)
675 for arg, value in self.final_config_dict.items()
676 if arg in self.parameters[category]
677 ]
678 )
679 args_info += "\n\n"
681 args_info += set_color("Other Hyper Parameters: \n", "magenta")
682 args_info += "\n".join(
683 [
684 (set_color("{}", "cyan") + " = " + set_color("{}", "yellow")).format(arg, value)
685 for arg, value in self.final_config_dict.items()
686 if arg
687 not in {_ for args in self.parameters.values() for _ in args}.union(
688 {"model", "dataset", "config_files"}
689 )
690 ]
691 )
692 args_info += "\n\n"
693 return args_info
695 def __repr__(self):
696 return self.__str__()
698 def compatibility_settings(self):
699 import numpy as np
701 np.bool = np.bool_
702 np.int = np.int_
703 if np.__version__.startswith("2."):
704 np.float = np.float64
705 np.complex = np.complex64
706 np.unicode = np.str_
707 np.unicode_ = np.unicode
708 else:
709 np.float = np.float_
710 np.complex = np.complex_
711 np.unicode = np.unicode_
713 np.object = np.object_
714 np.str = np.str_
715 with warnings.catch_warnings():
716 warnings.simplefilter("ignore")
717 if not hasattr(np, "long"):
718 np.long = np.int_
720 if not hasattr(np, "string_"):
721 np.string_ = np.bytes_