Coverage for hopwise/trainer/hyper_tuning.py: 82%

389 statements  

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

1# @Time : 2020/7/19 19:06 

2# @Author : Shanlei Mu 

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

4# @File : hyper_tuning.py 

5 

6# UPDATE: 

7# @Time : 2022/7/7, 2023/2/11 

8# @Author : Gaowei Zhang 

9# @Email : zgw15630559577@163.com 

10 

11# @Time : 2025 

12# @Author : Giacomo Medda 

13# @Email : giacomo.medda@unica.it 

14 

15"""hopwise.trainer.hyper_tuning 

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

17""" 

18 

19import os 

20import sys 

21from ast import literal_eval 

22from datetime import datetime 

23from enum import Enum 

24from functools import partial 

25 

26import numpy as np 

27 

28from hopwise.utils import dict2str 

29 

30 

31def _recursiveFindNodes(root, node_type="switch"): 

32 from hyperopt.pyll.base import Apply 

33 

34 nodes = [] 

35 if isinstance(root, (list, tuple)): 

36 for node in root: 

37 nodes.extend(_recursiveFindNodes(node, node_type)) 

38 elif isinstance(root, dict): 

39 for node in root.values(): 

40 nodes.extend(_recursiveFindNodes(node, node_type)) 

41 elif isinstance(root, (Apply)): 

42 if root.name == node_type: 

43 nodes.append(root) 

44 

45 for node in root.pos_args: 

46 if node.name == node_type: 

47 nodes.append(node) 

48 for _, node in root.named_args: 

49 if node.name == node_type: 

50 nodes.append(node) 

51 return nodes 

52 

53 

54def _parameters(space): 

55 # Analyze the domain instance to find parameters 

56 parameters = {} 

57 if isinstance(space, dict): 

58 space = list(space.values()) 

59 for node in _recursiveFindNodes(space, "switch"): 

60 # Find the name of this parameter 

61 paramNode = node.pos_args[0] 

62 assert paramNode.name == "hyperopt_param" 

63 paramName = paramNode.pos_args[0].obj 

64 

65 # Find all possible choices for this parameter 

66 values = [literal.obj for literal in node.pos_args[1:]] 

67 parameters[paramName] = np.array(range(len(values))) 

68 return parameters 

69 

70 

71def _spacesize(space): 

72 # Compute the number of possible combinations 

73 params = _parameters(space) 

74 return np.prod([len(values) for values in params.values()]) 

75 

76 

77class ExhaustiveSearchError(Exception): 

78 r"""ExhaustiveSearchError""" 

79 

80 pass 

81 

82 

83def exhaustive_search(new_ids, domain, trials, seed, nbMaxSucessiveFailures=1000): 

84 r"""This is for exhaustive search in HyperTuning.""" 

85 from hyperopt import pyll 

86 from hyperopt.base import miscs_update_idxs_vals 

87 

88 # Build a hash set for previous trials 

89 hashset = set( 

90 [ 

91 hash( 

92 frozenset( 

93 [ 

94 (key, value[0]) if len(value) > 0 else ((key, None)) 

95 for key, value in trial["misc"]["vals"].items() 

96 ] 

97 ) 

98 ) 

99 for trial in trials.trials 

100 ] 

101 ) 

102 

103 rng = np.random.RandomState(seed) 

104 rval = [] 

105 for _, new_id in enumerate(new_ids): 

106 newSample = False 

107 nbSucessiveFailures = 0 

108 while not newSample: 

109 # -- sample new specs, idxs, vals 

110 idxs, vals = pyll.rec_eval( 

111 domain.s_idxs_vals, 

112 memo={ 

113 domain.s_new_ids: [new_id], 

114 domain.s_rng: rng, 

115 }, 

116 ) 

117 new_result = domain.new_result() 

118 new_misc = dict(tid=new_id, cmd=domain.cmd, workdir=domain.workdir) 

