Skip to content

Commit fb9d8b9

Browse files
Jammy2211Jammy2211
authored andcommitted
fix issues with _use_jax now being private
1 parent 2cb4a04 commit fb9d8b9

5 files changed

Lines changed: 5 additions & 12 deletions

File tree

‎autofit/non_linear/analysis/analysis.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -109,7 +109,7 @@ def compute_latent_samples(self, samples: Samples, batch_size : Optional[int] =
109109

110110
compute_latent_for_model = functools.partial(self.compute_latent_variables, model=samples.model)
111111

112-
if self.use_jax:
112+
if self._use_jax:
113113
import jax
114114
start = time.time()
115115
logger.info("JAX: Applying vmap and jit to likelihood function for latent variables -- may take a few seconds.")
@@ -130,7 +130,7 @@ def batched_compute_latent(x):
130130
# batched JAX call on this chunk
131131
latent_values_batch = batched_compute_latent(batch)
132132

133-
if self.use_jax:
133+
if self._use_jax:
134134
import jax.numpy as jnp
135135
latent_values_batch = jnp.stack(latent_values_batch, axis=-1) # (batch, n_latents)
136136
mask = jnp.all(jnp.isfinite(latent_values_batch), axis=0)

‎autofit/non_linear/fitness.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -134,7 +134,7 @@ def __init__(
134134

135135
@property
136136
def _xp(self):
137-
if self.analysis.use_jax:
137+
if self.analysis._use_jax:
138138
import jax.numpy as jnp
139139
return jnp
140140
return np

‎autofit/non_linear/search/nest/dynesty/search/abstract.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -146,7 +146,7 @@ def _fit(
146146
"parallel"
147147
].get("force_x1_cpu")
148148
or self.kwargs.get("force_x1_cpu")
149-
or analysis.use_jax
149+
or analysis._use_jax
150150
):
151151
raise RuntimeError
152152

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

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -127,7 +127,7 @@ def _fit(self, model: AbstractPriorModel, analysis):
127127
if (
128128
self.config_dict.get("force_x1_cpu")
129129
or self.kwargs.get("force_x1_cpu")
130-
or analysis.use_jax
130+
or analysis._use_jax
131131
):
132132

133133
fitness = Fitness(

‎test_autofit/graphical/hierarchical/test_optimise.py‎

Lines changed: 0 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -11,13 +11,6 @@ def make_factor(hierarchical_factor):
1111
def test_optimise(factor):
1212
search = af.DynestyStatic(maxcall=100, dynamic_delta=False, delta=0.1,)
1313

14-
print(type(factor.analysis))
15-
print(type(factor.analysis))
16-
print(type(factor.analysis))
17-
print(type(factor.analysis))
18-
print(type(factor.analysis))
19-
20-
2114
_, status = search.optimise(
2215
factor.mean_field_approximation().factor_approximation(factor)
2316
)

0 commit comments

Comments
 (0)