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

1# @Time : 2020/6/28 

2# @Author : Zihan Lin 

3# @Email : linzihan.super@foxmail.com 

4 

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 

9 

10# @Time : 2025 

11# @Author : Giacomo Medda, Alessandro Soccol 

12# @Email : giacomo.medda@unica.it, alessandro.soccol@unica.it 

13 

14 

15"""hopwise.config.configurator 

16################################ 

17""" 

18 

19import os 

20import re 

21import sys 

22import warnings 

23from logging import getLogger 

24from typing import Literal 

25 

26import yaml 

27 

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) 

42 

43 

44class Config: 

45 """Configurator module that load the defined parameters. 

46 

47 Configurator module will first load the default parameters from the fixed properties in hopwise and then 

48 load parameters from the external input. 

49 

50 External input supports three kind of forms: config file, command line and parameter dictionaries. 

51 

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: 

54 

55 learning_rate: 0.001 

56 

57 train_batch_size: 2048 

58 

59 - command line: It should be in the format as '---learning_rate=0.001' 

60 

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} 

63 

64 Configuration module allows the above three kind of external input format to be used together, 

65 the priority order is as following: 

66 

67 command line > parameter dictionaries > config file 

68 

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. 

71 

72 Finally the learning_rate is equal to 0.02. 

73 """ 

74 

75 NESTED_KEY_SEPARATOR = "." 

76 

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

93 

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

104 

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 

111 

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 

129 

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] 

135 

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 

168 

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 

176 

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

182 

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 

202 

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 

209 

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) 

225 

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 

236 

237 return final_model, final_model_class, final_dataset 

238 

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 

245 

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

252 

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

266 

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 ] 

278 

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) 

305 

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) 

314 

315 # avoids parallelism issues with HuggingFace Tokenizers 

316 os.environ["TOKENIZERS_PARALLELISM"] = "false" 

317 

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 ] 

328 

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 

334 

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) 

343 

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

361 

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

365 

366 if train_stage is not None and train_stage in ["pretrain"]: 

367 self.final_config_dict["MODEL_INPUT_TYPE"] = InputType.PAIRWISE 

368 

369 metrics = self.final_config_dict["metrics"] 

370 if isinstance(metrics, str): 

371 self.final_config_dict["metrics"] = [metrics] 

372 

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

382 

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 ) 

388 

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 

391 

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

402 

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] 

413 

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

422 

423 self.final_config_dict["eval_type"] = eval_type_kg.pop() 

424 

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 ) 

430 

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

443 

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] 

448 

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 } 

457 

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 ) 

467 

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 

476 

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

486 

487 deep_dict_update(default_eval_args, self.final_config_dict["eval_args"]) 

488 

489 mode = default_eval_args["mode"] 

490 # backward compatible 

491 if isinstance(mode, str): 

492 default_eval_args["mode"] = {"valid": mode, "test": mode} 

493 

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 } 

501 

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

508 

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

515 

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 

518 

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

531 

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

536 

537 self.final_config_dict["path_sample_args"] = default_path_sample_args 

538 

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 

547 

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 

569 

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

583 

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 ) 

597 

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 

618 

619 def _set_torch_dtype(self): 

620 """ 

621 Convert a string dtype to a torch dtype. 

622 """ 

623 import torch 

624 

625 weight_precision = self.final_config_dict.get("weight_precision", "float32") 

626 

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

635 

636 self.final_config_dict["weight_precision"] = weight_precision 

637 

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 

644 

645 progress_bar = logger_module.ProgressBar(self.final_config_dict.get("progress_bar_rich", True)) 

646 setattr(logger_module, "_progress_bar", progress_bar) 

647 

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 

652 

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

659 

660 def __getitem__(self, item): 

661 return self.final_config_dict.get(item) 

662 

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 

667 

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" 

680 

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 

694 

695 def __repr__(self): 

696 return self.__str__() 

697 

698 def compatibility_settings(self): 

699 import numpy as np 

700 

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_ 

712 

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_ 

719 

720 if not hasattr(np, "string_"): 

721 np.string_ = np.bytes_