Merge pull request #706 from toshihikoyanase/support-numpy-scalar-in-user-attrs

Support `numpy` scalars in `Trial.user_attrs`
This commit is contained in:
c-bata
2023-11-27 14:28:13 +09:00
committed by GitHub
4 changed files with 53 additions and 4 deletions
@@ -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:
+27
View File
@@ -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()