From d3d5e42fea4124994d2d212a98f1692cfb947556 Mon Sep 17 00:00:00 2001 From: c-bata Date: Sun, 1 Jan 2023 20:28:44 +0900 Subject: [PATCH] Add markdown note for trials --- optuna_dashboard/_app.py | 30 ++++++++- optuna_dashboard/_note.py | 20 +++++- optuna_dashboard/_serializer.py | 1 + optuna_dashboard/ts/action.ts | 57 ++++++++++++++-- optuna_dashboard/ts/apiClient.ts | 15 ++++- optuna_dashboard/ts/components/Note.tsx | 66 +++++++++++++++++-- .../ts/components/StudyDetail.tsx | 4 +- .../ts/components/StudyDetailBeta.tsx | 4 +- optuna_dashboard/ts/components/TrialTable.tsx | 13 ++++ optuna_dashboard/ts/types/index.d.ts | 1 + 10 files changed, 192 insertions(+), 19 deletions(-) diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index d10bc9c1..df871b35 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -355,7 +355,7 @@ def create_app(storage: BaseStorage, debug: bool = False) -> Bottle: @app.put("/api/studies//note") @json_api_view - def save_note(study_id: int) -> BottleViewReturn: + def save_study_note(study_id: int) -> BottleViewReturn: req_note_ver = request.json.get("version", None) req_note_body = request.json.get("body", None) if req_note_ver is None or req_note_body is None: @@ -371,7 +371,33 @@ def create_app(storage: BaseStorage, debug: bool = False) -> Bottle: "note": note.get_note_from_system_attrs(system_attrs), } - note.save_note(storage, study_id, req_note_ver, req_note_body) + note.save_note_in_study(storage, study_id, req_note_ver, req_note_body) + response.status = 204 # No content + return {} + + @app.put("/api/trials//note") + @json_api_view + def save_trial_note(trial_id: int) -> BottleViewReturn: + trial = storage.get_trial(trial_id) + if trial.state.is_finished(): + response.status = 400 # Bad request + return {"reason": "Cannot update the finished trials"} + + req_note_ver = request.json.get("version", None) + req_note_body = request.json.get("body", None) + if req_note_ver is None or req_note_body is None: + response.status = 400 # Bad request + return {"reason": "Invalid request."} + + if not note.version_is_incremented(trial.system_attrs, req_note_ver): + response.status = 409 # Conflict + return { + "reason": "The text you are editing has changed. " + "Please copy your edits and refresh the page.", + "note": note.get_note_from_system_attrs(system_attrs), + } + + note.save_note_in_trial(storage, trial_id, req_note_ver, req_note_body) response.status = 204 # No content return {} diff --git a/optuna_dashboard/_note.py b/optuna_dashboard/_note.py index 00710276..bffbf3aa 100644 --- a/optuna_dashboard/_note.py +++ b/optuna_dashboard/_note.py @@ -41,7 +41,7 @@ def version_is_incremented(system_attrs: dict[str, Any], req_note_ver: int) -> b return req_note_ver == db_note_ver + 1 -def save_note(storage: BaseStorage, study_id: int, ver: int, body: str) -> None: +def save_note_in_study(storage: BaseStorage, study_id: int, ver: int, body: str) -> None: storage.set_study_system_attr(study_id, NOTE_VER_KEY, ver) attrs = split_body(body) @@ -59,6 +59,24 @@ def save_note(storage: BaseStorage, study_id: int, ver: int, body: str) -> None: storage.set_study_system_attr(study_id, f"{NOTE_STR_KEY_PREFIX}{i}", "") +def save_note_in_trial(storage: BaseStorage, trial_id: int, ver: int, body: str) -> None: + storage.set_trial_system_attr(trial_id, NOTE_VER_KEY, ver) + + attrs = split_body(body) + for k, v in attrs.items(): + storage.set_study_system_attr(trial_id, k, v) + + # Clear previous messages + all_note_attrs: dict[str, str] = { + key: value + for key, value in storage.get_trial_system_attrs(trial_id).items() + if key.startswith(NOTE_STR_KEY_PREFIX) + } + if len(all_note_attrs) > len(attrs): + for i in range(len(attrs), len(all_note_attrs)): + storage.set_trial_system_attr(trial_id, f"{NOTE_STR_KEY_PREFIX}{i}", "") + + def split_body(note_str: str) -> dict[str, str]: note_len = len(note_str) attrs = {} diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index 4e426849..41ddcce5 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -108,6 +108,7 @@ def serialize_frozen_trial(study_id: int, trial: FrozenTrial) -> dict[str, Any]: "params": [{"name": name, "value": str(value)} for name, value in trial.params.items()], "user_attrs": serialize_attrs(trial.user_attrs), "system_attrs": serialize_attrs(getattr(trial, "_system_attrs", {})), + "note": note.get_note_from_system_attrs(getattr(trial, "_system_attrs", {})) } serialized_intermediate_values: list[IntermediateValue] = [] diff --git a/optuna_dashboard/ts/action.ts b/optuna_dashboard/ts/action.ts index 335a815c..c6d707ac 100644 --- a/optuna_dashboard/ts/action.ts +++ b/optuna_dashboard/ts/action.ts @@ -6,7 +6,8 @@ import { getParamImportances, createNewStudyAPI, deleteStudyAPI, - saveNoteAPI, + saveStudyNoteAPI, + saveTrialNoteAPI, } from "./apiClient" import { graphVisibilityState, @@ -157,8 +158,8 @@ export const actionCreator = () => { localStorage.setItem(localStorageGraphVisibility, JSON.stringify(value)) } - const saveNote = (studyId: number, note: Note): Promise => { - return saveNoteAPI(studyId, note) + const saveStudyNote = (studyId: number, note: Note): Promise => { + return saveStudyNoteAPI(studyId, note) .then(() => { const newStudy = Object.assign({}, studyDetails[studyId]) newStudy.note = note @@ -183,6 +184,53 @@ export const actionCreator = () => { }) } + const saveTrialNote = ( + studyId: number, + trialId: number, + note: Note + ): Promise => { + return saveTrialNoteAPI(trialId, note) + .then(() => { + const newStudy = Object.assign({}, studyDetails[studyId]) + const trial = newStudy.trials.find((t) => t.trial_id === trialId) + if (trial === undefined) { + enqueueSnackbar(`Unexpected error happens. Please reload the page.`, { + variant: "error", + }) + return + } + trial.note = note + setStudyDetailState(studyId, newStudy) + enqueueSnackbar(`Success to save the note`, { + variant: "success", + }) + }) + .catch((err) => { + if (err.response.status === 409) { + const newStudy = Object.assign({}, studyDetails[studyId]) + const trial = newStudy.trials.find((t) => t.trial_id === trialId) + if (trial === undefined) { + enqueueSnackbar( + `Unexpected error happens. Please reload the page.`, + { + variant: "error", + } + ) + return + } + trial.note = err.response.data.note + setStudyDetailState(studyId, newStudy) + } + const reason = err.response?.data.reason + if (reason !== undefined) { + enqueueSnackbar(`Failed: ${reason}`, { + variant: "error", + }) + } + throw err + }) + } + return { updateStudyDetail, updateStudySummaries, @@ -191,7 +239,8 @@ export const actionCreator = () => { deleteStudy, getGraphVisibility, saveGraphVisibility, - saveNote, + saveStudyNote, + saveTrialNote, } } diff --git a/optuna_dashboard/ts/apiClient.ts b/optuna_dashboard/ts/apiClient.ts index 405e18f6..e90797aa 100644 --- a/optuna_dashboard/ts/apiClient.ts +++ b/optuna_dashboard/ts/apiClient.ts @@ -14,6 +14,7 @@ interface TrialResponse { params: TrialParam[] user_attrs: Attribute[] system_attrs: Attribute[] + note: Note } const convertTrialResponse = (res: TrialResponse): Trial => { @@ -33,6 +34,7 @@ const convertTrialResponse = (res: TrialResponse): Trial => { params: res.params, user_attrs: res.user_attrs, system_attrs: res.system_attrs, + note: res.note, } } @@ -183,7 +185,7 @@ export const deleteStudyAPI = (studyId: number) => { }) } -export const saveNoteAPI = ( +export const saveStudyNoteAPI = ( studyId: number, note: { version: number; body: string } ): Promise => { @@ -194,6 +196,17 @@ export const saveNoteAPI = ( }) } +export const saveTrialNoteAPI = ( + trialId: number, + note: { version: number; body: string } +): Promise => { + return axiosInstance + .put(`/api/trials/${trialId}/note`, note) + .then((res) => { + return + }) +} + interface ParamImportancesResponse { param_importances: ParamImportance[][] } diff --git a/optuna_dashboard/ts/components/Note.tsx b/optuna_dashboard/ts/components/Note.tsx index 8939ba2d..827212be 100644 --- a/optuna_dashboard/ts/components/Note.tsx +++ b/optuna_dashboard/ts/components/Note.tsx @@ -51,12 +51,48 @@ const CodeBlock: CodeComponent | ReactMarkdownNames = ({ ) } -export const Note: FC<{ +export const TrialNote: FC<{ + studyId: number + trialId: number + latestNote: Note + editable: boolean +}> = ({ studyId, trialId, latestNote, editable }) => { + return ( + + ) +} + +export const StudyNote: FC<{ studyId: number latestNote: Note minRows: number cardSx?: SxProps }> = ({ studyId, latestNote, minRows, cardSx }) => { + return ( + + ) +} + +export const NoteBase: FC<{ + studyId: number + trialId?: number + latestNote: Note + minRows: number + editable: boolean + cardSx?: SxProps +}> = ({ studyId, trialId, latestNote, minRows, editable, cardSx }) => { const theme = useTheme() const [renderMarkdown, setRenderMarkdown] = useState(true) const [saving, setSaving] = useState(false) @@ -85,8 +121,14 @@ export const Note: FC<{ body: textAreaRef.current ? textAreaRef.current.value : "", } setSaving(true) - action - .saveNote(studyId, newNote) + + let actionResponse: Promise + if (trialId === undefined) { + actionResponse = action.saveStudyNote(studyId, newNote) + } else { + actionResponse = action.saveTrialNote(studyId, trialId, newNote) + } + actionResponse .then(() => { setCurNote(newNote) setRenderMarkdown(true) @@ -105,11 +147,16 @@ export const Note: FC<{ setCurNote(latestNote) window.onbeforeunload = null } + let defaultBody: string + if (editable) { + defaultBody = + "*A markdown editor for taking a memo, related to the study. Click the 'Edit' button in the upper right corner to access the editor.*" + } else { + defaultBody = "" + } let content if (renderMarkdown) { - const defaultBody = - "*A markdown editor for taking a memo, related to the study. Click the 'Edit' button in the upper right corner to access the editor.*" content = ( ) : ( - setRenderMarkdown(false)}> + setRenderMarkdown(false)} + > ) diff --git a/optuna_dashboard/ts/components/StudyDetail.tsx b/optuna_dashboard/ts/components/StudyDetail.tsx index 2058424a..ad06611a 100644 --- a/optuna_dashboard/ts/components/StudyDetail.tsx +++ b/optuna_dashboard/ts/components/StudyDetail.tsx @@ -24,7 +24,7 @@ import { GraphIntermediateValues } from "./GraphIntermediateValues" import { GraphSlice } from "./GraphSlice" import { GraphHistory } from "./GraphHistory" import { GraphParetoFront } from "./GraphParetoFront" -import { Note } from "./Note" +import { StudyNote } from "./Note" import { actionCreator } from "../action" import { graphVisibilityState, @@ -235,7 +235,7 @@ export const StudyDetail: FC<{ {studyDetail !== null ? ( - { + const editable = + trials[index].state === "Running" || trials[index].state === "Waiting" return ( + {trials[index].note.body !== "" || editable ? ( + + + + ) : null} diff --git a/optuna_dashboard/ts/types/index.d.ts b/optuna_dashboard/ts/types/index.d.ts index 05d5fbd2..3ecdc25d 100644 --- a/optuna_dashboard/ts/types/index.d.ts +++ b/optuna_dashboard/ts/types/index.d.ts @@ -82,6 +82,7 @@ declare interface Trial { params: TrialParam[] user_attrs: Attribute[] system_attrs: Attribute[] + note: Note } declare interface StudySummary {