Add get_note() Python API

This commit is contained in:
c-bata
2023-03-29 17:47:20 +09:00
parent 957a8fd87d
commit 6892f31781
3 changed files with 42 additions and 2 deletions
+1
View File
@@ -1,6 +1,7 @@
from ._app import run_server # noqa
from ._app import wsgi # noqa
from ._named_objectives import set_objective_names # noqa
from ._note import get_note # noqa
from ._note import save_note # noqa
from ._objective_form_widget import ObjectiveChoiceWidget # noqa
from ._objective_form_widget import ObjectiveSliderWidget # noqa
+32 -2
View File
@@ -11,7 +11,6 @@ if TYPE_CHECKING:
from typing import Any
from typing import Optional
from typing import TypedDict
from typing import Union
NoteType = TypedDict(
"NoteType",
@@ -24,7 +23,7 @@ if TYPE_CHECKING:
SYSTEM_ATTR_MAX_LENGTH = 2045
def save_note(study_or_trial: Union[optuna.Study, optuna.Trial], body: str) -> None:
def save_note(study_or_trial: optuna.Study | optuna.Trial, body: str) -> None:
"""Save the note (Markdown format) to the Study or Trial.
Example:
@@ -70,6 +69,37 @@ def save_note(study_or_trial: Union[optuna.Study, optuna.Trial], body: str) -> N
save_note_with_version(storage, study_id, trial_id, next_ver, body)
def get_note(study_or_trial: optuna.Study | optuna.Trial) -> str:
"""Get the note (Markdown format) from the Study or Trial.
Example:
.. code-block:: python
import optuna
from optuna_dashboard import save_note, get_note
study = optuna.create_study()
save_note(study, "**Hello** World")
text = get_note(study)
print(text) # '**Hello** World'
"""
storage: BaseStorage
study_id: int
trial_id: Optional[int] = None
if isinstance(study_or_trial, optuna.Study):
storage = study_or_trial._storage
study_id = study_or_trial._study_id
else:
storage = study_or_trial.storage
study_id = study_or_trial.study._study_id
trial_id = study_or_trial._trial_id
system_attrs = storage.get_study_system_attrs(study_id)
note = get_note_from_system_attrs(system_attrs, trial_id)
return note["body"]
def note_ver_key(trial_id: Optional[int]) -> str:
prefix = "dashboard:note_ver"
if trial_id is None:
+9
View File
@@ -5,6 +5,7 @@ from unittest.mock import patch
import optuna
from optuna_dashboard import _note as note
from optuna_dashboard import get_note
from optuna_dashboard import save_note
@@ -29,6 +30,10 @@ class NoteTestCase(TestCase):
with self.subTest(body):
save_note(study, body)
system_attrs = study._storage.get_study_system_attrs(study._study_id)
actual = get_note(study)
assert actual == body
note_dict = note.get_note_from_system_attrs(system_attrs, None)
assert note_dict["body"] == body
assert note_dict["version"] == expected_ver
@@ -41,6 +46,10 @@ class NoteTestCase(TestCase):
with self.subTest(body):
save_note(trial, body)
system_attrs = study._storage.get_study_system_attrs(study._study_id)
actual = get_note(trial)
assert actual == body
note_dict = note.get_note_from_system_attrs(system_attrs, trial._trial_id)
assert note_dict["body"] == body
assert note_dict["version"] == expected_ver