119 miscs_update_idxs_vals([new_misc], idxs, vals) 

120 

121 # Compare with previous hashes 

122 h = hash(frozenset([(key, value[0]) if len(value) > 0 else ((key, None)) for key, value in vals.items()])) 

123 if h not in hashset: 

124 newSample = True 

125 else: 

126 # Duplicated sample, ignore 

127 nbSucessiveFailures += 1 

128 

129 if nbSucessiveFailures > nbMaxSucessiveFailures: 

130 # No more samples to produce 

131 return [] 

132 

133 rval.extend(trials.new_trial_docs([new_id], [None], [new_result], [new_misc])) 

134 return rval 

135 

136 

137class HyperTuning: 

138 r"""HyperTuning Class is used to manage the parameter tuning process of recommender system models. 

139 Given objective funciton, parameters range and optimization algorithm, using HyperTuning can find 

140 the best result among these parameters. 

141 

142 Note: 

143 HyperTuning provides three tuner tools: 

144 - hyperopt (https://github.com/hyperopt/hyperopt) 

145 - ray (https://docs.ray.io/en/latest/tune/index.html) 

146 - optuna (https://optuna.org/) 

147 

148 Thanks to sbrodeur for the exhaustive search code. 

149 https://github.com/hyperopt/hyperopt/issues/200 

150 """ 

151 

152 PARAMS_PER_ROW = 3 

153 TUNER_TYPES = Enum("TUNER_TYPES", {"HYPEROPT": "hyperopt", "RAY": "ray", "OPTUNA": "optuna"}) 

154 

155 def __init__( 

156 self, 

157 objective_function, 

158 tuner="optuna", 

159 space=None, 

160 params_file=None, 

161 params_dict=None, 

162 fixed_config_file_list=None, 

163 display_file=None, 

164 algo=None, 

165 max_evals=100, 

166 early_stop=10, 

167 output_path=None, 

168 timeout=None, 

169 show_progress=False, 

170 study_name=None, 

171 resume=False, 

172 ): 

173 self.tuner = self.TUNER_TYPES[tuner.upper()] 

174 self.best_score = None 

175 self.best_params = None 

176 self.best_test_result = None 

177 self.params2result = {} 

178 self.params_list = [] 

179 self.score_list = [] 

180 

181 self.show_progress = show_progress 

182 self.objective_function = objective_function 

183 self.max_evals = max_evals 

184 self.timeout = timeout 

185 self.fixed_config_file_list = fixed_config_file_list 

186 self.display_file = display_file 

187 self.output_path = output_path or "." 

188 if not os.path.exists(self.output_path): 

189 os.makedirs(self.output_path) 

190 self.study_name = study_name or f"hyper_{datetime.now().strftime('%d_%m_%Y_%H_%M_%S')}" 

191 self.resume = resume 

192 

193 if space: 

194 self.space = space 

195 elif params_file: 

196 self.space = self.build_space_from_file(params_file) 

197 elif params_dict: 

198 self.space = self.build_space_from_dict(params_dict) 

199 else: 

200 raise ValueError("at least one of `space`, `params_file` and `params_dict` is provided") 

201 

202 self.select_algo(algo) 

203 self.select_early_stop(early_stop) 

204 

205 def select_algo(self, algo): 

206 r"""Select the algorithm for hyperparameter tuning 

207 Args: 

208 algo (str or callable): the algorithm name or function 

209 """ 

210 if self.tuner == self.TUNER_TYPES.HYPEROPT: 

211 if algo is None: 

212 self.algo = partial(exhaustive_search, nbMaxSucessiveFailures=1000) 

213 self.max_evals = _spacesize(self.space) 

214 elif isinstance(algo, str): 

215 from hyperopt import anneal, rand, tpe 

216 

217 if algo == "exhaustive": 

218 self.algo = partial(exhaustive_search, nbMaxSucessiveFailures=1000) 

219 self.max_evals = _spacesize(self.space) 

