mirror of
https://github.com/wassname/optuna-dashboard.git
synced 2026-09-07 17:10:07 +08:00
Merge pull request #706 from toshihikoyanase/support-numpy-scalar-in-user-attrs
Support `numpy` scalars in `Trial.user_attrs`
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user