Support for recent changes to standalone_app

This commit is contained in:
porink0424 committed 2024-03-22 10:38:55 +09:00
1 parent 52a72aa5bd
commit 083946ce5c
3 files changed
+393 -347

No files matched your search

+16 -6
View File
@@ -1,11 +1,14 @@
declare const IS_VSCODE: boolean
type TrialValueNumber = number | "inf" | "-inf"
type TrialIntermediateValueNumber = number | "inf" | "-inf" | "nan"
type TrialState = "Running" | "Complete" | "Pruned" | "Fail" | "Waiting"
type TrialStateFinished = "Complete" | "Fail" | "Pruned"
type StudyDirection = "maximize" | "minimize" | "not_set"
type OptunaStorage = {
getStudies: () => Promise<StudySummary[]>
getStudy: (idx: number) => Promise<Study | null>
}
type FloatDistribution = {
type: "FloatDistribution"
low: number
@@ -22,14 +25,15 @@ type IntDistribution = {
log: boolean
}
type CategoricalChoiceType = null | boolean | number | string
type CategoricalDistribution = {
type: "CategoricalDistribution"
choices: { pytype: string; value: string }[]
choices: CategoricalChoiceType[]
}
type TrialIntermediateValue = {
step: number
value: TrialIntermediateValueNumber
value: number
}
type Distribution =
@@ -47,6 +51,12 @@ type AttributeSpec = {
sortable: boolean
}
type StudySummary = {
study_id: number
study_name: string
directions: StudyDirection[]
}
type Study = {
study_id: number
study_name: string
@@ -63,7 +73,7 @@ type Trial = {
number: number
study_id: number
state: TrialState
values?: TrialValueNumber[]
values?: number[]
params: TrialParam[]
intermediate_values: TrialIntermediateValue[]
user_attrs: Attribute[]
@@ -74,7 +84,7 @@ type Trial = {
type TrialParam = {
name: string
param_internal_value: number
param_external_value: string
param_external_value: CategoricalChoiceType
param_external_type: string
distribution: Distribution
}
+140 -142
View File
@@ -1,4 +1,4 @@
import { SetterOrUpdater } from "recoil";
import { SetterOrUpdater } from "recoil"
// JournalStorage
enum JournalOperation {
@@ -15,92 +15,92 @@ enum JournalOperation {
}
interface JournalOpBase {
op_code: JournalOperation;
workor_id: string;
op_code: JournalOperation
workor_id: string
}
interface JournalOpCreateStudy extends JournalOpBase {
study_name: string;
directions: number[]; // TODO(gen740): introduce Study Direction enum
study_name: string
directions: number[] // TODO(gen740): introduce Study Direction enum
}
interface JournalOpDeleteStudy extends JournalOpBase {
study_id: number;
study_id: number
}
interface JournalOpCreateTrial extends JournalOpBase {
study_id: number;
datetime_start?: string;
datetime_complete?: string;
distributions?: { [key: string]: string };
study_id: number
datetime_start?: string
datetime_complete?: string
distributions?: { [key: string]: string }
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
params?: { [key: string]: any };
params?: { [key: string]: any }
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
user_attrs?: { [key: string]: any };
user_attrs?: { [key: string]: any }
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
system_attrs?: { [key: string]: any };
state?: number;
intermediate_values?: { [key: string]: number };
value?: number;
values?: number[];
system_attrs?: { [key: string]: any }
state?: number
intermediate_values?: { [key: string]: number }
value?: number
values?: number[]
}
interface JournalOpSetTrialParam extends JournalOpBase {
trial_id: number;
param_name: string;
param_value_internal: number;
distribution: string;
trial_id: number
param_name: string
param_value_internal: number
distribution: string
}
interface JournalOpSetTrialStateValue extends JournalOpBase {
trial_id: number;
state: number;
values?: number[];
datetime_start?: string;
datetime_complete?: string;
trial_id: number
state: number
values?: number[]
datetime_start?: string
datetime_complete?: string
}
interface JournalOpSetTrialIntermediateValue extends JournalOpBase {
trial_id: number;
step: number;
intermediate_value: number;
trial_id: number
step: number
intermediate_value: number
}
interface JournalOpSetTrialUserAttr extends JournalOpBase {
trial_id: number;
trial_id: number
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
user_attr: { [key: string]: any };
user_attr: { [key: string]: any }
}
const trialStateNumToTrialState = (state: number): TrialState => {
switch (state) {
case 0:
return "Running";
return "Running"
case 1:
return "Complete";
return "Complete"
case 2:
return "Pruned";
return "Pruned"
case 3:
return "Fail";
return "Fail"
case 4:
return "Waiting";
return "Waiting"
default:
return "Running";
return "Running"
}
};
}
const parseDistribution = (distribution: string): Distribution => {
const distributionJson = JSON.parse(distribution);
const distributionJson = JSON.parse(distribution)
if (distributionJson.name === "IntDistribution") {
return {
...distributionJson.attributes,
type: "IntDistribution",
};
}
} else if (distributionJson.name === "FloatDistribution") {
return {
...distributionJson.attributes,
type: "FloatDistribution",
};
}
} else {
return {
// TODO(gen740): support other types
@@ -110,59 +110,59 @@ const parseDistribution = (distribution: string): Distribution => {
return {
pytype: "str",
value: choice.toString(),
};
}
}),
};
}
}
};
}
class JournalStorage {
private studies: Study[] = [];
private nextStudyId = 0;
private studyIdToTrialIDs: Map<number, number[]> = new Map();
private trialIdToStudyId: Map<number, number> = new Map();
private trialID = 0;
private studies: Study[] = []
private nextStudyId = 0
private studyIdToTrialIDs: Map<number, number[]> = new Map()
private trialIdToStudyId: Map<number, number> = new Map()
private trialID = 0
public getStudies(): Study[] {
for (const study of this.studies) {
const unionUserAttrs: Set<string> = new Set();
const unionSearchSpace: Set<string> = new Set();
let intersectionSearchSpace: string[] = [];
const unionUserAttrs: Set<string> = new Set()
const unionSearchSpace: Set<string> = new Set()
let intersectionSearchSpace: string[] = []
study.trials.forEach((trial, index) => {
for (const userAttr of trial.user_attrs) {
unionUserAttrs.add(userAttr.key);
unionUserAttrs.add(userAttr.key)
}
for (const param of trial.params) {
unionSearchSpace.add(param.name);
unionSearchSpace.add(param.name)
}
if (index === 0) {
intersectionSearchSpace = Array.from(unionSearchSpace);
intersectionSearchSpace = Array.from(unionSearchSpace)
} else {
intersectionSearchSpace = intersectionSearchSpace.filter((name) => {
return trial.params.some((param) => param.name === name);
});
return trial.params.some((param) => param.name === name)
})
}
});
})
study.union_user_attrs = Array.from(unionUserAttrs).map((key) => {
return {
key: key,
sortable: false,
};
});
}
})
study.union_search_space = Array.from(unionSearchSpace).map((name) => {
return {
name: name,
};
});
}
})
study.intersection_search_space = intersectionSearchSpace.map((name) => {
return {
name: name,
};
});
}
})
}
return this.studies;
return this.studies
}
public applyCreateStudy(log: JournalOpCreateStudy): void {
@@ -174,22 +174,20 @@ class JournalStorage {
intersection_search_space: [],
union_user_attrs: [],
trials: [],
});
this.nextStudyId++;
})
this.nextStudyId++
}
public applyDeleteStudy(log: JournalOpDeleteStudy): void {
this.studies = this.studies.filter(
(item) => item.study_id !== log.study_id,
);
this.studies = this.studies.filter((item) => item.study_id !== log.study_id)
}
public applyCreateTrial(log: JournalOpCreateTrial): void {
const thisStudy = this.studies.find(
(item) => item.study_id === log.study_id,
);
(item) => item.study_id === log.study_id
)
if (thisStudy === undefined) {
return;
return
}
const params: TrialParam[] =
@@ -197,32 +195,32 @@ class JournalStorage {
? []
: Object.entries(log.params).map(([name, value]) => {
// biome-ignore lint/style/noNonNullAssertion: <explanation>
const distribution = parseDistribution(log.distributions![name]);
const distribution = parseDistribution(log.distributions![name])
return {
name: name,
param_internal_value: value,
param_external_type: distribution.type,
param_external_value: (() => {
if (distribution.type === "FloatDistribution") {
return value.toString();
return value.toString()
} else if (distribution.type === "IntDistribution") {
return value.toString();
return value.toString()
} else {
return distribution.choices[value].value;
return distribution.choices[value]
}
})(),
distribution: distribution,
};
});
}
})
const userAtter = log.user_attrs
? Object.entries(log.user_attrs).map(([key, value]) => {
return {
key: key,
value: value,
};
}
})
: [];
: []
thisStudy.trials.push({
trial_id: this.trialID,
@@ -231,11 +229,11 @@ class JournalStorage {
state: trialStateNumToTrialState(log.state ?? 0),
values: (() => {
if (log.value !== undefined) {
return [log.value];
return [log.value]
} else if (log.values !== undefined) {
return log.values;
return log.values
} else {
return undefined;
return undefined
}
})(),
params: params,
@@ -247,36 +245,36 @@ class JournalStorage {
datetime_complete: log.datetime_complete
? new Date(log.datetime_complete)
: undefined,
});
})
this.studyIdToTrialIDs.set(
log.study_id,
this.studyIdToTrialIDs.get(log.study_id)?.concat([this.trialID]) ?? [
this.trialID,
],
);
this.trialIdToStudyId.set(this.trialID, log.study_id);
this.trialID++;
]
)
this.trialIdToStudyId.set(this.trialID, log.study_id)
this.trialID++
}
private getStudyAndTrial(trial_id: number): [Study?, Trial?] {
const study = this.studies.find(
(item) => item.study_id === this.trialIdToStudyId.get(trial_id),
);
(item) => item.study_id === this.trialIdToStudyId.get(trial_id)
)
if (study === undefined) {
return [undefined, undefined];
return [undefined, undefined]
}
const trial = study.trials.find((item) => item.trial_id === trial_id);
const trial = study.trials.find((item) => item.trial_id === trial_id)
if (trial === undefined) {
return [study, undefined];
return [study, undefined]
}
return [study, trial];
return [study, trial]
}
public applySetTrialParam(log: JournalOpSetTrialParam) {
const [thisStudy, thisTrial] = this.getStudyAndTrial(log.trial_id);
const [thisStudy, thisTrial] = this.getStudyAndTrial(log.trial_id)
if (thisStudy === undefined || thisTrial === undefined) {
return;
return
}
thisTrial.params.push({
name: log.param_name,
@@ -284,51 +282,51 @@ class JournalStorage {
param_external_type: "FloatDistribution",
param_external_value: log.param_value_internal.toString(),
distribution: parseDistribution(log.distribution),
});
})
}
public applySetTrialStateValues(log: JournalOpSetTrialStateValue): void {
const [thisStudy, thisTrial] = this.getStudyAndTrial(log.trial_id);
const [thisStudy, thisTrial] = this.getStudyAndTrial(log.trial_id)
if (thisStudy === undefined || thisTrial === undefined) {
return;
return
}
thisTrial.state = trialStateNumToTrialState(log.state);
thisTrial.values = log.values;
thisTrial.state = trialStateNumToTrialState(log.state)
thisTrial.values = log.values
thisTrial.datetime_start = log.datetime_start
? new Date(log.datetime_start)
: undefined;
: undefined
thisTrial.datetime_complete = log.datetime_complete
? new Date(log.datetime_complete)
: undefined;
: undefined
}
public applySetTrialIntermediateValue(
log: JournalOpSetTrialIntermediateValue,
log: JournalOpSetTrialIntermediateValue
) {
const [thisStudy, thisTrial] = this.getStudyAndTrial(log.trial_id);
const [thisStudy, thisTrial] = this.getStudyAndTrial(log.trial_id)
if (thisStudy === undefined || thisTrial === undefined) {
return;
return
}
thisTrial.intermediate_values.push({
step: log.step,
value: log.intermediate_value,
});
})
}
public applySetTrialUserAttr(log: JournalOpSetTrialUserAttr) {
const [thisStudy, thisTrial] = this.getStudyAndTrial(log.trial_id);
const [thisStudy, thisTrial] = this.getStudyAndTrial(log.trial_id)
if (thisStudy === undefined || thisTrial === undefined) {
return;
return
}
for (const [key, value] of Object.entries(log.user_attr)) {
const index = thisTrial.user_attrs.findIndex((item) => item.key === key);
const index = thisTrial.user_attrs.findIndex((item) => item.key === key)
if (index !== -1) {
thisTrial.user_attrs[index].value = value.toString();
thisTrial.user_attrs[index].value = value.toString()
} else {
thisTrial.user_attrs.push({
key: key,
value: value.toString(),
});
})
}
}
}
@@ -336,58 +334,58 @@ class JournalStorage {
export const loadJournalStorage = (
arrayBuffer: ArrayBuffer,
setter: SetterOrUpdater<Study[]>,
setter: SetterOrUpdater<Study[]>
): void => {
const decoder = new TextDecoder("utf-8");
const logs = decoder.decode(arrayBuffer).split("\n");
const decoder = new TextDecoder("utf-8")
const logs = decoder.decode(arrayBuffer).split("\n")
const journalStorage = new JournalStorage();
const journalStorage = new JournalStorage()
for (const log of logs) {
if (log === "") {
continue;
continue
}
const parsedLog: JournalOpBase = JSON.parse(log);
const parsedLog: JournalOpBase = JSON.parse(log)
switch (parsedLog.op_code) {
case JournalOperation.CREATE_STUDY:
journalStorage.applyCreateStudy(parsedLog as JournalOpCreateStudy);
break;
journalStorage.applyCreateStudy(parsedLog as JournalOpCreateStudy)
break
case JournalOperation.DELETE_STUDY:
journalStorage.applyDeleteStudy(parsedLog as JournalOpDeleteStudy);
break;
journalStorage.applyDeleteStudy(parsedLog as JournalOpDeleteStudy)
break
case JournalOperation.SET_STUDY_USER_ATTR:
// Unsupported
break;
break
case JournalOperation.SET_STUDY_SYSTEM_ATTR:
// Unsupported
break;
break
case JournalOperation.CREATE_TRIAL:
journalStorage.applyCreateTrial(parsedLog as JournalOpCreateTrial);
break;
journalStorage.applyCreateTrial(parsedLog as JournalOpCreateTrial)
break
case JournalOperation.SET_TRIAL_PARAM:
journalStorage.applySetTrialParam(parsedLog as JournalOpSetTrialParam);
break;
journalStorage.applySetTrialParam(parsedLog as JournalOpSetTrialParam)
break
case JournalOperation.SET_TRIAL_STATE_VALUES:
journalStorage.applySetTrialStateValues(
parsedLog as JournalOpSetTrialStateValue,
);
break;
parsedLog as JournalOpSetTrialStateValue
)
break
case JournalOperation.SET_TRIAL_INTERMEDIATE_VALUE:
journalStorage.applySetTrialIntermediateValue(
parsedLog as JournalOpSetTrialIntermediateValue,
);
break;
parsedLog as JournalOpSetTrialIntermediateValue
)
break
case JournalOperation.SET_TRIAL_USER_ATTR:
journalStorage.applySetTrialUserAttr(
parsedLog as JournalOpSetTrialUserAttr,
);
break;
parsedLog as JournalOpSetTrialUserAttr
)
break
case JournalOperation.SET_TRIAL_SYSTEM_ATTR:
// Unsupported
break;
break
}
}
const studies = journalStorage.getStudies();
setter((prev) => [...prev, ...studies]);
};
const studies = journalStorage.getStudies()
setter((prev) => [...prev, ...studies])
}
+237 -199
View File
@@ -1,87 +1,118 @@
import sqlite3InitModule from "@sqlite.org/sqlite-wasm";
import { SetterOrUpdater } from "recoil";
type SQLite3DB = {
exec(options: {
sql: string;
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
callback: (...args: any[]) => void;
}): void;
};
// @ts-ignore
import sqlite3InitModule from "@sqlite.org/sqlite-wasm"
import { SetterOrUpdater } from "recoil"
export const loadSQLite3Storage = async (
arrayBuffer: ArrayBuffer,
setter: SetterOrUpdater<Study[]>,
setter: SetterOrUpdater<Study[]>
) => {
const sqlite3 = await sqlite3InitModule({
const sqlite3Storage = new SQLite3Storage(arrayBuffer)
const studySummaries = await sqlite3Storage.getStudies()
const studies = (
await Promise.all(
studySummaries.map((summary) => sqlite3Storage.getStudy(summary.study_id))
)
).filter((s) => s !== null) as Study[]
setter((prev) => [...prev, ...studies])
}
type SQLite3DB = {
exec(options: {
sql: string
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
print: (...args: any): void => {
console.log(args);
},
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
printErr: (...args: any): void => {
console.log(args);
},
// @ts-ignore
});
const p = sqlite3.wasm.allocFromTypedArray(arrayBuffer);
const db = new sqlite3.oo1.DB();
const rc = sqlite3.capi.sqlite3_deserialize(
// @ts-ignore
db.pointer,
"main",
p,
arrayBuffer.byteLength,
arrayBuffer.byteLength,
sqlite3.capi.SQLITE_DESERIALIZE_FREEONCLOSE,
);
db.checkRc(rc);
try {
const schemaVersion = getSchemaVersion(db);
if (!isSupportedSchema(schemaVersion)) {
return;
}
const studies = getStudies(db, schemaVersion);
setter((prev) => [...prev, ...studies]);
} finally {
db.close();
callback: (...args: any[]) => void
}): void
}
export class SQLite3Storage implements OptunaStorage {
db: Promise<SQLite3DB>
summaries_cache: StudySummary[] | null
constructor(arrayBuffer: ArrayBuffer) {
this.db = this.initDB(arrayBuffer)
this.summaries_cache = null
}
};
async initDB(arrayBuffer: ArrayBuffer): Promise<SQLite3DB> {
return sqlite3InitModule({
print: console.log,
printErr: console.log,
// @ts-ignore
}).then((sqlite3) => {
const p = sqlite3.wasm.allocFromTypedArray(arrayBuffer)
const db = new sqlite3.oo1.DB()
const rc = sqlite3.capi.sqlite3_deserialize(
// @ts-ignore
db.pointer,
"main",
p,
arrayBuffer.byteLength,
arrayBuffer.byteLength,
sqlite3.capi.SQLITE_DESERIALIZE_FREEONCLOSE
)
db.checkRc(rc)
return db
})
}
getStudies = async (): Promise<StudySummary[]> => {
const db = await this.db
this.summaries_cache = getStudySummaries(db)
return this.summaries_cache
}
getStudy = async (idx: number): Promise<Study | null> => {
const db = await this.db
const schemaVersion = getSchemaVersion(db)
if (!isSupportedSchema(schemaVersion)) {
return null
}
if (this.summaries_cache === null) {
this.summaries_cache = getStudySummaries(db)
}
const summary = this.summaries_cache[idx]
if (summary === undefined) {
return null
}
return getStudy(db, schemaVersion, summary)
}
}
const getSchemaVersion = (db: SQLite3DB): string => {
let schemaVersion = "";
let schemaVersion = ""
db.exec({
sql: "SELECT version_num FROM alembic_version LIMIT 1",
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
callback: (vals: any[]) => {
schemaVersion = vals[0];
schemaVersion = vals[0]
},
});
return schemaVersion;
};
})
return schemaVersion
}
const isSupportedSchema = (schemaVersion: string): boolean => {
const lowestVersion = "v2.6.0.a"; // supported: "v3.2.0.a", "v3.0.0.{a,b,c,d}", "v2.6.0.a"
if (schemaVersion === lowestVersion) return true;
return isGreaterSchemaVersion(schemaVersion, lowestVersion);
};
const lowestVersion = "v2.6.0.a" // supported: "v3.2.0.a", "v3.0.0.{a,b,c,d}", "v2.6.0.a"
if (schemaVersion === lowestVersion) return true
return isGreaterSchemaVersion(schemaVersion, lowestVersion)
}
const isGreaterSchemaVersion = (
leftVersion: string,
rightVersion: string,
rightVersion: string
): boolean => {
// return leftVersion > rightVersion
const leftSuffix = leftVersion.split(".").reverse()[0];
const rightSuffix = rightVersion.split(".").reverse()[0];
const leftSuffix = leftVersion.split(".").reverse()[0]
const rightSuffix = rightVersion.split(".").reverse()[0]
const leftVersion_ = leftVersion.replace(/\D/g, "")
const rightVersion_ = rightVersion.replace(/\D/g, "")
const left = Number(leftVersion.replace(/\D/g, ""));
const right = Number(rightVersion.replace(/\D/g, ""));
if (left === right) return leftSuffix > rightSuffix;
return left > right;
};
const left = Number(leftVersion_)
const right = Number(rightVersion_)
if (left === right) return leftSuffix > rightSuffix
return left > right
}
const getStudies = (db: SQLite3DB, schemaVersion: string): Study[] => {
const studies: Study[] = [];
const getStudySummaries = (db: SQLite3DB): StudySummary[] => {
const summaries: StudySummary[] = []
db.exec({
sql:
"SELECT s.study_id, s.study_name, sd.direction, sd.objective" +
@@ -89,81 +120,91 @@ const getStudies = (db: SQLite3DB, schemaVersion: string): Study[] => {
" ON s.study_id = sd.study_id ORDER BY sd.study_direction_id",
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
callback: (vals: any[]) => {
const studyId = vals[0];
const studyName = vals[1];
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, schemaVersion);
const union_search_space: SearchSpaceItem[] = [];
const union_user_attrs: AttributeSpec[] = [];
let intersection_search_space: Set<SearchSpaceItem> = new Set();
trials.forEach((trial) => {
const userAttrs = getTrialUserAttributes(db, trial.trial_id);
userAttrs.forEach((attr) => {
if (union_user_attrs.findIndex((s) => s.key === attr.key) === -1) {
union_user_attrs.push({ key: attr.key, sortable: false });
}
});
const params = getTrialParams(db, trial.trial_id);
const param_names = new Set<string>();
params.forEach((param) => {
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;
trial.user_attrs = userAttrs;
});
vals[2] === "MINIMIZE" ? "minimize" : "maximize"
const objective = vals[3]
if (objective === 0) {
studies.push({
summaries.push({
study_id: studyId,
study_name: studyName,
directions: [direction],
union_search_space: union_search_space,
intersection_search_space: Array.from(intersection_search_space),
union_user_attrs: union_user_attrs,
trials: trials,
});
return;
})
return
}
const index = studies.findIndex((s) => s.study_id === studyId);
studies[index].directions.push(direction);
const index = summaries.findIndex((s) => s.study_id === studyId)
summaries[index].directions.push(direction)
},
});
return studies;
};
})
return summaries
}
const getStudy = (
db: SQLite3DB,
schemaVersion: string,
summary: StudySummary
): Study => {
const study: Study = {
study_id: summary.study_id,
study_name: summary.study_name,
directions: summary.directions,
union_search_space: [],
intersection_search_space: [],
union_user_attrs: [],
trials: [],
}
let intersection_search_space: Set<SearchSpaceItem> = new Set()
study.trials = getTrials(db, summary.study_id, schemaVersion)
for (const trial of study.trials) {
const userAttrs = getTrialUserAttributes(db, trial.trial_id)
for (const attr of userAttrs) {
if (study.union_user_attrs.findIndex((s) => s.key === attr.key) === -1) {
study.union_user_attrs.push({ key: attr.key, sortable: false })
}
}
const params = getTrialParams(db, trial.trial_id)
const param_names = new Set<string>()
for (const param of params) {
param_names.add(param.name)
if (
study.union_search_space.findIndex((s) => s.name === param.name) === -1
) {
study.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
trial.user_attrs = userAttrs
}
study.intersection_search_space = Array.from(intersection_search_space)
return study
}
const getTrials = (
db: SQLite3DB,
studyId: number,
schemaVersion: string,
schemaVersion: string
): Trial[] => {
const trials: 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`,
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
callback: (vals: any[]) => {
const trialId = vals[0];
const trialId = vals[0]
const state: TrialState =
vals[2] === "COMPLETE"
? "Complete"
@@ -173,7 +214,7 @@ const getTrials = (
? "Running"
: vals[2] === "WAITING"
? "Waiting"
: "Fail";
: "Fail"
const trial: Trial = {
trial_id: trialId,
number: vals[1],
@@ -183,25 +224,25 @@ const getTrials = (
intermediate_values: getTrialIntermediateValues(
db,
trialId,
schemaVersion,
schemaVersion
),
params: [], // Set this column later
user_attrs: [], // Set this column later
datetime_start: vals[3],
datetime_complete: vals[4],
};
trials.push(trial);
}
trials.push(trial)
},
});
return trials;
};
})
return trials
}
const getTrialValues = (
db: SQLite3DB,
trialId: number,
schemaVersion: string,
): TrialValueNumber[] => {
const values: TrialValueNumber[] = [];
schemaVersion: string
): number[] => {
const values: number[] = []
if (isGreaterSchemaVersion(schemaVersion, "v3.0.0.c")) {
db.exec({
sql: `SELECT value, value_type FROM trial_values WHERE trial_id = ${trialId} ORDER BY objective`,
@@ -209,62 +250,62 @@ const getTrialValues = (
callback: (vals: any[]) => {
values.push(
vals[1] === "INF_NEG"
? "-inf"
? -Infinity
: vals[1] === "INF_POS"
? "+inf"
: vals[0],
);
? Infinity
: vals[0]
)
},
});
})
} else {
db.exec({
sql: `SELECT value FROM trial_values WHERE trial_id = ${trialId} ORDER BY objective`,
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
callback: (vals: any[]) => {
values.push(vals[0]);
values.push(vals[0])
},
});
})
}
return values;
};
return values
}
const getTrialParams = (db: SQLite3DB, trialId: number): TrialParam[] => {
const params: TrialParam[] = [];
const params: TrialParam[] = []
db.exec({
sql: `SELECT param_name, param_value, distribution_json FROM trial_params WHERE trial_id = ${trialId}`,
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
callback: (vals: any[]) => {
const distribution = parseDistributionJSON(vals[2]);
const distribution = parseDistributionJSON(vals[2])
params.push({
name: vals[0],
param_internal_value: vals[1],
param_external_type: distribution.type,
param_external_value: paramInternalValueToExternalValue(
distribution,
vals[1],
vals[1]
),
distribution: distribution,
});
})
},
});
return params;
};
})
return params
}
const paramInternalValueToExternalValue = (
distribution: Distribution,
internalValue: number,
): string => {
internalValue: number
): CategoricalChoiceType => {
if (distribution.type === "FloatDistribution") {
return internalValue.toString();
} else if (distribution.type === "IntDistribution") {
return internalValue.toString();
} else {
return distribution.choices[internalValue].value;
return internalValue.toString()
}
};
if (distribution.type === "IntDistribution") {
return internalValue.toString()
}
return distribution.choices[internalValue]
}
const parseDistributionJSON = (t: string): Distribution => {
const parsed = JSON.parse(t);
const parsed = JSON.parse(t)
if (parsed.name === "FloatDistribution") {
return {
type: "FloatDistribution",
@@ -272,76 +313,73 @@ const parseDistributionJSON = (t: string): Distribution => {
high: parsed.attributes.high as number,
step: parsed.attributes.step as number,
log: parsed.attributes.log as boolean,
};
} else if (parsed.name === "UniformDistribution") {
}
}
if (parsed.name === "UniformDistribution") {
return {
type: "FloatDistribution",
low: parsed.attributes.low as number,
high: parsed.attributes.high as number,
step: null,
log: false,
};
} else if (parsed.name === "LogUniformDistribution") {
}
}
if (parsed.name === "LogUniformDistribution") {
return {
type: "FloatDistribution",
low: parsed.attributes.low as number,
high: parsed.attributes.high as number,
step: null,
log: true,
};
} else if (parsed.name === "DiscreteUniformDistribution") {
}
}
if (parsed.name === "DiscreteUniformDistribution") {
return {
type: "FloatDistribution",
low: parsed.attributes.low as number,
high: parsed.attributes.high as number,
step: parsed.attributes.q,
log: false,
};
} else if (parsed.name === "IntDistribution") {
}
}
if (parsed.name === "IntDistribution") {
return {
type: "IntDistribution",
low: parsed.attributes.low as number,
high: parsed.attributes.high as number,
step: parsed.attributes.step as number,
log: parsed.attributes.log as boolean,
};
} else if (parsed.name === "IntUniformDistribution") {
}
}
if (parsed.name === "IntUniformDistribution") {
return {
type: "IntDistribution",
low: parsed.attributes.low as number,
high: parsed.attributes.high as number,
step: parsed.attributes.step as number,
log: false,
};
} else if (parsed.name === "IntLogUniformDistribution") {
}
}
if (parsed.name === "IntLogUniformDistribution") {
return {
type: "IntDistribution",
low: parsed.attributes.low as number,
high: parsed.attributes.high as number,
step: parsed.attributes.step as number,
log: true,
};
} else {
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
const choices = parsed.attributes.choices.map((value: any) => {
// TODO(c-bata): Support other types
return {
pytype: "str",
value: value.toString(),
};
});
return {
type: "CategoricalDistribution",
choices: choices,
};
}
}
};
return {
type: "CategoricalDistribution",
choices: parsed.attributes.choices,
}
}
const getTrialUserAttributes = (
db: SQLite3DB,
trialId: number,
trialId: number
): Attribute[] => {
const attrs: Attribute[] = [];
const attrs: Attribute[] = []
db.exec({
sql: `SELECT key, value_json FROM trial_user_attributes WHERE trial_id = ${trialId}`,
// biome-ignore lint/suspicious/noExplicitAny: <explanation>
@@ -349,18 +387,18 @@ const getTrialUserAttributes = (
attrs.push({
key: vals[0],
value: vals[1],
});
})
},
});
return attrs;
};
})
return attrs
}
const getTrialIntermediateValues = (
db: SQLite3DB,
trialId: number,
schemaVersion: string,
schemaVersion: string
): TrialIntermediateValue[] => {
const values: TrialIntermediateValue[] = [];
const values: TrialIntermediateValue[] = []
if (isGreaterSchemaVersion(schemaVersion, "v3.0.0.c")) {
db.exec({
sql: `SELECT step, intermediate_value, intermediate_value_type FROM trial_intermediate_values WHERE trial_id = ${trialId} ORDER BY step`,
@@ -370,15 +408,15 @@ const getTrialIntermediateValues = (
step: vals[0],
value:
vals[2] === "INF_NEG"
? "-inf"
? -Infinity
: vals[2] === "INF_POS"
? "+inf"
? Infinity
: vals[2] === "NAN"
? "nan"
? NaN
: vals[1],
});
})
},
});
})
} else {
db.exec({
sql: `SELECT step, intermediate_value FROM trial_intermediate_values WHERE trial_id = ${trialId} ORDER BY step`,
@@ -387,9 +425,9 @@ const getTrialIntermediateValues = (
values.push({
step: vals[0],
value: vals[1],
});
})
},
});
})
}
return values;
};
return values
}