220 elif algo == "random": 

221 self.algo = rand.suggest 

222 elif algo == "bayes": 

223 self.algo = tpe.suggest 

224 elif algo == "anneal": 

225 if sys.version_info >= (3, 12): 

226 raise RuntimeError( 

227 "hyperopt's `anneal` algorithm is not supported on Python >= 3.12: its " 

228 "sampler calls int() on a 1-d numpy array, which numpy>=2 (bundled with " 

229 "Python 3.12) rejects with 'TypeError: only 0-dimensional arrays can be " 

230 "converted to Python scalars'. This is an upstream hyperopt bug present in " 

231 "all current releases. Use a different algo (e.g. 'bayes', 'random', " 

232 "'exhaustive'), or run on Python < 3.12 with numpy < 2." 

233 ) 

234 self.algo = anneal.suggest 

235 else: 

236 raise ValueError(f"Illegal algo [{algo}]") 

237 else: 

238 self.algo = algo 

239 elif self.tuner == self.TUNER_TYPES.RAY: 

240 from ray.tune import schedulers, search 

241 

242 if algo is None: 

243 self.algo = {"search_alg": None, "scheduler": "async_hyperband"} 

244 elif isinstance(algo, str): 

245 if "-" in algo: 

246 search_alg, scheduler = algo.split("-") 

247 search_alg = search.SEARCH_ALG_IMPORT.get(search_alg, None) 

248 self.algo = {"search_alg": search_alg, "scheduler": scheduler} 

249 else: 

250 search_alg = search.SEARCH_ALG_IMPORT.get(algo, None) 

251 scheduler = algo if algo in schedulers.SCHEDULER_IMPORT else None 

252 self.algo = { 

253 "search_alg": search_alg, 

254 "scheduler": scheduler, 

255 } 

256 elif self.tuner == self.TUNER_TYPES.OPTUNA: 

257 import optuna 

258 

259 if algo is None: 

260 self.algo = { 

261 "sampler": None, 

262 "pruner": optuna.pruners.MedianPruner(), 

263 } 

264 elif isinstance(algo, str): 

265 grid_space = {k: (v[2] if v[0] == "choice" else list(v[2:])) for k, v in self.space.items()} 

266 

267 if "-" in algo: 

268 sampler, pruner = algo.split("-") 

269 sampler_args = [grid_space] if sampler == "GridSampler" else [] 

270 self.algo = { 

271 "sampler": getattr(optuna.samplers, sampler)(*sampler_args), 

272 "pruner": getattr(optuna.pruners, pruner)(), 

273 } 

274 else: 

275 sampler = getattr(optuna.samplers, algo) if hasattr(optuna.samplers, algo) else None 

276 pruner = getattr(optuna.pruners, algo) if hasattr(optuna.pruners, algo) else None 

277 if sampler is not None and sampler is optuna.samplers.GridSampler: 

278 sampler_args = [grid_space] 

279 else: 

280 sampler_args = [] 

281 self.algo = { 

282 "sampler": sampler(*sampler_args) if sampler is not None else None, 

283 "pruner": pruner() if pruner is not None else None, 

284 } 

285 

286 def select_early_stop(self, early_stop_steps): 

287 from hyperopt.early_stop import no_progress_loss 

288 

289 self.early_stop_fn = no_progress_loss(early_stop_steps) 

290 

291 def _get_tuner_distributions(self): 

292 if self.tuner == self.TUNER_TYPES.HYPEROPT: 

293 from hyperopt import hp 

294 

295 def choice(name, values): 

296 return hp.choice(name, values) 

297 

298 def uniform(name, low, high): 

299 return hp.uniform(name, low, high) 

300 

301 def quniform(name, low, high, q): 

302 return hp.quniform(name, low, high, q) 

303 

304 def loguniform(name, low, high): 

305 return hp.loguniform(name, low, high) 

306 elif self.tuner == self.TUNER_TYPES.RAY: 

