diff --git a/optuna_dashboard/_cached_extra_study_property.py b/optuna_dashboard/_cached_extra_study_property.py index 1e4fb1a7..919853f4 100644 --- a/optuna_dashboard/_cached_extra_study_property.py +++ b/optuna_dashboard/_cached_extra_study_property.py @@ -1,6 +1,7 @@ from __future__ import annotations import copy +import numbers import threading from typing import Any from typing import List @@ -9,7 +10,6 @@ from typing import Set from typing import Tuple from typing import TYPE_CHECKING -import numpy as np from optuna.distributions import BaseDistribution from optuna.trial import FrozenTrial from optuna.trial import TrialState @@ -86,11 +86,11 @@ class _CachedExtraStudyProperty: self._cursor = next_cursor - def _is_sortable_value(self, v: Any) -> bool: - return not isinstance(v, bool) and isinstance(v, (int, float, np.integer, np.floating)) - def _update_user_attrs(self, trial: FrozenTrial) -> None: - current_user_attrs = {k: self._is_sortable_value(v) for k, v in trial.user_attrs.items()} + current_user_attrs = { + k: not isinstance(v, bool) and isinstance(v, numbers.Real) + for k, v in trial.user_attrs.items() + } for attr_name, current_is_sortable in current_user_attrs.items(): is_sortable = self._union_user_attrs.get(attr_name) if is_sortable is None: diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index 2730c328..dbcdc991 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -2,11 +2,11 @@ from __future__ import annotations from datetime import datetime import json +import numbers from typing import Any from typing import TYPE_CHECKING from typing import Union -import numpy as np from optuna.distributions import BaseDistribution from optuna.distributions import CategoricalDistribution from optuna.study import StudySummary @@ -104,8 +104,8 @@ def serialize_attrs(attrs: dict[str, Any]) -> list[Attribute]: value = "" elif isinstance(v, str): value = v - elif isinstance(v, (np.floating, np.integer)): - value = str(v.item()) + elif isinstance(v, numbers.Real): + value = str(v) else: value = json.dumps(v) value = value[:MAX_ATTR_LENGTH] if len(value) > MAX_ATTR_LENGTH else value