Merge pull request #237 from c-bata/fix-neg-inf

Fix bug when given -inf or nan values
This commit is contained in:
Masashi Shibata
2022-05-25 21:52:28 +09:00
committed by GitHub
11 changed files with 106 additions and 47 deletions
+32 -23
View File
@@ -1,11 +1,11 @@
import json
import math
from typing import Any
from typing import Dict
from typing import List
from typing import Tuple
from typing import Union
import numpy as np
from optuna.distributions import BaseDistribution
from optuna.study import StudySummary
from optuna.trial import FrozenTrial
@@ -14,8 +14,10 @@ from . import _note as note
try:
from typing import Literal
from typing import TypedDict
except ImportError:
from typing_extensions import Literal # type: ignore
from typing_extensions import TypedDict
@@ -31,14 +33,7 @@ IntermediateValue = TypedDict(
"IntermediateValue",
{
"step": int,
"value": Union[float, str],
},
)
TrialParam = TypedDict(
"TrialParam",
{
"name": str,
"value": str,
"value": Union[float, Literal["inf", "-inf", "nan"]],
},
)
@@ -56,17 +51,6 @@ def serialize_attrs(attrs: Dict[str, Any]) -> List[Attribute]:
return serialized
def serialize_intermediate_values(values: Dict[int, float]) -> List[IntermediateValue]:
return [
{"step": step, "value": "inf" if math.isinf(value) else value}
for step, value in values.items()
]
def serialize_trial_params(params: Dict[str, Any]) -> List[TrialParam]:
return [{"name": name, "value": str(value)} for name, value in params.items()]
def serialize_study_summary(summary: StudySummary) -> Dict[str, Any]:
serialized = {
"study_id": summary._study_id,
@@ -118,14 +102,39 @@ def serialize_frozen_trial(study_id: int, trial: FrozenTrial) -> Dict[str, Any]:
"study_id": study_id,
"number": trial.number,
"state": trial.state.name.capitalize(),
"intermediate_values": serialize_intermediate_values(trial.intermediate_values),
"params": serialize_trial_params(trial.params),
"params": [
{"name": name, "value": str(value)} for name, value in trial.params.items()
],
"user_attrs": serialize_attrs(trial.user_attrs),
"system_attrs": serialize_attrs(trial.system_attrs),
}
serialized_intermediate_values: List[IntermediateValue] = []
for step, value in trial.intermediate_values.items():
serialized_value: Union[float, Literal["nan", "inf", "-inf"]]
if np.isnan(value):
serialized_value = "nan"
elif np.isposinf(value):
serialized_value = "inf"
elif np.isneginf(value):
serialized_value = "-inf"
else:
assert np.isfinite(value)
serialized_value = value
serialized_intermediate_values.append({"step": step, "value": serialized_value})
serialized["intermediate_values"] = serialized_intermediate_values
if trial.values is not None:
serialized["values"] = ["inf" if math.isinf(v) else v for v in trial.values]
serialized_values: List[Union[float, Literal["inf", "-inf"]]] = []
for v in trial.values:
assert not np.isnan(v), "Should not detect nan value"
if np.isposinf(v):
serialized_values.append("inf")
elif np.isneginf(v):
serialized_values.append("-inf")
else:
serialized_values.append(v)
serialized["values"] = serialized_values
if trial.datetime_start is not None:
serialized["datetime_start"] = trial.datetime_start.isoformat()
+1 -1
View File
@@ -7,7 +7,7 @@ interface TrialResponse {
study_id: number
number: number
state: TrialState
values?: (number | "inf")[]
values?: TrialValueNumber[]
intermediate_values: TrialIntermediateValue[]
datetime_start?: string
datetime_complete?: string
+2 -1
View File
@@ -66,7 +66,8 @@ const filterFunc = (trial: Trial, objectiveId: number): boolean => {
return (
trial.state === "Complete" &&
trial.values !== undefined &&
trial.values[objectiveId] !== "inf"
trial.values[objectiveId] !== "inf" &&
trial.values[objectiveId] !== "-inf"
)
}
@@ -180,7 +180,9 @@ const filterFunc = (trial: Trial, objectiveId: number): boolean => {
return false
}
return (
trial.values.length > objectiveId && trial.values[objectiveId] !== "inf"
trial.values.length > objectiveId &&
trial.values[objectiveId] !== "inf" &&
trial.values[objectiveId] !== "-inf"
)
}
@@ -58,7 +58,9 @@ const plotIntermediateValue = (trials: Trial[], mode: string) => {
t.state == "Running"
)
const plotData: Partial<plotly.PlotData>[] = filteredTrials.map((trial) => {
const values = trial.intermediate_values.filter((iv) => iv.value !== "inf")
const values = trial.intermediate_values.filter(
(iv) => iv.value !== "inf" && iv.value !== "-inf" && iv.value !== "nan"
)
return {
x: values.map((iv) => iv.step),
y: values.map((iv) => iv.value),
@@ -71,7 +71,9 @@ const filterFunc = (trial: Trial, objectiveId: number): boolean => {
return false
}
return (
trial.values.length > objectiveId && trial.values[objectiveId] !== "inf"
trial.values.length > objectiveId &&
trial.values[objectiveId] !== "inf" &&
trial.values[objectiveId] !== "-inf"
)
}
@@ -90,7 +90,7 @@ const filterFunc = (trial: Trial, directions: StudyDirection[]): boolean => {
trial.state === "Complete" &&
trial.values !== undefined &&
trial.values.length === directions.length &&
trial.values.every((v) => v !== "inf")
trial.values.every((v) => v !== "inf" && v !== "-inf")
)
}
@@ -130,7 +130,9 @@ const filterFunc = (
return false
}
return (
trial.values.length > objectiveId && trial.values[objectiveId] !== "inf"
trial.values.length > objectiveId &&
trial.values[objectiveId] !== "inf" &&
trial.values[objectiveId] !== "-inf"
)
}
+41 -9
View File
@@ -515,13 +515,18 @@ export const TrialTable: FC<{ studyDetail: StudyDetail | null }> = ({
if (firstVal === secondVal) {
return 0
} else if (firstVal && secondVal) {
return firstVal < secondVal ? 1 : -1
} else if (firstVal) {
}
if (firstVal === undefined) {
return -1
} else {
} else if (secondVal === undefined) {
return 1
}
if (firstVal === "-inf" || secondVal === "inf") {
return 1
} else if (secondVal === "-inf" || firstVal === "inf") {
return -1
}
return firstVal < secondVal ? 1 : -1
},
toCellValue: (i) => {
if (trials[i].values === undefined) {
@@ -542,13 +547,18 @@ export const TrialTable: FC<{ studyDetail: StudyDetail | null }> = ({
if (firstVal === secondVal) {
return 0
} else if (firstVal && secondVal) {
return firstVal < secondVal ? 1 : -1
} else if (firstVal) {
}
if (firstVal === undefined) {
return -1
} else {
} else if (secondVal === undefined) {
return 1
}
if (firstVal === "-inf" || secondVal === "inf") {
return 1
} else if (secondVal === "-inf" || firstVal === "inf") {
return -1
}
return firstVal < secondVal ? 1 : -1
},
toCellValue: (i) => {
if (trials[i].values === undefined) {
@@ -647,7 +657,29 @@ export const TrialTable: FC<{ studyDetail: StudyDetail | null }> = ({
const collapseIntermediateValueColumns: DataGridColumn<TrialIntermediateValue>[] =
[
{ field: "step", label: "Step", sortable: true },
{ field: "value", label: "Value", sortable: true },
{
field: "value",
label: "Value",
sortable: true,
less: (firstEl, secondEl): number => {
const firstVal = firstEl.value
const secondVal = secondEl.value
if (firstVal === secondVal) {
return 0
}
if (firstVal === "nan") {
return -1
} else if (secondVal === "nan") {
return 1
}
if (firstVal === "-inf" || secondVal === "inf") {
return 1
} else if (secondVal === "-inf" || firstVal === "inf") {
return -1
}
return firstVal < secondVal ? 1 : -1
},
},
]
const collapseAttrColumns: DataGridColumn<Attribute>[] = [
{ field: "key", label: "Key", sortable: true },
+4 -2
View File
@@ -7,6 +7,8 @@ declare const APP_BAR_TITLE: string
declare const API_ENDPOINT: string
declare const URL_PREFIX: string
type TrialValueNumber = number | "inf" | "-inf"
type TrialIntermediateValueNumber = number | "inf" | "-inf" | "nan"
type TrialState = "Running" | "Complete" | "Pruned" | "Fail" | "Waiting"
type StudyDirection = "maximize" | "minimize" | "not_set"
type Distribution =
@@ -21,7 +23,7 @@ type Distribution =
declare interface TrialIntermediateValue {
step: number
value: number | "inf"
value: TrialIntermediateValueNumber
}
declare interface TrialParam {
@@ -55,7 +57,7 @@ declare interface Trial {
study_id: number
number: number
state: TrialState
values?: (number | "inf")[]
values?: TrialValueNumber[]
intermediate_values: TrialIntermediateValue[]
datetime_start?: Date
datetime_complete?: Date
+13 -6
View File
@@ -1,6 +1,5 @@
import argparse
import asyncio
import math
import os
import sys
import threading
@@ -87,13 +86,15 @@ def create_dummy_storage() -> optuna.storages.InMemoryStorage:
study.optimize(objective_single_dynamic, n_trials=50)
# Single objective study with 'inf' value
# Single objective study with 'inf', '-inf', or 'nan' value
study = optuna.create_study(study_name="single-inf", storage=storage)
def objective_single_inf(trial: optuna.Trial) -> float:
x = trial.suggest_float("x", -10, 10)
if x > 0:
return math.inf
if trial.number % 3 == 0:
return float("inf")
elif trial.number % 3 == 1:
return float("-inf")
else:
return x**2
@@ -152,13 +153,19 @@ def create_dummy_storage() -> optuna.storages.InMemoryStorage:
study.optimize(objective_prune_without_report, n_trials=100)
# Single objective pruned after reported 'inf' value
# Single objective pruned after reported 'inf', '-inf', or 'nan'
study = optuna.create_study(study_name="single-inf-report", storage=storage)
def objective_single_inf_report(trial: optuna.Trial) -> float:
x = trial.suggest_float("x", -10, 10)
if trial.number % 3 == 0:
trial.report(float("inf"), 1)
elif trial.number % 3 == 1:
trial.report(float("-inf"), 1)
else:
trial.report(float("nan"), 1)
if x > 0:
trial.report(math.inf, 1)
raise optuna.TrialPruned()
else:
return x**2