307 from ray import tune 

308 

309 def choice(name, values): 

310 return tune.choice(values) 

311 

312 def uniform(name, low, high): 

313 return tune.uniform(low, high) 

314 

315 def quniform(name, low, high, q): 

316 return tune.quniform(low, high, q) 

317 

318 def loguniform(name, low, high): 

319 return tune.uniform(np.exp(low), np.exp(high)) 

320 elif self.tuner == self.TUNER_TYPES.OPTUNA: 

321 

322 def choice(name, values): 

323 return "choice", name, values 

324 

325 def uniform(name, low, high): 

326 return "uniform", name, low, high 

327 

328 def quniform(name, low, high, q): 

329 return "quniform", name, low, high, q 

330 

331 def loguniform(name, low, high): 

332 return "loguniform", name, low, high 

333 

334 return choice, uniform, quniform, loguniform 

335 

336 def build_space_from_file(self, file): 

337 choice, uniform, quniform, loguniform = self._get_tuner_distributions() 

338 

339 space = self._build_space_from_file( 

340 file, 

341 choice_fn=choice, 

342 uniform_fn=uniform, 

343 quniform_fn=quniform, 

344 loguniform_fn=loguniform, 

345 ) 

346 

347 return space 

348 

349 def build_space_from_dict(self, config_dict): 

350 choice, uniform, quniform, loguniform = self._get_tuner_distributions() 

351 

352 space = self._build_space_from_dict( 

353 config_dict, 

354 choice_fn=choice, 

355 uniform_fn=uniform, 

356 quniform_fn=quniform, 

357 loguniform_fn=loguniform, 

358 ) 

359 

360 return space 

361 

362 @staticmethod 

363 def _build_space_from_file( 

364 file, 

365 choice_fn=None, 

366 uniform_fn=None, 

367 quniform_fn=None, 

368 loguniform_fn=None, 

369 ): 

370 config_dict = {} 

371 with open(file) as fp: 

372 for line in fp: 

373 para_list = line.strip().split(" ") 

374 if len(para_list) < HyperTuning.PARAMS_PER_ROW: 

375 continue 

376 para_name, para_type, para_value = ( 

377 para_list[0], 

378 para_list[1], 

379 "".join(para_list[2:]), 

380 ) 

381 if para_type == "choice": 

382 config_dict.setdefault("choice", {}) 

383 config_dict["choice"][para_name] = literal_eval(para_value) 

384 elif para_type == "uniform": 

385 config_dict.setdefault("uniform", {}) 

386 low, high = para_value.strip().split(",") 

387 config_dict["uniform"][para_name] = (float(low), float(high)) 

388 elif para_type == "quniform": 

389 config_dict.setdefault("quniform", {}) 

390 low, high, q = para_value.strip().split(",") 

391 config_dict["quniform"][para_name] = (float(low), float(high), float(q)) 

392 elif para_type == "loguniform": 

393 config_dict.setdefault("loguniform", {}) 

394 low, high = para_value.strip().split(",") 

395 config_dict["loguniform"][para_name] = (float(low), float(high)) 

396 else: 

397 raise ValueError(f"Illegal param type [{para_type}]") 

398 

399 space = HyperTuning._build_space_from_dict( 

400 config_dict, 

401 choice_fn=choice_fn, 

402 uniform_fn=uniform_fn, 

403 quniform_fn=quniform_fn, 

404 loguniform_fn=loguniform_fn, 

405 ) 

406 return space 

407 

408 @staticmethod 

409 def _build_space_from_dict( 

410 config_dict, 

411 choice_fn=None, 

412 uniform_fn=None, 

413 quniform_fn=None, 

414 loguniform_fn=None, 

415 ): 

416 space = {} 

417 for para_type in config_dict: 

418 if para_type == "choice": 

419 for para_name in config_dict["choice"]: 

420 para_value = config_dict["choice"][para_name] 

421 space[para_name] = choice_fn(para_name, para_value) 

