Files
optuna-dashboard/optuna_dashboard/serializer.py
T
2021-04-04 15:23:40 +09:00

147 lines
4.3 KiB
Python

from typing import Any, Dict, List, Tuple
from optuna.distributions import BaseDistribution
from optuna.study import StudySummary
from optuna.trial import FrozenTrial
try:
from typing import TypedDict
except ImportError:
from typing_extensions import TypedDict
MAX_ATTR_LENGTH = 128
Attribute = TypedDict(
"Attribute",
{
"key": str,
"value": str,
},
)
IntermediateValue = TypedDict(
"IntermediateValue",
{
"step": int,
"value": float,
},
)
TrialParam = TypedDict(
"TrialParam",
{
"name": str,
"value": str,
},
)
def serialize_attrs(attrs: Dict[str, Any]) -> List[Attribute]:
serialized: List[Attribute] = []
for k, v in attrs.items():
value: str
if isinstance(v, str):
value = v[:MAX_ATTR_LENGTH] if len(v) > MAX_ATTR_LENGTH else v
elif isinstance(v, (bool, float, int)):
value = str(v)
elif isinstance(v, bytes):
value = "<binary object>"
elif v is None:
value = "None"
else: # unsupported type
continue
serialized.append({"key": k, "value": value})
return serialized
def serialize_intermediate_values(values: Dict[int, float]) -> List[IntermediateValue]:
return [{"step": step, "value": value} for step, value in values.items()]
def serialize_trial_params(params: Dict[str, Any]) -> List[TrialParam]:
return [{"name": name, "value": str(value)} for name, value in params.items()]
def serialize_study_summary(summary: StudySummary) -> Dict[str, Any]:
serialized = {
"study_id": summary._study_id,
"study_name": summary.study_name,
"directions": [d.name.lower() for d in summary.directions],
"user_attrs": serialize_attrs(summary.user_attrs),
"system_attrs": serialize_attrs(summary.system_attrs),
}
if summary.datetime_start is not None:
serialized["datetime_start"] = (summary.datetime_start.isoformat(),)
if summary.best_trial:
serialized["best_trial"] = serialize_frozen_trial(
summary._study_id, summary.best_trial
)
return serialized
def serialize_study_detail(
summary: StudySummary,
trials: List[FrozenTrial],
intersection: List[Tuple[str, BaseDistribution]],
union: List[Tuple[str, BaseDistribution]],
) -> Dict[str, Any]:
serialized: Dict[str, Any] = {
"name": summary.study_name,
"directions": [d.name.lower() for d in summary.directions],
}
if summary.datetime_start is not None:
serialized["datetime_start"] = summary.datetime_start.isoformat()
if summary.best_trial is not None:
serialized["best_trial"] = serialize_frozen_trial(
summary._study_id, summary.best_trial
)
serialized["trials"] = [
serialize_frozen_trial(summary._study_id, trial) for trial in trials
]
serialized["intersection_search_space"] = serialize_search_space(intersection)
serialized["union_search_space"] = serialize_search_space(union)
return serialized
def serialize_frozen_trial(study_id: int, trial: FrozenTrial) -> Dict[str, Any]:
serialized = {
"trial_id": trial._trial_id,
"study_id": study_id,
"number": trial.number,
"state": trial.state.name.capitalize(),
"intermediate_values": serialize_intermediate_values(trial.intermediate_values),
"params": serialize_trial_params(trial.params),
"user_attrs": serialize_attrs(trial.user_attrs),
"system_attrs": serialize_attrs(trial.system_attrs),
}
if trial.values is not None:
serialized["values"] = trial.values
if trial.datetime_start is not None:
serialized["datetime_start"] = trial.datetime_start.isoformat()
if trial.datetime_complete is not None:
serialized["datetime_complete"] = trial.datetime_complete.isoformat()
return serialized
def serialize_search_space(
search_space: List[Tuple[str, BaseDistribution]]
) -> List[Dict[str, Any]]:
serialized = []
for param_name, distribution in search_space:
serialized.append(
{
"name": param_name,
"type": distribution.__class__.__name__,
"attributes": distribution._asdict(),
}
)
return serialized