Skip to content

Commit bee332c

Browse files
committed
♻️ Refactor benchmark runner to use proper config objects
- Replace intermediate dict with ExperimentConfig dataclass holding DiffusionConfig - Generate DiffusionConfig objects directly in generate_experiments() - Extract config values from config object instead of dict unpacking - Update BenchmarkMetrics.seed type to int | None to match DiffusionConfig
1 parent d380d66 commit bee332c

2 files changed

Lines changed: 40 additions & 44 deletions

File tree

‎benchmark/metrics.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -24,7 +24,7 @@ class BenchmarkMetrics:
2424
num_steps: int
2525
param_dim: int
2626
sigma_m: float
27-
seed: int
27+
seed: int | None
2828

2929
# Performance metrics
3030
best_fitness: float

‎benchmark/runner.py‎

Lines changed: 39 additions & 43 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,9 @@
11
"""Grid search runner with multiprocessing support."""
22

33
import time
4+
from dataclasses import dataclass
45
from itertools import product
56
from multiprocessing import Pool, cpu_count
6-
from typing import Any
77

88
import numpy as np
99
from tqdm import tqdm
@@ -15,51 +15,45 @@
1515
evaluate_population_fitness,
1616
)
1717
from devol import DiffusionConfig, DiffusionEvolution
18+
from devol.config import ScheduleConfig, ScheduleType
1819

1920

20-
def run_single_experiment(params: dict[str, Any]) -> BenchmarkMetrics:
21-
"""Run a single experiment with given parameters.
21+
@dataclass
22+
class ExperimentConfig:
23+
"""Configuration for a single benchmark experiment."""
24+
25+
config: DiffusionConfig
26+
fitness_fn: FitnessFunction
27+
28+
29+
def run_single_experiment(experiment: ExperimentConfig) -> BenchmarkMetrics:
30+
"""Run a single experiment with given configuration.
2231
2332
Args:
24-
params: Dictionary with experiment parameters
33+
experiment: ExperimentConfig with DiffusionConfig and fitness function
2534
2635
Returns:
2736
BenchmarkMetrics object with results
2837
"""
29-
schedule_type = params["schedule_type"]
30-
population_size = params["population_size"]
31-
num_steps = params["num_steps"]
32-
param_dim = params["param_dim"]
33-
sigma_m = params["sigma_m"]
34-
seed = params["seed"]
35-
fitness_fn = params["fitness_fn"]
36-
37-
config = DiffusionConfig(
38-
population_size=population_size,
39-
num_steps=num_steps,
40-
param_dim=param_dim,
41-
sigma_m=sigma_m,
42-
schedule={"type": schedule_type},
43-
seed=seed,
44-
)
38+
config = experiment.config
39+
fitness_fn = experiment.fitness_fn
4540

4641
start_time = time.time()
4742
algo = DiffusionEvolution(config, fitness_fn)
48-
final_pop = algo.run()
43+
final_pop = algo.run(initial_population=None)
4944
runtime = time.time() - start_time
5045

5146
best_individual, best_fitness = algo.get_best_individual()
5247
fitness_values = evaluate_population_fitness(final_pop, fitness_fn)
53-
5448
final_diversity = calculate_population_diversity(fitness_values)
5549

5650
return BenchmarkMetrics(
57-
schedule_type=schedule_type,
58-
population_size=population_size,
59-
num_steps=num_steps,
60-
param_dim=param_dim,
61-
sigma_m=sigma_m,
62-
seed=seed,
51+
schedule_type=config.schedule.type.value,
52+
population_size=config.population_size,
53+
num_steps=config.num_steps,
54+
param_dim=config.param_dim,
55+
sigma_m=config.sigma_m,
56+
seed=config.seed,
6357
best_fitness=float(best_fitness),
6458
mean_fitness=float(np.mean(fitness_values)),
6559
std_fitness=float(np.std(fitness_values)),
@@ -75,7 +69,7 @@ class GridSearchRunner:
7569
def __init__(
7670
self,
7771
fitness_fn: FitnessFunction,
78-
schedule_types: list[str],
72+
schedule_types: list[ScheduleType | str],
7973
population_sizes: list[int],
8074
num_steps_list: list[int],
8175
param_dims: list[int],
@@ -96,35 +90,37 @@ def __init__(
9690
n_workers: Number of parallel workers (default: CPU count)
9791
"""
9892
self.fitness_fn = fitness_fn
99-
self.schedule_types = schedule_types
93+
self.schedule_types = [
94+
ScheduleType(s) if isinstance(s, str) else s for s in schedule_types
95+
]
10096
self.population_sizes = population_sizes
10197
self.num_steps_list = num_steps_list
10298
self.param_dims = param_dims
10399
self.sigma_m_values = sigma_m_values
104100
self.seeds = seeds
105101
self.n_workers = n_workers or cpu_count()
106102

107-
def generate_experiments(self) -> list[dict[str, Any]]:
103+
def generate_experiments(self) -> list[ExperimentConfig]:
108104
"""Generate all experiment configurations."""
109-
configs = []
110-
for schedule, pop_size, steps, dim, sigma, seed in product(
105+
experiments = []
106+
for schedule_type, pop_size, steps, dim, sigma, seed in product(
111107
self.schedule_types,
112108
self.population_sizes,
113109
self.num_steps_list,
114110
self.param_dims,
115111
self.sigma_m_values,
116112
self.seeds,
117113
):
118-
configs.append({
119-
"schedule_type": schedule,
120-
"population_size": pop_size,
121-
"num_steps": steps,
122-
"param_dim": dim,
123-
"sigma_m": sigma,
124-
"seed": seed,
125-
"fitness_fn": self.fitness_fn,
126-
})
127-
return configs
114+
config = DiffusionConfig(
115+
population_size=pop_size,
116+
num_steps=steps,
117+
param_dim=dim,
118+
sigma_m=sigma,
119+
schedule=ScheduleConfig(type=schedule_type),
120+
seed=seed,
121+
)
122+
experiments.append(ExperimentConfig(config=config, fitness_fn=self.fitness_fn))
123+
return experiments
128124

129125
def run(self, verbose: bool = True) -> list[BenchmarkMetrics]:
130126
"""Run all experiments in parallel.

0 commit comments

Comments
 (0)