diff --git a/optuna_dashboard/ts/components/StudyDetail.tsx b/optuna_dashboard/ts/components/StudyDetail.tsx index f832ec7f..7f9528d9 100644 --- a/optuna_dashboard/ts/components/StudyDetail.tsx +++ b/optuna_dashboard/ts/components/StudyDetail.tsx @@ -1,7 +1,9 @@ import ChevronRightIcon from "@mui/icons-material/ChevronRight" +import DownloadIcon from "@mui/icons-material/Download" import HomeIcon from "@mui/icons-material/Home" import { Box, + Button, Card, CardContent, IconButton, @@ -13,6 +15,8 @@ import React, { FC, useEffect, useMemo } from "react" import { Link, useParams } from "react-router-dom" import { useRecoilValue } from "recoil" +import { TrialTable } from "@optuna/react" +import * as Optuna from "@optuna/types" import { actionCreator } from "../action" import { useConstants } from "../constantsProvider" import { @@ -34,7 +38,6 @@ import { PreferentialHistory } from "./Preferential/PreferentialHistory" import { PreferentialTrials } from "./Preferential/PreferentialTrials" import { StudyHistory } from "./StudyHistory" import { TrialList } from "./TrialList" -import { TrialTable } from "./TrialTable" export const useURLVars = (): number => { const { studyId } = useParams<{ studyId: string }>() @@ -59,6 +62,19 @@ export const StudyDetail: FC<{ const reloadInterval = useRecoilValue(reloadIntervalState) const studyName = useStudyName(studyId) const isPreferential = useStudyIsPreferential(studyId) + const study: Optuna.Study | null = studyDetail + ? { + id: studyDetail.id, + name: studyDetail.name, + directions: studyDetail.directions, + union_search_space: studyDetail.union_search_space, + intersection_search_space: studyDetail.intersection_search_space, + union_user_attrs: studyDetail.union_user_attrs, + datetime_start: studyDetail.datetime_start, + trials: studyDetail.trials, + metric_names: studyDetail.metric_names, + } + : null const title = studyName !== null ? `${studyName} (id=${studyId})` : `Study #${studyId}` @@ -166,7 +182,10 @@ export const StudyDetail: FC<{ ) } else if (page === "trialList") { content = - } else if (page === "trialTable") { + } else if (page === "trialTable" && study !== null) { + const linkURL = (studyId: number, trialNumber: number) => { + return url_prefix + `/studies/${studyId}/trials?numbers=${trialNumber}` + } content = ( - + + diff --git a/optuna_dashboard/ts/components/TrialTable.tsx b/optuna_dashboard/ts/components/TrialTable.tsx deleted file mode 100644 index e7197bfa..00000000 --- a/optuna_dashboard/ts/components/TrialTable.tsx +++ /dev/null @@ -1,161 +0,0 @@ -import DownloadIcon from "@mui/icons-material/Download" -import LinkIcon from "@mui/icons-material/Link" -import { Button, IconButton, useTheme } from "@mui/material" -import React, { FC } from "react" - -import { DataGrid } from "@optuna/react" - -import { Link } from "react-router-dom" -import { StudyDetail, Trial } from "ts/types/optuna" - -import { - ColumnDef, - FilterFn, - Row, - createColumnHelper, -} from "@tanstack/react-table" -import { useConstants } from "../constantsProvider" - -const multiValueFilter: FilterFn = ( - row: Row, - columnId: string, - filterValue: string[] -) => { - const rowValue = row.getValue(columnId) as string - return !filterValue.includes(rowValue) -} - -export const TrialTable: FC<{ - studyDetail: StudyDetail | null -}> = ({ studyDetail }) => { - const { url_prefix } = useConstants() - - const theme = useTheme() - const trials: Trial[] = studyDetail !== null ? studyDetail.trials : [] - const metricNames: string[] = studyDetail?.metric_names || [] - - const columnHelper = createColumnHelper() - // eslint-disable-next-line @typescript-eslint/no-explicit-any - const columns: ColumnDef[] = [ - columnHelper.accessor("number", { - header: "Number", - enableColumnFilter: false, - }), - columnHelper.accessor("state", { - header: "State", - enableSorting: false, - enableColumnFilter: true, - filterFn: multiValueFilter, - }), - ] - if (studyDetail === null || studyDetail.directions.length === 1) { - columns.push( - columnHelper.accessor("values", { - header: "Value", - enableSorting: true, - enableColumnFilter: false, - sortUndefined: "last", - }) - ) - } else { - columns.push( - ...studyDetail.directions.map((s, objectiveId) => - columnHelper.accessor((row) => row["values"]?.[objectiveId], { - id: `values_${objectiveId}`, - header: - metricNames.length === studyDetail?.directions.length - ? metricNames[objectiveId] - : `Objective ${objectiveId}`, - enableSorting: true, - enableColumnFilter: false, - sortUndefined: "last", - }) - ) - ) - } - const isDynamicSpace = - studyDetail?.union_search_space.length !== - studyDetail?.intersection_search_space.length - studyDetail?.union_search_space.forEach((s) => { - const sortable = s.distribution.type !== "CategoricalDistribution" - const filterChoices: (string | null)[] | undefined = - s.distribution.type === "CategoricalDistribution" - ? s.distribution.choices.map((c) => c?.toString() ?? "null") - : undefined - const hasMissingValue = trials.some( - (t) => !t.params.some((p) => p.name === s.name) - ) - if (filterChoices !== undefined && isDynamicSpace && hasMissingValue) { - filterChoices.push(null) - } - columns.push( - columnHelper.accessor( - (row) => - row["params"].find((p) => p.name === s.name)?.param_external_value || - null, - { - id: `params_${s.name}`, - header: `Param ${s.name}`, - enableSorting: sortable, - sortUndefined: "last", - enableColumnFilter: filterChoices !== undefined, - filterFn: multiValueFilter, - } - ) - ) - }) - - studyDetail?.union_user_attrs.forEach((attr_spec) => { - columns.push( - columnHelper.accessor( - (row) => - row["user_attrs"].find((a) => a.key === attr_spec.key)?.value || null, - { - id: `user_attrs_${attr_spec.key}`, - header: `UserAttribute ${attr_spec.key}`, - enableSorting: attr_spec.sortable, - enableColumnFilter: false, - sortUndefined: "last", - } - ) - ) - }) - columns.push( - columnHelper.accessor((row) => row, { - header: "Detail", - cell: (info) => ( - - - - ), - enableSorting: false, - enableColumnFilter: false, - }) - ) - - return ( - <> - - - - ) -} diff --git a/tslib/react/src/components/TrialTable.tsx b/tslib/react/src/components/TrialTable.tsx index 3acb2bdb..fe148e06 100644 --- a/tslib/react/src/components/TrialTable.tsx +++ b/tslib/react/src/components/TrialTable.tsx @@ -1,3 +1,5 @@ +import { Link as LinkIcon } from "@mui/icons-material" +import { IconButton } from "@mui/material" import { FC } from "react" import * as Optuna from "@optuna/types" @@ -22,8 +24,12 @@ const multiValueFilter: FilterFn = ( export const TrialTable: FC<{ study: Optuna.Study initialRowsPerPage?: number -}> = ({ study, initialRowsPerPage }) => { + // biome-ignore lint/suspicious/noExplicitAny: Any react component. + linkComponent?: React.ComponentType + linkURL?: (studyId: number, trialNumber: number) => string +}> = ({ study, initialRowsPerPage, linkComponent, linkURL }) => { const trials: Optuna.Trial[] = study.trials + const metricNames: string[] = study.metric_names || [] const columnHelper = createColumnHelper() // biome-ignore lint/suspicious/noExplicitAny: It is difficult to specify this type. @@ -53,7 +59,10 @@ export const TrialTable: FC<{ ...study.directions.map((_s, objectiveId) => columnHelper.accessor((row) => row.values?.[objectiveId], { id: `values_${objectiveId}`, - header: `Objective ${objectiveId}`, + header: + metricNames.length === study.directions.length + ? metricNames[objectiveId] + : `Objective ${objectiveId}`, enableSorting: true, enableColumnFilter: false, sortUndefined: "last", @@ -61,9 +70,22 @@ export const TrialTable: FC<{ ) ) } + const isDynamicSpace = + study.union_search_space.length !== study.intersection_search_space.length - if (study?.union_search_space != null) { + if (study.union_search_space != null) { for (const s of study.union_search_space) { + const sortable = s.distribution.type !== "CategoricalDistribution" + const filterChoices: (string | null)[] | undefined = + s.distribution.type === "CategoricalDistribution" + ? s.distribution.choices.map((c) => c?.toString() ?? "null") + : undefined + const hasMissingValue = trials.some( + (t) => !t.params.some((p) => p.name === s.name) + ) + if (filterChoices !== undefined && isDynamicSpace && hasMissingValue) { + filterChoices.push(null) + } columns.push( columnHelper.accessor( (row) => { @@ -77,9 +99,9 @@ export const TrialTable: FC<{ { id: `params_${s.name}`, header: `Param ${s.name}`, - enableSorting: true, + enableSorting: sortable, sortUndefined: "last", - enableColumnFilter: false, + enableColumnFilter: filterChoices !== undefined, filterFn: multiValueFilter, } ) @@ -87,7 +109,7 @@ export const TrialTable: FC<{ } } - if (study?.union_search_space != null) { + if (study.union_search_space != null) { for (const attr_spec of study.union_user_attrs) { columns.push( columnHelper.accessor( @@ -108,6 +130,26 @@ export const TrialTable: FC<{ ) } } + if (linkComponent !== undefined && linkURL !== undefined) { + columns.push( + columnHelper.accessor((row) => row, { + header: "Detail", + cell: (info) => ( + + + + ), + enableSorting: false, + enableColumnFilter: false, + }) + ) + } return (