diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index eaab1ba6..812201ad 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -222,10 +222,8 @@ def create_app( ) = get_cached_extra_study_property(study_id, trials) plotly_graph_objects = get_plotly_graph_objects(system_attrs) - trials_id2number = {trial._trial_id: trial.number for trial in trials} - skipped_trials = [ - trials_id2number[trial_id] for trial_id in get_skipped_trial_ids(system_attrs) - ] + skipped_trial_ids = get_skipped_trial_ids(system_attrs) + skipped_trial_numbers = [t.number for t in trials if t._trial_id in skipped_trial_ids] return serialize_study_detail( summary, best_trials, @@ -235,7 +233,7 @@ def create_app( union_user_attrs, has_intermediate_values, plotly_graph_objects, - skipped_trials, + skipped_trial_numbers, ) @app.get("/api/studies//param_importances") diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index e1840df9..81544ea8 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -136,7 +136,7 @@ def serialize_study_detail( union_user_attrs: list[tuple[str, bool]], has_intermediate_values: bool, plotly_graph_objects: dict[str, str], - skipped_trials: list[int], + skipped_trial_numbers: list[int], ) -> dict[str, Any]: serialized: dict[str, Any] = { "name": summary.study_name, @@ -174,7 +174,7 @@ def serialize_study_detail( if serialized["is_preferential"]: serialized["preference_history"] = serialize_preference_history(system_attrs) serialized["preferences"] = get_preferences(system_attrs) - serialized["skipped_trials"] = skipped_trials + serialized["skipped_trial_numbers"] = skipped_trial_numbers serialized["plotly_graph_objects"] = [ {"id": id_, "graph_object": graph_object} for id_, graph_object in plotly_graph_objects.items() diff --git a/optuna_dashboard/ts/action.ts b/optuna_dashboard/ts/action.ts index 0e1c9420..41afdadb 100644 --- a/optuna_dashboard/ts/action.ts +++ b/optuna_dashboard/ts/action.ts @@ -16,6 +16,8 @@ import { deleteArtifactAPI, reportPreferenceAPI, skipPreferentialTrialAPI, + removePreferentialHistoryAPI, + restorePreferentialHistoryAPI, reportFeedbackComponentAPI, } from "./apiClient" import { @@ -587,11 +589,11 @@ export const actionCreator = () => { } const updatePreference = ( - study_id: number, + studyId: number, candidates: number[], clicked: number ) => { - reportPreferenceAPI(study_id, candidates, clicked).catch((err) => { + reportPreferenceAPI(studyId, candidates, clicked).catch((err) => { const reason = err.response?.data.reason enqueueSnackbar(`Failed to report preference. Reason: ${reason}`, { variant: "error", @@ -609,7 +611,6 @@ export const actionCreator = () => { console.log(err) }) } - const updateFeedbackComponent = ( studyId: number, compoennt_type: FeedbackComponentType @@ -632,6 +633,52 @@ export const actionCreator = () => { }) } + const removePreferentialHistory = (studyId: number, historyId: string) => { + removePreferentialHistoryAPI(studyId, historyId) + .then(() => { + const newStudy = Object.assign({}, studyDetails[studyId]) + newStudy.preference_history = newStudy.preference_history?.map((h) => + h.id === historyId ? { ...h, is_removed: true } : h + ) + const removed = newStudy.preference_history + ?.filter((h) => h.id === historyId) + .pop()?.preferences + newStudy.preferences = newStudy.preferences?.filter( + (p) => !removed?.some((r) => r[0] === p[0] && r[1] === p[1]) + ) + setStudyDetailState(studyId, newStudy) + }) + .catch((err) => { + const reason = err.response?.data.reason + + enqueueSnackbar(`Failed to switch history. Reason: ${reason}`, { + variant: "error", + }) + console.log(err) + }) + } + const restorePreferentialHistory = (studyId: number, historyId: string) => { + restorePreferentialHistoryAPI(studyId, historyId) + .then(() => { + const newStudy = Object.assign({}, studyDetails[studyId]) + newStudy.preference_history = newStudy.preference_history?.map((h) => + h.id === historyId ? { ...h, is_removed: false } : h + ) + const restored = newStudy.preference_history + ?.filter((h) => h.id === historyId) + .pop()?.preferences + newStudy.preferences = newStudy.preferences?.concat(restored ?? []) + setStudyDetailState(studyId, newStudy) + }) + .catch((err) => { + const reason = err.response?.data.reason + enqueueSnackbar(`Failed to switch history. Reason: ${reason}`, { + variant: "error", + }) + console.log(err) + }) + } + return { updateAPIMeta, updateStudyDetail, @@ -653,6 +700,8 @@ export const actionCreator = () => { saveTrialUserAttrs, updatePreference, skipPreferentialTrial, + removePreferentialHistory, + restorePreferentialHistory, updateFeedbackComponent, } } diff --git a/optuna_dashboard/ts/apiClient.ts b/optuna_dashboard/ts/apiClient.ts index 60c461db..c0065bd3 100644 --- a/optuna_dashboard/ts/apiClient.ts +++ b/optuna_dashboard/ts/apiClient.ts @@ -55,10 +55,9 @@ const convertTrialResponse = (res: TrialResponse): Trial => { } } -interface PreferenceHistoryResponce { +interface PreferenceHistoryResponse { history: { id: string - preference_id: string candidates: number[] clicked: number mode: PreferenceFeedbackMode @@ -69,11 +68,10 @@ interface PreferenceHistoryResponce { } const convertPreferenceHistory = ( - res: PreferenceHistoryResponce + res: PreferenceHistoryResponse ): PreferenceHistory => { return { id: res.history.id, - preference_id: res.history.preference_id, candidates: res.history.candidates, clicked: res.history.clicked, feedback_mode: res.history.mode, @@ -99,10 +97,10 @@ interface StudyDetailResponse { objective_names?: string[] form_widgets?: FormWidgets preferences?: [number, number][] - preference_history?: PreferenceHistoryResponce[] + preference_history?: PreferenceHistoryResponse[] plotly_graph_objects: PlotlyGraphObject[] feedback_component_type: FeedbackComponentType - skipped_trials?: number[] + skipped_trial_numbers?: number[] } export const getStudyDetailAPI = ( @@ -144,7 +142,7 @@ export const getStudyDetailAPI = ( convertPreferenceHistory ), plotly_graph_objects: res.data.plotly_graph_objects, - skipped_trials: res.data.skipped_trials ?? [], + skipped_trial_numbers: res.data.skipped_trial_numbers ?? [], } }) } @@ -379,6 +377,27 @@ export const skipPreferentialTrialAPI = ( }) } +export const removePreferentialHistoryAPI = ( + studyId: number, + historyUuid: string +): Promise => { + return axiosInstance + .delete(`/api/studies/${studyId}/preference/${historyUuid}`) + .then(() => { + return + }) +} +export const restorePreferentialHistoryAPI = ( + studyId: number, + historyUuid: string +): Promise => { + return axiosInstance + .post(`/api/studies/${studyId}/preference/${historyUuid}`) + .then(() => { + return + }) +} + export const reportFeedbackComponentAPI = ( studyId: number, component_type: FeedbackComponentType diff --git a/optuna_dashboard/ts/components/PreferenceHistory.tsx b/optuna_dashboard/ts/components/PreferenceHistory.tsx index e9c4948c..6aa67317 100644 --- a/optuna_dashboard/ts/components/PreferenceHistory.tsx +++ b/optuna_dashboard/ts/components/PreferenceHistory.tsx @@ -10,12 +10,15 @@ import { import ClearIcon from "@mui/icons-material/Clear" import IconButton from "@mui/material/IconButton" import OpenInFullIcon from "@mui/icons-material/OpenInFull" +import RestoreFromTrashIcon from "@mui/icons-material/RestoreFromTrash" +import DeleteIcon from "@mui/icons-material/Delete" import Modal from "@mui/material/Modal" import { red } from "@mui/material/colors" import { TrialListDetail } from "./TrialList" import { getArtifactUrlPath } from "./PreferentialTrials" import { formatDate } from "../dateUtil" +import { actionCreator } from "../action" import { useStudyDetailValue } from "../state" import { PreferentialOutputComponent } from "./PreferentialOutputComponent" @@ -157,33 +160,75 @@ const CandidateTrial: FC<{ ) } -const ChoiceTrials: FC<{ choice: PreferenceHistory; trials: Trial[] }> = ({ - choice, - trials, -}) => { +const ChoiceTrials: FC<{ + choice: PreferenceHistory + trials: Trial[] + studyId: number +}> = ({ choice, trials, studyId }) => { + const [isRemoved, setRemoved] = useState(choice.is_removed) const theme = useTheme() const worst_trials = new Set([choice.clicked]) + const action = actionCreator() return ( - - {formatDate(choice.timestamp)} - + + {formatDate(choice.timestamp)} + + {choice.is_removed ? ( + { + setRemoved(false) + action.restorePreferentialHistory(studyId, choice.id) + }} + sx={{ + margin: `auto ${theme.spacing(2)}`, + }} + > + + + ) : ( + { + setRemoved(true) + action.removePreferentialHistory(studyId, choice.id) + }} + sx={{ + margin: `auto ${theme.spacing(2)}`, + }} + > + + + )} + + {choice.candidates.map((trial_num, index) => ( = ({ key={choice.id} choice={choice} trials={studyDetail.trials} + studyId={studyDetail.id} /> ))} diff --git a/optuna_dashboard/ts/components/PreferentialGraph.tsx b/optuna_dashboard/ts/components/PreferentialGraph.tsx index 4f2c1059..e30974c1 100644 --- a/optuna_dashboard/ts/components/PreferentialGraph.tsx +++ b/optuna_dashboard/ts/components/PreferentialGraph.tsx @@ -203,6 +203,7 @@ export const PreferentialGraph: FC<{ if (!studyDetail.is_preferential || studyDetail.preferences === undefined) return const preferences = reductionPreference(studyDetail.preferences) + const trialNodes = Array.from(new Set(preferences.flat())) const graph: ElkNode = { id: "root", layoutOptions: { @@ -211,8 +212,8 @@ export const PreferentialGraph: FC<{ "elk.layered.spacing.nodeNodeBetweenLayers": nodeMargin.toString(), "elk.spacing.nodeNode": nodeMargin.toString(), }, - children: studyDetail.trials.map((trial) => ({ - id: `${trial.number}`, + children: trialNodes.map((trial) => ({ + id: `${trial}`, targetPosition: "top", sourcePosition: "bottom", width: nodeWidth, @@ -229,7 +230,7 @@ export const PreferentialGraph: FC<{ .then((layoutedGraph) => { setNodes( layoutedGraph.children?.map((node, index) => { - const trial = studyDetail.trials[index] + const trial = studyDetail.trials[trialNodes[index]] return { id: `${trial.number}`, type: "note", diff --git a/optuna_dashboard/ts/components/PreferentialTrials.tsx b/optuna_dashboard/ts/components/PreferentialTrials.tsx index 27f3a8d5..b3934a80 100644 --- a/optuna_dashboard/ts/components/PreferentialTrials.tsx +++ b/optuna_dashboard/ts/components/PreferentialTrials.tsx @@ -21,10 +21,11 @@ import { import IconButton from "@mui/material/IconButton" import OpenInFullIcon from "@mui/icons-material/OpenInFull" import ReplayIcon from "@mui/icons-material/Replay" +import { red } from "@mui/material/colors" +import UndoIcon from "@mui/icons-material/Undo" import ClearIcon from "@mui/icons-material/Clear" import SettingsIcon from "@mui/icons-material/Settings" import FullscreenIcon from "@mui/icons-material/Fullscreen" -import red from "@mui/material/colors/red" import { actionCreator } from "../action" import { TrialListDetail } from "./TrialList" @@ -399,6 +400,8 @@ export const PreferentialTrials: FC<{ studyDetail: StudyDetail | null }> = ({ studyDetail, }) => { const theme = useTheme() + const action = actionCreator() + const [undoHistoryFlag, setUndoHistoryFlag] = useState(false) const [openThreejsArtifactModal, renderThreejsArtifactModal] = useThreejsArtifactModal() const [displayTrials, setDisplayTrials] = useState({ @@ -414,8 +417,9 @@ export const PreferentialTrials: FC<{ studyDetail: StudyDetail | null }> = ({ const hiddenTrials = new Set( studyDetail.preference_history - ?.map((p) => p.clicked) - .concat(studyDetail.skipped_trials) ?? [] + ?.filter((h) => !h.is_removed) + .map((p) => p.clicked) + .concat(studyDetail.skipped_trial_numbers) ?? [] ) const activeTrials = studyDetail.trials.filter( (t) => @@ -468,34 +472,75 @@ export const PreferentialTrials: FC<{ studyDetail: StudyDetail | null }> = ({ } }) } + const visibleTrial = (num: number) => { + setDisplayTrials((prev) => { + const index = prev.clicked.findIndex((n) => n === num) + if (index === -1) { + return prev + } + const clicked = [...prev.clicked] + clicked[index] = -1 + return { + display: prev.display, + clicked: clicked, + } + }) + } + const latestHistoryId = + studyDetail?.preference_history?.filter((h) => !h.is_removed).pop()?.id ?? + null return ( - - setSettingShown(true)} - > - - - - Which trial is the worst? - + + + + Which trial is the worst? + + + + + + {displayTrials.display.map((t, index) => { const trial = activeTrials.find((trial) => trial.number === t) @@ -553,9 +598,10 @@ export const PreferentialTrials: FC<{ studyDetail: StudyDetail | null }> = ({ setDetailTrial(null)} > diff --git a/optuna_dashboard/ts/types/index.d.ts b/optuna_dashboard/ts/types/index.d.ts index e23c22e3..3e883e14 100644 --- a/optuna_dashboard/ts/types/index.d.ts +++ b/optuna_dashboard/ts/types/index.d.ts @@ -218,7 +218,7 @@ type StudyDetail = { preferences?: [number, number][] preference_history?: PreferenceHistory[] plotly_graph_objects: PlotlyGraphObject[] - skipped_trials: number[] + skipped_trial_numbers: number[] } type StudyDetails = { @@ -230,7 +230,6 @@ type StudyParamImportance = { } type PreferenceHistory = { id: string - preference_id: string candidates: number[] clicked: number feedback_mode: PreferenceFeedbackMode