422 elif para_type == "uniform": 

423 for para_name in config_dict["uniform"]: 

424 para_value = config_dict["uniform"][para_name] 

425 low = para_value[0] 

426 high = para_value[1] 

427 space[para_name] = uniform_fn(para_name, float(low), float(high)) 

428 elif para_type == "quniform": 

429 for para_name in config_dict["quniform"]: 

430 para_value = config_dict["quniform"][para_name] 

431 low = para_value[0] 

432 high = para_value[1] 

433 q = para_value[2] 

434 space[para_name] = quniform_fn(para_name, float(low), float(high), float(q)) 

435 elif para_type == "loguniform": 

436 for para_name in config_dict["loguniform"]: 

437 para_value = config_dict["loguniform"][para_name] 

438 low = para_value[0] 

439 high = para_value[1] 

440 space[para_name] = loguniform_fn(para_name, float(low), float(high)) 

441 else: 

442 raise ValueError(f"Illegal param type [{para_type}]") 

443 return space 

444 

445 def build_optuna_space(self, trial): 

446 r"""Build the space for optuna 

447 

448 Args: 

449 trial (optuna.trial): the trial object 

450 """ 

451 params = {} 

452 for para_name in self.space: 

453 para_type, _, *para_value = self.space[para_name] 

454 if para_type == "choice": 

455 para_value = para_value[0] 

456 params[para_name] = trial.suggest_categorical(para_name, para_value) 

457 elif para_type == "uniform": 

458 low = para_value[0] 

459 high = para_value[1] 

460 params[para_name] = trial.suggest_float(para_name, low, high) 

461 elif para_type == "quniform": 

462 low = para_value[0] 

463 high = para_value[1] 

464 q = para_value[2] 

465 params[para_name] = trial.suggest_float(para_name, low, high, step=q) 

466 elif para_type == "loguniform": 

467 low = para_value[0] 

468 high = para_value[1] 

469 

470 params[para_name] = np.exp(trial.suggest_float(para_name, low, high)) 

471 else: 

472 raise ValueError(f" Illegal param type [{para_type}]") 

473 

474 return params 

475 

476 @staticmethod 

477 def params2str(params): 

478 r"""Convert dict to str 

479 

480 Args: 

481 params (dict): parameters dict 

482 Returns: 

483 str: parameters string 

484 """ 

485 params_str = "" 

486 for param_name in params: 

487 params_str += param_name + ":" + str(params[param_name]) + ", " 

488 return params_str[:-2] 

489 

490 @staticmethod 

491 def _print_result(result_dict: dict): 

492 print("current best valid score: %.4f" % result_dict["best_valid_score"]) 

493 print("current best valid result:") 

494 print(result_dict["best_valid_result"]) 

495 print("current test result:") 

496 print(result_dict["test_result"]) 

497 print() 

498 

499 def export_result(self, output_path=None): 

500 r"""Write the searched parameters and corresponding results to the file 

501 

502 Args: 

503 output_path (str): the output file 

504 

505 """ 

506 output_path = output_path or self.output_path 

507 output_file = os.path.join(output_path, self.study_name + ".txt") 

508 

509 with open(output_file, "w") as fp: 

510 fp.write("***Best trial***\n") 

511 fp.write("Best valid score: %.4f\n" % self.best_score) 

512 fp.write("Best parameters: " + dict2str(self.best_params) + "\n") 

513 fp.write("Best valid result:\n" + dict2str(self.best_valid_result) + "\n") 

514 fp.write("Best test result:\n" + dict2str(self.best_test_result) + "\n\n") 

515 

516 for params in self.params2result: 

517 fp.write(params + "\n") 

518 fp.write("Valid result:\n" + dict2str(self.params2result[params]["best_valid_result"]) + "\n") 

519 

520 fp.write("Test result:\n" + dict2str(self.params2result[params]["test_result"]) + "\n\n") 

521 

