Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
1 change: 1 addition & 0 deletions apps/frontend/lib/db_types.ts
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ export enum HyperparameterTypes {
export interface GenericHyperparameter {
name: string;
type: HyperparameterTypes;
useDefault: boolean;
}

export interface ParamGroupHyperparameter extends GenericHyperparameter {
Expand Down
2 changes: 1 addition & 1 deletion apps/runner/modules/configs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading