diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index 64ff7cf3..ef785cbf 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -83,6 +83,7 @@ if TYPE_CHECKING: MAX_ATTR_LENGTH = 1024 +CONSTRAINTS_KEY = "constraints" def serialize_attrs(attrs: dict[str, Any]) -> list[Attribute]: @@ -187,6 +188,7 @@ def serialize_frozen_trial( ), "note": note.get_note_from_system_attrs(study_system_attrs, trial._trial_id), "artifacts": list_trial_artifacts(study_system_attrs, trial._trial_id), + "constraints": trial_system_attrs.get(CONSTRAINTS_KEY, []), } serialized_intermediate_values: list[IntermediateValue] = [] diff --git a/optuna_dashboard/ts/apiClient.ts b/optuna_dashboard/ts/apiClient.ts index b6793cb2..fc75b942 100644 --- a/optuna_dashboard/ts/apiClient.ts +++ b/optuna_dashboard/ts/apiClient.ts @@ -30,6 +30,7 @@ interface TrialResponse { system_attrs: Attribute[] note: Note artifacts: Artifact[] + constraints: number[] } const convertTrialResponse = (res: TrialResponse): Trial => { @@ -52,6 +53,7 @@ const convertTrialResponse = (res: TrialResponse): Trial => { system_attrs: res.system_attrs, note: res.note, artifacts: res.artifacts, + constraints: res.constraints, } } diff --git a/optuna_dashboard/ts/components/GraphParetoFront.tsx b/optuna_dashboard/ts/components/GraphParetoFront.tsx index 18f5f896..f7afdf06 100644 --- a/optuna_dashboard/ts/components/GraphParetoFront.tsx +++ b/optuna_dashboard/ts/components/GraphParetoFront.tsx @@ -108,9 +108,10 @@ const makeScatterObject = ( objectiveXId: number, objectiveYId: number, hovertemplate: string, - dominated: boolean + dominated: boolean, + feasible: boolean ): Partial => { - const marker = makeMarker(trials, dominated) + const marker = makeMarker(trials, dominated, feasible) return { x: trials.map((t) => t.values![objectiveXId] as number), y: trials.map((t) => t.values![objectiveYId] as number), @@ -124,9 +125,10 @@ const makeScatterObject = ( const makeMarker = ( trials: Trial[], - dominated: boolean + dominated: boolean, + feasible: boolean ): Partial => { - if (dominated) { + if (feasible && dominated) { return { line: { width: 0.5, color: "Grey" }, // @ts-ignore @@ -137,7 +139,7 @@ const makeMarker = ( title: "Trial", }, } - } else { + } else if (feasible && !dominated) { return { line: { width: 0.5, color: "Grey" }, // @ts-ignore @@ -149,6 +151,11 @@ const makeMarker = ( xpad: 80, }, } + } else { + return { + // @ts-ignore + color: "#cccccc", + } } } @@ -183,8 +190,18 @@ const plotParetoFront = ( return } - const normalizedValues: number[][] = [] + const feasibleTrials: Trial[] = [] + const InfeasibleTrials: Trial[] = [] filteredTrials.forEach((t) => { + if (t.constraints.every((c) => c <= 0)) { + feasibleTrials.push(t) + } else { + InfeasibleTrials.push(t) + } + }) + + const normalizedValues: number[][] = [] + feasibleTrials.forEach((t) => { if (t.values && t.values.length === study.directions.length) { const trialValues = t.values.map((v, i) => { return study.directions[i] === "minimize" @@ -210,17 +227,29 @@ const plotParetoFront = ( const plotData: Partial[] = [ makeScatterObject( - filteredTrials.filter((t, i) => dominatedTrials[i]), + feasibleTrials.filter((t, i) => dominatedTrials[i]), objectiveXId, objectiveYId, - "%{text}Trial", + InfeasibleTrials.length === 0 + ? "%{text}Trial" + : "%{text}Feasible Trial", + true, true ), makeScatterObject( - filteredTrials.filter((t, i) => !dominatedTrials[i]), + feasibleTrials.filter((t, i) => !dominatedTrials[i]), objectiveXId, objectiveYId, "%{text}Best Trial", + false, + true + ), + makeScatterObject( + InfeasibleTrials, + objectiveXId, + objectiveYId, + "%{text}Infeasible Trial", + false, false ), ] diff --git a/optuna_dashboard/ts/types/index.d.ts b/optuna_dashboard/ts/types/index.d.ts index 72ba1dea..040a1fc3 100644 --- a/optuna_dashboard/ts/types/index.d.ts +++ b/optuna_dashboard/ts/types/index.d.ts @@ -112,6 +112,7 @@ type Trial = { }[] user_attrs: Attribute[] system_attrs: Attribute[] + constraints: number[] note: Note artifacts: Artifact[] }