11"""Grid search runner with multiprocessing support."""
22
33import time
4+ from dataclasses import dataclass
45from itertools import product
56from multiprocessing import Pool , cpu_count
6- from typing import Any
77
88import numpy as np
99from tqdm import tqdm
1515 evaluate_population_fitness ,
1616)
1717from 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