522 if self.tuner == self.TUNER_TYPES.OPTUNA: 

523 if not hasattr(self, "study"): 

524 raise ValueError("Optuna study not created. Call `run` method first.") 

525 

526 optuna_df = self.study.trials_dataframe() 

527 fp.write("Optuna trials dataframe:\n") 

528 optuna_df.to_string(fp, index=False) 

529 

530 def trial(self, params): 

531 r"""Given a set of parameters, return results and optimization status 

532 

533 Args: 

534 params (dict): the parameter dictionary 

535 """ 

536 config_dict = params.copy() 

537 params_str = self.params2str(params) 

538 self.params_list.append(params_str) 

539 print("running parameters:", config_dict) 

540 result_dict = self.objective_function(config_dict, self.fixed_config_file_list, saved=False) 

541 self.params2result[params_str] = result_dict 

542 model, score, bigger = ( 

543 result_dict["model"], 

544 result_dict["best_valid_score"], 

545 result_dict["valid_score_bigger"], 

546 ) 

547 self.model = model 

548 self.score_list.append(score) 

549 

550 if not self.best_score or (bigger and score > self.best_score) or (not bigger and score < self.best_score): 

551 self.best_score = score 

552 self.best_params = params 

553 self.best_valid_result = result_dict["best_valid_result"] 

554 self.best_test_result = result_dict["test_result"] 

555 self._print_result(result_dict) 

556 

557 if bigger: 

558 score = -score 

559 

560 return {**result_dict, "hyper_score": score} 

561 

562 def plot_hyper(self): 

563 import pandas as pd 

564 import plotly.graph_objs as go 

565 from plotly.offline import plot 

566 

567 data_dict = {"valid_score": self.score_list, "params": self.params_list} 

568 trial_df = pd.DataFrame(data_dict) 

569 trial_df["trial_number"] = trial_df.index + 1 

570 trial_df["trial_number"] = trial_df["trial_number"].astype(dtype=np.str) 

571 

572 trace = go.Scatter( 

573 x=trial_df["trial_number"], 

574 y=trial_df["valid_score"], 

575 text=trial_df["params"], 

576 mode="lines+markers", 

577 marker=dict(color="green"), 

578 showlegend=True, 

579 textposition="top center", 

580 name=self.model + " tuning process", 

581 ) 

582 

583 data = [trace] 

584 layout = go.Layout( 

585 title="hyperparams_tuning", 

586 xaxis=dict(title="trials"), 

587 yaxis=dict(title="valid_score"), 

588 ) 

589 fig = go.Figure(data=data, layout=layout) 

590 

591 plot(fig, filename=self.display_file) 

592 

593 def run(self): 

594 r"""Begin to search the best parameters""" 

595 if self.resume: 

596 print("\n# Resume " + "-" * 40) 

597 print(f"Resuming from {os.path.join(self.output_path, self.study_name)} if exists.\n") 

598 

599 if self.tuner == self.TUNER_TYPES.HYPEROPT: 

600 import hyperopt 

601 

602 def hyperopt_objective(params): 

603 try: 

604 result_dict = self.trial(params) 

605 except Exception as e: 

606 print(f"Error occurred during trial: {e}") 

607 import traceback 

608 

609 traceback.print_exc() 

610 return {"loss": np.nan, "status": hyperopt.STATUS_FAIL} 

611 

612 return {"loss": result_dict["hyper_score"], "status": hyperopt.STATUS_OK} 

613 

614 if os.path.exists(os.path.join(self.output_path, self.study_name)) and not self.resume: 

615 raise FileExistsError( 

616 f"File {os.path.join(self.output_path, self.study_name)} already exists. " 

617 "Please remove it or set `--resume True` to continue." 

618 ) 

619 

620 hyperopt.fmin( 

621 hyperopt_objective, 

622 self.space, 

623 algo=self.algo, 

624 max_evals=self.max_evals, 

625 early_stop_fn=self.early_stop_fn, 

626 timeout=self.timeout, 

627 trials_save_file=os.path.join(self.output_path, self.study_name), 

628 ) 

