From c007b89e30525029e6e7dedaf50d9bb75c5beb1b Mon Sep 17 00:00:00 2001 From: c-bata Date: Sat, 7 Jan 2023 00:25:40 +0900 Subject: [PATCH] Fix broken tests at Optuna 2.10 --- optuna_dashboard/_serializer.py | 107 ++++++++++++++++++-------------- 1 file changed, 59 insertions(+), 48 deletions(-) diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index 0d8e6920..a20f6f01 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -8,8 +8,6 @@ from typing import Union import numpy as np from optuna.distributions import BaseDistribution from optuna.distributions import CategoricalDistribution -from optuna.distributions import FloatDistribution -from optuna.distributions import IntDistribution from optuna.study import StudySummary from optuna.trial import FrozenTrial @@ -209,23 +207,71 @@ def serialize_frozen_trial( def serialize_distribution(distribution: BaseDistribution) -> DistributionJSON: - distribution = normalize_distribution(distribution) - if isinstance(distribution, FloatDistribution): + if distribution.__class__.__name__ == "FloatDistribution": + # Added from Optuna v3.0 return { "type": "FloatDistribution", - "low": distribution.low, - "high": distribution.high, - "step": distribution.step, - "log": distribution.log, + "low": getattr(distribution, "low"), + "high": getattr(distribution, "high"), + "step": getattr(distribution, "step"), + "log": getattr(distribution, "log"), } - if isinstance(distribution, IntDistribution): + if distribution.__class__.__name__ == "UniformDistribution": + # Deprecated from Optuna v3.0 + return { + "type": "FloatDistribution", + "low": getattr(distribution, "low"), + "high": getattr(distribution, "high"), + "step": 0, + "log": False, + } + if distribution.__class__.__name__ == "LogUniformDistribution": + # Deprecated from Optuna v3.0 + return { + "type": "FloatDistribution", + "low": getattr(distribution, "low"), + "high": getattr(distribution, "high"), + "step": 0, + "log": True, + } + if distribution.__class__.__name__ == "DiscreteUniformDistribution": + # Deprecated from Optuna v3.0 + return { + "type": "FloatDistribution", + "low": getattr(distribution, "low"), + "high": getattr(distribution, "high"), + "step": getattr(distribution, "q"), + "log": False, + } + + if distribution.__class__.__name__ == "IntDistribution": + # Added from Optuna v3.0 return { "type": "IntDistribution", - "low": distribution.low, - "high": distribution.high, - "step": distribution.step, - "log": distribution.log, + "low": getattr(distribution, "low"), + "high": getattr(distribution, "high"), + "step": getattr(distribution, "step"), + "log": getattr(distribution, "log"), } + if distribution.__class__.__name__ == "IntUniformDistribution": + # Deprecated from Optuna v3.0 + return { + "type": "IntDistribution", + "low": getattr(distribution, "low"), + "high": getattr(distribution, "high"), + "step": getattr(distribution, "step"), + "log": False, + } + if distribution.__class__.__name__ == "IntLogUniformDistribution": + # Deprecated from Optuna v3.0 + return { + "type": "IntDistribution", + "low": getattr(distribution, "low"), + "high": getattr(distribution, "high"), + "step": getattr(distribution, "step"), + "log": True, + } + if isinstance(distribution, CategoricalDistribution): return { "type": "CategoricalDistribution", @@ -237,41 +283,6 @@ def serialize_distribution(distribution: BaseDistribution) -> DistributionJSON: raise ValueError(f"Unexpected distribution {str(distribution)}") -def normalize_distribution(distribution: BaseDistribution) -> BaseDistribution: - if distribution.__class__.__name__ == "UniformDistribution": - return FloatDistribution( - low=getattr(distribution, "low"), - high=getattr(distribution, "high"), - ) - elif distribution.__class__.__name__ == "LogUniformDistribution": - return FloatDistribution( - low=getattr(distribution, "low"), - high=getattr(distribution, "high"), - log=True, - ) - elif distribution.__class__.__name__ == "DiscreteUniformDistribution": - return FloatDistribution( - low=getattr(distribution, "low"), - high=getattr(distribution, "high"), - step=getattr(distribution, "q"), - ) - elif distribution.__class__.__name__ == "IntUniformDistribution": - return IntDistribution( - low=getattr(distribution, "low"), - high=getattr(distribution, "high"), - step=getattr(distribution, "step"), - ) - elif distribution.__class__.__name__ == "IntLogUniformDistribution": - return IntDistribution( - low=getattr(distribution, "low"), - high=getattr(distribution, "high"), - step=getattr(distribution, "step"), - log=True, - ) - else: - return distribution - - def serialize_search_space( search_space: list[tuple[str, BaseDistribution]] ) -> list[dict[str, Any]]: