diff --git a/apps/frontend/app/components/flows/AddExperiment/stepComponents/ParamStep.tsx b/apps/frontend/app/components/flows/AddExperiment/stepComponents/ParamStep.tsx index 3ed89ba5..6e94e361 100644 --- a/apps/frontend/app/components/flows/AddExperiment/stepComponents/ParamStep.tsx +++ b/apps/frontend/app/components/flows/AddExperiment/stepComponents/ParamStep.tsx @@ -5,87 +5,73 @@ import { InputSection } from '../../../InputSection'; //import { formList } from '@mantine/form'; import { HyperparametersCollection, HyperparameterTypes, IntegerHyperparameter } from '../../../../../lib/db_types'; import { useDebounce } from "use-debounce"; +import { final } from 'pino'; + +/** + * Helper to get the number of possible values for a single hyperparameter + */ +function getParamRange(param: any): number { + switch (param.type) { + case HyperparameterTypes.INTEGER: + return Math.floor((param.max - param.min) / param.step) + 1; + case HyperparameterTypes.FLOAT: + return Math.floor(((param.max - param.min) / param.step) + 1e-9) + 1; + case HyperparameterTypes.BOOLEAN: + return 2; + case HyperparameterTypes.STRING_LIST: + return param.values ? param.values.length : 0; + case HyperparameterTypes.STRING: + return 1; + case HyperparameterTypes.PARAM_GROUP: + const values = Object.values((param.values || {}) as any[][]); + return values.length > 0 ? values[0].length : 0; + default: + return 1; + } +} -function calcPermutations(parameters: HyperparametersCollection) { - var noDefaultCount = 1; - var defaultCount = 0; - - var countDefaults = 0; - var totalObjs = 0; - - var allInts = true; - - if (parameters.hyperparameters.length > 0) { - - parameters.hyperparameters.forEach(hyperparameter => { - totalObjs++; - if (hyperparameter.type == HyperparameterTypes.INTEGER || hyperparameter.type == HyperparameterTypes.FLOAT) { - - if (isNaN(hyperparameter.step) || hyperparameter.step == 0) { - hyperparameter.step = 1; - return -1; - } - - if (hyperparameter.type == HyperparameterTypes.FLOAT) - allInts = false; - - let hyper = hyperparameter; - let numObjs = 0; - - for (let i = hyper.min * 100; i <= hyper.max * 100; i += hyper.step * 100) { - numObjs++; - } - - - if (hyper.default == -1) { - noDefaultCount = noDefaultCount * numObjs; - } else { - defaultCount = defaultCount + numObjs; - countDefaults++; - } - - } - else if (hyperparameter.type == HyperparameterTypes.BOOLEAN) { - if (!hyperparameter.useDefault) { - noDefaultCount = noDefaultCount * 2; - } - else { - defaultCount = defaultCount + 1; - countDefaults++; - } - } - else if (hyperparameter.type == HyperparameterTypes.STRING_LIST) { - if (hyperparameter.default == '-1') { - noDefaultCount = noDefaultCount * hyperparameter.values.length; - } - else { - defaultCount = defaultCount + hyperparameter.values.length; - countDefaults++; - } - } - else if (hyperparameter.type == HyperparameterTypes.PARAM_GROUP) { - let hyper = hyperparameter; - let numObjs = 0; - for (let key in hyper.values) { - numObjs = hyper.values[key].length; - break; - } - - noDefaultCount = noDefaultCount * numObjs; - - } - }); - - if (totalObjs < 3 && allInts && countDefaults > 0) { - const total = (noDefaultCount + defaultCount) - 1; - return total; - } - else { - const total = (noDefaultCount * defaultCount) - (noDefaultCount * (countDefaults - 1)); - return total; - } - } - +function calcPermutations(parameters: HyperparametersCollection): number { + const params = parameters.hyperparameters; + + const paramGroups = params.filter(p => p.type === HyperparameterTypes.PARAM_GROUP); + const normalParams = params.filter(p => p.type !== HyperparameterTypes.PARAM_GROUP); + + const hasValidDefault = (p: any): boolean => { + const def = p.default; + return p.useDefault || (def !== -1 && def !== "-1" && def !== '' && def !== undefined && def !== null); + }; + + const D = normalParams.filter(hasValidDefault); + const F = normalParams.filter(p => !hasValidDefault(p)); + + // 1. Calculate Product of Free Parameters + let freeProduct = 1; + for (const p of F) { + freeProduct *= getParamRange(p); + } + + // 2. Calculate Combinations of Constrained Parameters (One-at-a-time deviation) + // Formula: 1 (all at default) + Sum of (range - 1) for each param + let constrainedCombinations = 1; + for (const p of D) { + const range = getParamRange(p); + if (range > 1) { + constrainedCombinations += (range - 1); + } + } + + let total = freeProduct * constrainedCombinations; + + // 3. Handle Param Groups (Sum of group lengths multiplied by current total) + if (paramGroups.length > 0) { + let groupTotal = 0; + for (const pg of paramGroups) { + groupTotal += getParamRange(pg); + } + total *= groupTotal; + } + + return total; } export const ParameterOptions = ['integer', 'float', 'bool', 'stringlist', 'paramgroup'] as const; diff --git a/apps/frontend/lib/db_types.ts b/apps/frontend/lib/db_types.ts index 2dfa37c5..eaf0eec0 100644 --- a/apps/frontend/lib/db_types.ts +++ b/apps/frontend/lib/db_types.ts @@ -20,6 +20,7 @@ export enum HyperparameterTypes { export interface GenericHyperparameter { name: string; type: HyperparameterTypes; + useDefault: boolean; } export interface ParamGroupHyperparameter extends GenericHyperparameter { diff --git a/apps/runner/modules/configs.py b/apps/runner/modules/configs.py index a8ab5796..74fa5fb3 100644 --- a/apps/runner/modules/configs.py +++ b/apps/runner/modules/configs.py @@ -70,7 +70,7 @@ def generate_permutations(parameters, paramgroup=None): explogger.info("paramgroup vals: %s", str(paramgroup)) for param in parameters: - if param["default"] != -1 and param["default"] != "-1" and param["default"] != '': + if param["useDefault"] or (param["default"] != -1 and param["default"] != "-1" and param["default"] != ''): default_vals[param["name"]] = [param["default"]] else: default_vals[param["name"]] = expand_values(param)