Merge pull request #574 from c-bata/improve-sqlite3-wasm-loader

Add some improvements on SQLite3 WASM loader
This commit is contained in:
c-bata
2023-08-27 01:52:21 +09:00
committed by GitHub
2 changed files with 191 additions and 148 deletions
+183 -144
View File
@@ -2,6 +2,14 @@
import sqlite3InitModule from "@sqlite.org/sqlite-wasm"
import { SetterOrUpdater } from "recoil"
type SQLite3DB = {
exec(options: {
sql: string
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (...args: any[]) => void
}): void
}
export const loadStorage = (
arrayBuffer: ArrayBuffer,
setter: SetterOrUpdater<Study[]>
@@ -30,155 +38,186 @@ export const loadStorage = (
)
db.checkRc(rc)
try {
// Check version_info table
let supported = true
db.exec({
sql: "SELECT schema_version FROM version_info LIMIT 1",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
if (vals[0] != 12) {
supported = false
}
},
})
if (!supported) {
if (!isSupportedSchema(db)) {
return
}
// Get studies
const studies: Study[] = []
db.exec({
sql:
"SELECT s.study_id, s.study_name, sd.direction, sd.objective" +
" FROM studies AS s INNER JOIN study_directions AS sd" +
" ON s.study_id = sd.study_id ORDER BY sd.study_direction_id",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const study_id = vals[0]
const study_name = vals[1]
const direction: StudyDirection = vals[2].toLowerCase()
const objective = vals[3]
let index = 0
if (objective === 0) {
studies.push({
study_id: study_id,
study_name: study_name,
directions: [direction],
union_search_space: [],
intersection_search_space: [],
user_attrs: [],
system_attrs: [],
trials: [],
})
} else {
index = studies.findIndex((s) => s.study_id === study_id)
studies[index].directions.push(direction)
}
},
})
studies.forEach((s) => {
db.exec({
sql:
"SELECT t.trial_id, t.number, t.study_id, t.state, t.datetime_start, t.datetime_complete," +
" tv.objective, tv.value, tv.value_type" +
" FROM trials AS t LEFT JOIN trial_values AS tv ON tv.trial_id = t.trial_id" +
` WHERE t.study_id = ${s.study_id}` +
" ORDER BY t.number",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const state: TrialState =
vals[3] === "COMPLETE"
? "Complete"
: vals[3] === "PRUNED"
? "Pruned"
: vals[3] === "RUNNING"
? "Running"
: vals[3] === "WAITING"
? "Waiting"
: "Fail"
const trial: Trial = {
trial_id: vals[0],
number: vals[1],
study_id: vals[2],
state: state,
params: [],
intermediate_values: [],
user_attrs: [],
system_attrs: [],
}
s.trials.push(trial)
},
})
const union_search_space: SearchSpaceItem[] = []
let intersection_search_space: Set<SearchSpaceItem> = new Set()
s.trials.forEach((trial) => {
const params: TrialParam[] = []
const param_names = new Set<string>()
db.exec({
sql:
"SELECT param_name, param_value" +
` FROM trial_params WHERE trial_id = ${trial.trial_id}`,
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const param_name = vals[0]
params.push({
name: param_name,
param_internal_value: vals[1],
})
param_names.add(param_name)
if (
union_search_space.findIndex((s) => s.name === param_name) == -1
) {
union_search_space.push({ name: param_name })
}
},
})
if (intersection_search_space.size === 0) {
param_names.forEach((s) => {
intersection_search_space.add({
name: s,
})
})
} else {
intersection_search_space = new Set(
Array.from(intersection_search_space).filter((s) =>
param_names.has(s.name)
)
)
}
trial.params = params
const values: TrialValueNumber[] = []
db.exec({
sql:
"SELECT value, value_type" +
` FROM trial_values WHERE trial_id = ${trial.trial_id}` +
" ORDER BY objective",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
values.push(
vals[1] === "INF_NEG"
? "-inf"
: vals[1] === "INF_POS"
? "+inf"
: vals[0]
)
},
})
if (s.directions.length === values.length) {
trial.values = values
}
})
s.union_search_space = union_search_space
s.intersection_search_space = Array.from(intersection_search_space)
})
const studies = getStudies(db)
setter((prev) => [...prev, ...studies])
} finally {
db.close()
}
})
}
const isSupportedSchema = (db: SQLite3DB): boolean => {
let supported = true
db.exec({
sql: "SELECT schema_version FROM version_info LIMIT 1",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
if (vals[0] != 12) {
supported = false
}
},
})
return supported
}
const getStudies = (db: SQLite3DB): Study[] => {
const studies: Study[] = []
db.exec({
sql:
"SELECT s.study_id, s.study_name, sd.direction, sd.objective" +
" FROM studies AS s INNER JOIN study_directions AS sd" +
" ON s.study_id = sd.study_id ORDER BY sd.study_direction_id",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const studyId = vals[0]
const studyName = vals[1]
const direction: StudyDirection =
vals[2] === "MINIMIZE" ? "minimize" : "maximize"
const objective = vals[3]
const trials = getTrials(db, studyId)
const union_search_space: SearchSpaceItem[] = []
let intersection_search_space: Set<SearchSpaceItem> = new Set()
trials.forEach((trial) => {
const params: TrialParam[] = []
const param_names = new Set<string>()
db.exec({
sql:
"SELECT param_name, param_value" +
` FROM trial_params WHERE trial_id = ${trial.trial_id}`,
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const param_name = vals[0]
// TODO(c-bata): Support param_external_value
params.push({
name: param_name,
param_internal_value: vals[1],
})
param_names.add(param_name)
if (
union_search_space.findIndex((s) => s.name === param_name) == -1
) {
union_search_space.push({ name: param_name })
}
},
})
if (intersection_search_space.size === 0) {
param_names.forEach((s) => {
intersection_search_space.add({
name: s,
})
})
} else {
intersection_search_space = new Set(
Array.from(intersection_search_space).filter((s) =>
param_names.has(s.name)
)
)
}
trial.params = params
})
if (objective === 0) {
studies.push({
study_id: studyId,
study_name: studyName,
directions: [direction],
union_search_space: union_search_space,
intersection_search_space: Array.from(intersection_search_space),
trials: trials,
})
return
}
const index = studies.findIndex((s) => s.study_id === studyId)
studies[index].directions.push(direction)
},
})
return studies
}
const getTrials = (db: SQLite3DB, studyId: number): Trial[] => {
const trials: Trial[] = []
db.exec({
sql:
"SELECT trial_id, number, state, datetime_start, datetime_complete FROM trials" +
` WHERE study_id = ${studyId} ORDER BY number`,
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const trialId = vals[0]
const state: TrialState =
vals[2] === "COMPLETE"
? "Complete"
: vals[2] === "PRUNED"
? "Pruned"
: vals[2] === "RUNNING"
? "Running"
: vals[2] === "WAITING"
? "Waiting"
: "Fail"
const trial: Trial = {
trial_id: trialId,
number: vals[1],
study_id: studyId,
state: state,
values: getTrialValues(db, trialId),
params: [], // Set this column later
intermediate_values: getTrialIntermediateValues(db, trialId),
datetime_start: vals[3],
datetime_complete: vals[4],
}
trials.push(trial)
},
})
return trials
}
const getTrialValues = (db: SQLite3DB, trialId: number): TrialValueNumber[] => {
const values: TrialValueNumber[] = []
db.exec({
sql:
"SELECT value, value_type" +
` FROM trial_values WHERE trial_id = ${trialId}` +
" ORDER BY objective",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
values.push(
vals[1] === "INF_NEG"
? "-inf"
: vals[1] === "INF_POS"
? "+inf"
: vals[0]
)
},
})
return values
}
const getTrialIntermediateValues = (
db: SQLite3DB,
trialId: number
): TrialIntermediateValue[] => {
const values: TrialIntermediateValue[] = []
db.exec({
sql:
"SELECT step, intermediate_value, intermediate_value_type" +
` FROM trial_values WHERE trial_id = ${trialId}` +
" ORDER BY step",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
values.push({
step: vals[0],
value:
vals[2] === "INF_NEG"
? "-inf"
: vals[2] === "INF_POS"
? "+inf"
: vals[1],
})
},
})
return values
}
+8 -4
View File
@@ -27,6 +27,11 @@ type CategoricalDistribution = {
choices: { pytype: string; value: string }[]
}
type TrialIntermediateValue = {
step: number
value: TrialIntermediateValueNumber
}
type Distribution =
| FloatDistribution
| IntDistribution
@@ -41,10 +46,8 @@ type Study = {
study_id: number
study_name: string
directions: StudyDirection[]
user_attrs: Attribute[]
union_search_space: SearchSpaceItem[]
intersection_search_space: SearchSpaceItem[]
system_attrs: Attribute[]
datetime_start?: Date
trials: Trial[]
}
@@ -59,13 +62,14 @@ type Trial = {
intermediate_values: TrialIntermediateValue[]
datetime_start?: Date
datetime_complete?: Date
user_attrs: Attribute[]
system_attrs: Attribute[]
}
type TrialParam = {
name: string
param_internal_value: number
// param_external_value: string
// param_external_type: string
// distribution: Distribution
}
type SearchSpaceItem = {