class AutoTuner(Builder):
"""
[`Builder`][solverpy_learn.builder.builder.Builder] that trains a model by
running Optuna-based hyperparameter tuning (`autotune.prettytuner`) over
the training/development strategy files, then saves the best model.
"""
def __init__(
self,
setup: Setup,
tuneargs: (dict[str, Any] | None) = None,
templates: (list[str] | None) = None,
):
assert "evals" in setup
assert "dataname" in setup["evals"]
Builder.__init__(self, setup["evals"]["dataname"])
self._setup = setup
self._tuneargs: dict[str, Any] = TUNEARGS | (tuneargs or {})
self._templates = templates or []
def represent(self) -> dict[str, Any]:
return dict(
cls=f"{self.__class__.__module__}.{self.__class__.__name__}",
dataname=self._dataname,
tuneargs=self._tuneargs,
templates=self._templates,
)
def path(self, modelfile: str = "model.lgb") -> str:
if modelfile:
return os.path.join(super().path(), modelfile)
else:
return super().path()
def build(self, talker: Talker = Talker()) -> None:
trains = self._setup["evals"]
devels = self._setup["devels"] if "devels" in self._setup else trains
assert "plugin" in trains
assert "plugin" in devels
assert "refs" in trains
report = markdown.newline() + markdown.heading(f"Building model `{self._dataname}`", level=2)
reporter.add(report)
logger.info(f"Building model: {self._dataname}")
logger.debug(f'using trains: {trains["plugin"].path()}')
f_model = self.path()
if os.path.exists(f_model):
logger.info(f"Skipped model building; model {self._dataname} exists.")
self._strats = self.applies(trains["refs"], self._dataname)
return
f_train = trains["plugin"].path()
f_test = devels["plugin"].path()
use_builder = ("atpeval" in self._tuneargs) and self._tuneargs["atpeval"]
started_at = time.monotonic()
logger.info(f"Tunning learning params: train={f_train} test={f_test}")
logger.info(resource_summary("main", started_at))
logger.debug(usage(f"before tuning: {self._dataname}"))
try:
ret = autotune.prettytuner(
talker=talker,
f_train=f_train,
f_test=f_test,
d_tmp=self.path("opt"),
builder=self if use_builder else None,
**self._tuneargs,
)
except KeyboardInterrupt:
raise
finally:
logger.info(resource_summary("main", started_at))
logger.debug(usage(f"after tuning: {self._dataname}"))
#f_best = ret[3]
(_, _, _, f_best, _, _, pos, neg) = ret
(pos, neg) = (int(pos), int(neg))
shutil.copyfile(f_best, f_model)
#self._models = [f_model]
self._strats = self.applies(trains["refs"], self._dataname)
progress.build(self._dataname, *ret)
logger.info(f"Model {self._dataname} built.")