From 6ed7b5ef0e4ad04618bfbec70efc17f0e1ce1745 Mon Sep 17 00:00:00 2001 From: porink0424 Date: Fri, 10 May 2024 11:49:43 +0900 Subject: [PATCH] Support metric_names for tslib/storage --- tslib/storage/src/journal.ts | 19 ++++++++++++++++++- tslib/storage/src/sqlite.ts | 22 ++++++++++++++++++++++ 2 files changed, 40 insertions(+), 1 deletion(-) diff --git a/tslib/storage/src/journal.ts b/tslib/storage/src/journal.ts index 29e59964..e09536ef 100644 --- a/tslib/storage/src/journal.ts +++ b/tslib/storage/src/journal.ts @@ -29,6 +29,13 @@ interface JournalOpDeleteStudy extends JournalOpBase { study_id: number } +interface JournalOpSetStudySystemAttr extends JournalOpBase { + study_id: number + system_attr: { + "study:metric_names": string[] + } +} + interface JournalOpCreateTrial extends JournalOpBase { study_id: number datetime_start?: string @@ -179,6 +186,14 @@ class JournalStorage { this.studies = this.studies.filter((item) => item.id !== log.study_id) } + public applyStudySystemAttr(log: JournalOpSetStudySystemAttr): void { + const thisStudy = this.studies.find((item) => item.id === log.study_id) + if (thisStudy === undefined) { + return + } + thisStudy.metric_names = log.system_attr["study:metric_names"] + } + public applyCreateTrial(log: JournalOpCreateTrial): void { const thisStudy = this.studies.find((item) => item.id === log.study_id) if (thisStudy === undefined) { @@ -393,7 +408,9 @@ const loadJournalStorage = ( // Unsupported break case JournalOperation.SET_STUDY_SYSTEM_ATTR: - // Unsupported + journalStorage.applyStudySystemAttr( + parsedLog as JournalOpSetStudySystemAttr + ) break case JournalOperation.CREATE_TRIAL: journalStorage.applyCreateTrial(parsedLog as JournalOpCreateTrial) diff --git a/tslib/storage/src/sqlite.ts b/tslib/storage/src/sqlite.ts index 4756060f..543e7dba 100644 --- a/tslib/storage/src/sqlite.ts +++ b/tslib/storage/src/sqlite.ts @@ -144,6 +144,11 @@ const getStudy = ( trials: [], } + const studySystemAttrs = getStudySystemAttributes(db, summary.id) + if (studySystemAttrs !== undefined) { + study.metric_names = studySystemAttrs.metric_names + } + let intersection_search_space: Set = new Set() study.trials = getTrials(db, summary.id, schemaVersion) for (const trial of study.trials) { @@ -367,6 +372,23 @@ const parseDistributionJSON = (t: string): Optuna.Distribution => { } } +const getStudySystemAttributes = ( + db: SQLite3DB, + studyId: number +) => { + let attrs: { metric_names: string[] } | undefined + db.exec({ + sql: `SELECT key, value_json FROM study_system_attributes WHERE study_id = ${studyId} AND key = 'dashboard:objective_names'`, + // biome-ignore lint/suspicious/noExplicitAny: + callback: (vals: any[]) => { + attrs = { + metric_names: JSON.parse(vals[1]), + } + }, + }) + return attrs +} + const getTrialUserAttributes = ( db: SQLite3DB, trialId: number