Skip to content

Commit 2adef52

Browse files
Jammy2211Jammy2211
authored andcommitted
spawn parallel
1 parent 057f25c commit 2adef52

3 files changed

Lines changed: 34 additions & 8 deletions

File tree

autofit/non_linear/fitness.py

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
import numpy as np_excplicit
12
import logging
23
import os
34
from typing import Optional
@@ -188,10 +189,17 @@ def __setstate__(self, state):
188189

189190
@cached_property
190191
def _call(self):
191-
logger.info("Compiling fitness function for JAX...")
192-
debug.print("aaa")
192+
debug.print(str(jax_wrapper.use_jax))
193+
debug.print("Compiling fitness function for JAX...")
193194
return jax_wrapper.jit(self.call)
194195

196+
def call_numpy_wrapper(self, parameters):
197+
198+
figure_of_merit = self.__call__(parameters=np_excplicit.array(parameters))
199+
200+
return figure_of_merit.item()
201+
202+
195203
@cached_property
196204
def _grad(self):
197205
return jax_wrapper.grad(self._call)

autofit/non_linear/search/abstract_search.py

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
import gc
44
import logging
55
import multiprocessing as mp
6+
import numpy as np
67
import os
78
import time
89
import warnings
@@ -31,7 +32,6 @@
3132
)
3233
from autofit.graphical.utils import Status
3334
from autofit.mapper.prior_model.abstract import AbstractPriorModel
34-
from autofit.mapper.prior_model.collection import Collection
3535
from autofit.mapper.model import ModelInstance
3636
from autofit.non_linear.initializer import Initializer
3737
from autofit.non_linear.fitness import Fitness
@@ -995,7 +995,11 @@ def perform_update(
995995
parameters = samples.max_log_likelihood(as_instance=False)
996996

997997
start = time.time()
998-
fitness(parameters)
998+
figure_of_merit = fitness(parameters)
999+
1000+
# account for asynchronous JAX calls
1001+
np.array(figure_of_merit)
1002+
9991003
log_likelihood_function_time = time.time() - start
10001004

10011005
if jax_wrapper.use_jax:

autofit/non_linear/search/nest/nautilus/search.py

Lines changed: 18 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -19,12 +19,22 @@
1919

2020
logger = logging.getLogger(__name__)
2121

22+
import time
2223

2324
def prior_transform(cube, model):
2425
return model.vector_from_unit_vector(unit_vector=cube)
2526

2627
def prior_transform_vectorized(cube, model):
27-
return np.array([model.vector_from_unit_vector(row) for row in cube])
28+
29+
start = time.time()
30+
31+
trans = np.array([model.vector_from_unit_vector(row) for row in cube])
32+
33+
end = time.time()
34+
35+
print(f"Time taken for prior transform vectorized: {end - start:.4f} seconds")
36+
37+
return trans
2838

2939
class Nautilus(abstract_nest.AbstractNest):
3040
__identifier_fields__ = (
@@ -141,7 +151,7 @@ def _fit(self, model: AbstractPriorModel, analysis):
141151
if (
142152
self.config_dict.get("force_x1_cpu")
143153
or self.kwargs.get("force_x1_cpu")
144-
or jax_wrapper.use_jax
154+
# or jax_wrapper.use_jax
145155
):
146156
search_internal = self.fit_x1_cpu(
147157
fitness=fitness,
@@ -227,7 +237,7 @@ def fit_x1_cpu(self, fitness, model, analysis):
227237
func = jax.vmap(fitness)
228238
prior_t = prior_transform_vectorized
229239
else:
230-
func = fitness.__call__
240+
func = fitness.call_numpy_wrapper
231241
prior_t = prior_transform
232242

233243
search_internal = self.sampler_cls(
@@ -264,12 +274,16 @@ def fit_multiprocessing(self, fitness, model, analysis):
264274
the log likelihood the search maximizes.
265275
"""
266276

277+
# from dask.distributed import Client
278+
# client = Client(processes=True, n_workers=4, threads_per_worker=1)
279+
267280
search_internal = self.sampler_cls(
268281
prior=prior_transform,
269-
likelihood=fitness.__call__,
282+
likelihood=fitness.call_numpy_wrapper,
270283
n_dim=model.prior_count,
271284
prior_kwargs={"model": model},
272285
filepath=self.checkpoint_file,
286+
# pool=client,
273287
pool=self.number_of_cores,
274288
**self.config_dict_search,
275289
)

0 commit comments

Comments
 (0)