Modify the logic of storages

This commit is contained in:
porink0424
2024-06-21 18:36:57 +09:00
parent cc8fe6a6d4
commit e215b1bf10
2 changed files with 98 additions and 38 deletions
+55 -17
View File
@@ -1,6 +1,36 @@
import * as Optuna from "@optuna/types"
import { OptunaStorage } from "./storage"
// TODO(porink0424): Refactor to common function with sqlite.ts (current workaround duplicates code due to missing file extensions in tsc build output).
const isDistributionEqual = (
a: Optuna.Distribution,
b: Optuna.Distribution
) => {
if (a.type !== b.type) {
return false
}
if (a.type === "IntDistribution" || a.type === "FloatDistribution") {
if (b.type !== "IntDistribution" && b.type !== "FloatDistribution") {
throw new Error("Invalid distribution type")
}
return (
a.low === b.low &&
a.high === b.high &&
a.step === b.step &&
a.log === b.log
)
}
if (a.type === "CategoricalDistribution") {
if (b.type !== "CategoricalDistribution") {
throw new Error("Invalid distribution type")
}
return JSON.stringify(a.choices) === JSON.stringify(b.choices)
}
throw new Error("Invalid distribution type")
}
// JournalStorage
enum JournalOperation {
CREATE_STUDY = 0,
@@ -137,30 +167,42 @@ class JournalStorage {
public getStudies(): Optuna.Study[] {
for (const study of this.studies) {
const unionUserAttrs: Set<string> = new Set()
const unionSearchSpace: Set<string> = new Set()
let intersectionSearchSpace: string[] = []
const nameToSearchSpaceItem: Map<string, Optuna.SearchSpaceItem> =
new Map()
const unionSearchSpace: Optuna.SearchSpaceItem[] = []
let intersectionSearchSpace: Optuna.SearchSpaceItem[] = []
study.trials.forEach((trial, index) => {
for (const userAttr of trial.user_attrs) {
unionUserAttrs.add(userAttr.key)
}
for (const param of trial.params) {
unionSearchSpace.add(param.name)
if (!nameToSearchSpaceItem.has(param.name)) {
nameToSearchSpaceItem.set(param.name, {
if (
!unionSearchSpace.some(
(item) =>
item.name === param.name &&
isDistributionEqual(item.distribution, param.distribution)
)
) {
unionSearchSpace.push({
name: param.name,
distribution: param.distribution,
})
}
}
if (index === 0) {
intersectionSearchSpace = Array.from(unionSearchSpace)
intersectionSearchSpace = [...unionSearchSpace]
} else {
intersectionSearchSpace = intersectionSearchSpace.filter((name) => {
return trial.params.some((param) => param.name === name)
})
intersectionSearchSpace = intersectionSearchSpace.filter(
(searchSpaceItem) => {
return trial.params.some(
(param) =>
param.name === searchSpaceItem.name &&
isDistributionEqual(
param.distribution,
searchSpaceItem.distribution
)
)
}
)
}
})
study.union_user_attrs = Array.from(unionUserAttrs).map((key) => {
@@ -169,12 +211,8 @@ class JournalStorage {
sortable: false,
}
})
study.union_search_space = Array.from(unionSearchSpace).map((name) => {
return nameToSearchSpaceItem.get(name) as Optuna.SearchSpaceItem
})
study.intersection_search_space = intersectionSearchSpace.map((name) => {
return nameToSearchSpaceItem.get(name) as Optuna.SearchSpaceItem
})
study.union_search_space = unionSearchSpace
study.intersection_search_space = intersectionSearchSpace
}
return this.studies
+43 -21
View File
@@ -3,6 +3,36 @@ import * as Optuna from "@optuna/types"
import sqlite3InitModule from "@sqlite.org/sqlite-wasm"
import { OptunaStorage } from "./storage"
// TODO(porink0424): Refactor to common function with journal.ts (current workaround duplicates code due to missing file extensions in tsc build output).
const isDistributionEqual = (
a: Optuna.Distribution,
b: Optuna.Distribution
) => {
if (a.type !== b.type) {
return false
}
if (a.type === "IntDistribution" || a.type === "FloatDistribution") {
if (b.type !== "IntDistribution" && b.type !== "FloatDistribution") {
throw new Error("Invalid distribution type")
}
return (
a.low === b.low &&
a.high === b.high &&
a.step === b.step &&
a.log === b.log
)
}
if (a.type === "CategoricalDistribution") {
if (b.type !== "CategoricalDistribution") {
throw new Error("Invalid distribution type")
}
return JSON.stringify(a.choices) === JSON.stringify(b.choices)
}
throw new Error("Invalid distribution type")
}
type SQLite3DB = {
exec(options: {
sql: string
@@ -149,7 +179,7 @@ const getStudy = (
study.metric_names = studySystemAttrs.metric_names
}
let intersection_search_space: Set<Optuna.SearchSpaceItem> = new Set()
let intersectionSearchSpace: Optuna.SearchSpaceItem[] = []
study.trials = getTrials(db, summary.id, schemaVersion)
for (const trial of study.trials) {
const userAttrs = getTrialUserAttributes(db, trial.trial_id)
@@ -165,16 +195,7 @@ const getStudy = (
}
const params = getTrialParams(db, trial.trial_id)
const paramNames = new Set<string>()
const paramNameToSearchSpaceItem = new Map<string, Optuna.SearchSpaceItem>()
for (const param of params) {
paramNames.add(param.name)
if (paramNameToSearchSpaceItem.has(param.name)) {
paramNameToSearchSpaceItem.set(param.name, {
name: param.name,
distribution: param.distribution,
})
}
if (
study.union_search_space.findIndex((s) => s.name === param.name) === -1
) {
@@ -184,23 +205,24 @@ const getStudy = (
})
}
}
if (intersection_search_space.size === 0) {
for (const s of paramNames) {
intersection_search_space.add(
paramNameToSearchSpaceItem.get(s) as Optuna.SearchSpaceItem
)
}
if (intersectionSearchSpace.length === 0) {
intersectionSearchSpace = params.map((param) => ({
name: param.name,
distribution: param.distribution,
}))
} else {
intersection_search_space = new Set(
Array.from(intersection_search_space).filter((s) =>
paramNames.has(s.name)
intersectionSearchSpace = intersectionSearchSpace.filter((item) => {
return params.some(
(param) =>
item.name === param.name &&
isDistributionEqual(item.distribution, param.distribution)
)
)
})
}
trial.params = params
trial.user_attrs = userAttrs
}
study.intersection_search_space = Array.from(intersection_search_space)
study.intersection_search_space = intersectionSearchSpace
return study
}