diff --git a/tslib/storage/src/journal.ts b/tslib/storage/src/journal.ts index 79a746d1..05c0779d 100644 --- a/tslib/storage/src/journal.ts +++ b/tslib/storage/src/journal.ts @@ -139,6 +139,8 @@ class JournalStorage { const unionUserAttrs: Set = new Set() const unionSearchSpace: Set = new Set() let intersectionSearchSpace: string[] = [] + const nameToSearchSpaceItem: Map = + new Map() study.trials.forEach((trial, index) => { for (const userAttr of trial.user_attrs) { @@ -146,6 +148,12 @@ class JournalStorage { } for (const param of trial.params) { unionSearchSpace.add(param.name) + if (!nameToSearchSpaceItem.has(param.name)) { + nameToSearchSpaceItem.set(param.name, { + name: param.name, + distribution: param.distribution, + }) + } } if (index === 0) { intersectionSearchSpace = Array.from(unionSearchSpace) @@ -162,14 +170,10 @@ class JournalStorage { } }) study.union_search_space = Array.from(unionSearchSpace).map((name) => { - return { - name: name, - } + return nameToSearchSpaceItem.get(name) as Optuna.SearchSpaceItem }) study.intersection_search_space = intersectionSearchSpace.map((name) => { - return { - name: name, - } + return nameToSearchSpaceItem.get(name) as Optuna.SearchSpaceItem }) } diff --git a/tslib/storage/src/sqlite.ts b/tslib/storage/src/sqlite.ts index cb4be1ff..11912ac2 100644 --- a/tslib/storage/src/sqlite.ts +++ b/tslib/storage/src/sqlite.ts @@ -165,24 +165,35 @@ const getStudy = ( } const params = getTrialParams(db, trial.trial_id) - const param_names = new Set() + const paramNames = new Set() + const paramNameToSearchSpaceItem = new Map() for (const param of params) { - param_names.add(param.name) + 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 ) { - study.union_search_space.push({ name: param.name }) + study.union_search_space.push({ + name: param.name, + distribution: param.distribution, + }) } } if (intersection_search_space.size === 0) { - // biome-ignore lint/complexity/noForEach: - param_names.forEach((s) => { - intersection_search_space.add({ name: s }) - }) + for (const s of paramNames) { + intersection_search_space.add( + paramNameToSearchSpaceItem.get(s) as Optuna.SearchSpaceItem + ) + } } else { intersection_search_space = new Set( Array.from(intersection_search_space).filter((s) => - param_names.has(s.name) + paramNames.has(s.name) ) ) }