629 elif self.tuner == self.TUNER_TYPES.RAY: 

630 import ray 

631 from ray import tune 

632 

633 def ray_objective(params): 

634 result_dict = self.trial(params) 

635 tune.report({"hyper_score": result_dict["hyper_score"]}) 

636 

637 return result_dict 

638 

639 if not ray.is_initialized(): 

640 # Don't let Ray snapshot the cwd into a working_dir: its packager applies 

641 # .gitignore excludes (e.g. `dataset/`) but drops the `!` re-includes, so 

642 # hopwise/properties/dataset/*.yaml is missing from the copy and trial workers 

643 # load an incomplete config (numerical_features=None -> crash). With no 

644 # working_dir, workers import the installed hopwise instead. 

645 ray.init(runtime_env={"working_dir": None}) 

646 tune.register_trainable("ray-trial", ray_objective) 

647 if self.algo["scheduler"] is not None: 

648 scheduler = tune.create_scheduler( 

649 self.algo["scheduler"], 

650 metric="hyper_score", 

651 mode="min", 

652 max_t=self.max_evals, 

653 grace_period=1, 

654 reduction_factor=2, 

655 ) 

656 else: 

657 scheduler = None 

658 

659 if self.algo["search_alg"] is not None: 

660 from ray.tune import search 

661 

662 if self.algo["search_alg"]() is search.BasicVariantGenerator: 

663 search_alg = search.BasicVariantGenerator() 

664 else: 

665 search_alg = self.algo["search_alg"]()( 

666 metric="hyper_score", 

667 mode="min", 

668 ) 

669 else: 

670 search_alg = None 

671 

672 tune.run( 

673 ray_objective, 

674 config=self.space, 

675 num_samples=self.max_evals, 

676 scheduler=scheduler, 

677 search_alg=search_alg, 

678 storage_path=self.output_path, 

679 name=self.study_name, 

680 log_to_file=self.study_name, 

681 resume=self.resume, 

682 ) 

683 elif self.tuner == self.TUNER_TYPES.OPTUNA: 

684 import optuna 

685 

686 try: 

687 self.study = optuna.create_study( 

688 direction="minimize", 

689 study_name=self.study_name, 

690 storage=f"sqlite:///{os.path.join(self.output_path, self.study_name)}.db", 

691 pruner=self.algo["pruner"], 

692 sampler=self.algo["sampler"], 

693 load_if_exists=self.resume, 

694 ) 

695 except optuna.exceptions.DuplicatedStudyError as e: 

696 raise optuna.exceptions.DuplicatedStudyError( 

697 f"Study {os.path.join(self.output_path, self.study_name)} already exists. " 

698 "Please use --resume True to load an existing study checkpoint." 

699 ) from e 

700 

701 def optuna_objective(trial): 

702 params = self.build_optuna_space(trial) 

703 

704 def trial_callback(epoch_idx, valid_score): 

705 trial.report(valid_score, epoch_idx) 

706 

707 if trial.should_prune(): 

708 raise optuna.exceptions.TrialPruned() 

709 

710 if isinstance(self.objective_function, partial): 

711 self.objective_function = partial( 

712 self.objective_function.func, 

713 callback_fn=trial_callback, 

714 ) 

715 else: 

716 self.objective_function = partial( 

717 self.objective_function, 

718 callback_fn=trial_callback, 

719 ) 

720 

721 result_dict = self.trial(params) 

722 

723 for key, value in result_dict["test_result"].items(): 

724 trial.set_user_attr(key, value) 

725 

726 return result_dict["hyper_score"] 

727 

728 self.study.optimize( 

729 optuna_objective, 

730 n_trials=self.max_evals, 

731 timeout=self.timeout, 

732 ) 

733 

734 if self.display_file is not None: 

735 self.plot_hyper()