From ecf5ab80f7ed3a858b51640ec75e886dd06b828f Mon Sep 17 00:00:00 2001 From: Toshihiko Yanase Date: Wed, 22 Nov 2023 11:48:22 +0900 Subject: [PATCH 1/5] Support numpy scalars in Trial.user_attrs --- .../_cached_extra_study_property.py | 11 ++++---- optuna_dashboard/_serializer.py | 2 ++ .../test_cached_extra_study_property.py | 23 ++++++++++++++-- python_tests/test_serializers.py | 27 +++++++++++++++++++ 4 files changed, 56 insertions(+), 7 deletions(-) diff --git a/optuna_dashboard/_cached_extra_study_property.py b/optuna_dashboard/_cached_extra_study_property.py index 2f27fa64..1e4fb1a7 100644 --- a/optuna_dashboard/_cached_extra_study_property.py +++ b/optuna_dashboard/_cached_extra_study_property.py @@ -2,12 +2,14 @@ from __future__ import annotations import copy import threading +from typing import Any from typing import List from typing import Optional 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 @@ -84,12 +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: - # TODO(c-bata): Support numpy-specific number types. - current_user_attrs = { - k: not isinstance(v, bool) and isinstance(v, (int, float)) - for k, v in trial.user_attrs.items() - } + current_user_attrs = {k: self._is_sortable_value(v) 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 b5a3b305..9b0cce64 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -104,6 +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 = v.item() else: value = json.dumps(v) value = value[:MAX_ATTR_LENGTH] if len(value) > MAX_ATTR_LENGTH else value diff --git a/python_tests/test_cached_extra_study_property.py b/python_tests/test_cached_extra_study_property.py index c02006a7..dcd5bc5e 100644 --- a/python_tests/test_cached_extra_study_property.py +++ b/python_tests/test_cached_extra_study_property.py @@ -4,6 +4,7 @@ from typing import Any from unittest import TestCase import warnings +import numpy as np import optuna from optuna import create_trial from optuna.distributions import BaseDistribution @@ -254,11 +255,29 @@ class _CachedExtraStudyPropertyUserAttrs(TestCase): def test_infer_sortable(self) -> None: user_attrs_list: list[dict[str, Any]] = [ - {"a": 1, "b": 1, "c": 1, "d": "a", "e": 1, "f": True}, + { + "a": 1, + "b": 1, + "c": 1, + "d": "a", + "e": 1, + "f": True, + "g": np.float128(1.1), + "h": np.int64(2), + }, {"a": 2, "b": "a", "c": "a", "d": "a"}, {"a": 3, "b": None, "c": 3, "d": "a", "e": 3}, ] - expected = {"a": True, "b": False, "c": False, "d": False, "e": True, "f": False} + expected = { + "a": True, + "b": False, + "c": False, + "d": False, + "e": True, + "f": False, + "g": True, + "h": True, + } trials = [] for user_attrs in user_attrs_list: diff --git a/python_tests/test_serializers.py b/python_tests/test_serializers.py index 0d1546d2..8c64c5d8 100644 --- a/python_tests/test_serializers.py +++ b/python_tests/test_serializers.py @@ -2,6 +2,7 @@ from __future__ import annotations import sys +import numpy as np import optuna from optuna_dashboard._serializer import serialize_attrs from optuna_dashboard._serializer import serialize_study_detail @@ -25,6 +26,32 @@ def test_serialize_dict() -> None: assert len(serialized) <= 1 +def test_serialize_numpy_integer() -> None: + serialized = serialize_attrs( + { + "int8": np.int8(1), + "int16": np.int16(1), + "int32": np.int32(1), + "int64": np.int64(1), + } + ) + assert len(serialized) == 4 + assert all([v["value"] == 1 for v in serialized]) + + +def test_serialize_numpy_floating() -> None: + serialized = serialize_attrs( + { + "float16": np.float16(1.0), + "float32": np.float32(1.0), + "float64": np.float64(1.0), + "float128": np.float128(1.0), + } + ) + assert len(serialized) == 4 + assert all([v["value"] == 1.0 for v in serialized]) + + @pytest.mark.skipif(sys.version_info < (3, 8), reason="BoTorch dropped Python3.7 support") def test_get_study_detail_is_preferential() -> None: storage = optuna.storages.InMemoryStorage() From 5810ef7760a12164f480e10e5f4f2ed2858a458d Mon Sep 17 00:00:00 2001 From: Toshihiko Yanase Date: Wed, 22 Nov 2023 12:08:35 +0900 Subject: [PATCH 2/5] Cast to str. --- optuna_dashboard/_serializer.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index 9b0cce64..2730c328 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -105,7 +105,7 @@ def serialize_attrs(attrs: dict[str, Any]) -> list[Attribute]: elif isinstance(v, str): value = v elif isinstance(v, (np.floating, np.integer)): - value = v.item() + value = str(v.item()) else: value = json.dumps(v) value = value[:MAX_ATTR_LENGTH] if len(value) > MAX_ATTR_LENGTH else value From 0aec86dbb1db6fe50bb80df8dea4c633947efca2 Mon Sep 17 00:00:00 2001 From: Toshihiko Yanase Date: Wed, 22 Nov 2023 13:30:02 +0900 Subject: [PATCH 3/5] Fix expected values. --- python_tests/test_serializers.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/python_tests/test_serializers.py b/python_tests/test_serializers.py index 8c64c5d8..d1bdf59b 100644 --- a/python_tests/test_serializers.py +++ b/python_tests/test_serializers.py @@ -36,7 +36,7 @@ def test_serialize_numpy_integer() -> None: } ) assert len(serialized) == 4 - assert all([v["value"] == 1 for v in serialized]) + assert all([v["value"] == "1" for v in serialized]) def test_serialize_numpy_floating() -> None: @@ -49,7 +49,7 @@ def test_serialize_numpy_floating() -> None: } ) assert len(serialized) == 4 - assert all([v["value"] == 1.0 for v in serialized]) + assert all([v["value"] == "1.0" for v in serialized]) @pytest.mark.skipif(sys.version_info < (3, 8), reason="BoTorch dropped Python3.7 support") From 7401fae2d66e2e5b73f298293bb34e056582f821 Mon Sep 17 00:00:00 2001 From: Toshihiko Yanase Date: Wed, 22 Nov 2023 13:59:40 +0900 Subject: [PATCH 4/5] Check if value is numbers.Real --- optuna_dashboard/_cached_extra_study_property.py | 10 +++++----- optuna_dashboard/_serializer.py | 6 +++--- 2 files changed, 8 insertions(+), 8 deletions(-) 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 From ad4d29f70b2f6c2c7e6fbca52f75b9264ac011ae Mon Sep 17 00:00:00 2001 From: Toshihiko Yanase Date: Wed, 22 Nov 2023 15:22:23 +0900 Subject: [PATCH 5/5] Fix import lines --- optuna_dashboard/_cached_extra_study_property.py | 1 - optuna_dashboard/_serializer.py | 1 + 2 files changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/_cached_extra_study_property.py b/optuna_dashboard/_cached_extra_study_property.py index 919853f4..24e22372 100644 --- a/optuna_dashboard/_cached_extra_study_property.py +++ b/optuna_dashboard/_cached_extra_study_property.py @@ -3,7 +3,6 @@ from __future__ import annotations import copy import numbers import threading -from typing import Any from typing import List from typing import Optional from typing import Set diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index dbcdc991..7030abec 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -7,6 +7,7 @@ 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