mirror of
https://github.com/wassname/optuna-dashboard.git
synced 2026-09-09 11:28:14 +08:00
Support numpy scalars in Trial.user_attrs
This commit is contained in:
@@ -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