Skip to content

Commit 6913693

Browse files
authored
Merge pull request #1153 from rhayes777/feature/declarative_deterministic
feature/declarative deterministic
2 parents 40f5b25 + a4a525e commit 6913693

5 files changed

Lines changed: 88 additions & 10 deletions

File tree

.gitignore

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
graph.html
12
phase_output_path.zip
23
*.pickle
34
example/

autofit/mapper/prior/arithmetic/compound.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -84,6 +84,9 @@ def __init__(self, left, right):
8484
self.left = left
8585
self.right = right
8686

87+
def __repr__(self):
88+
return str(self)
89+
8790
def dict(self) -> dict:
8891
from autofit import ModelObject
8992

autofit/mapper/prior_model/prior_model.py

Lines changed: 14 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -406,12 +406,11 @@ def __setattr__(self, key, value):
406406
logger.exception(key)
407407

408408
def __getattr__(self, item):
409+
if item in ("_is_frozen", "tuple_prior_tuples"):
410+
return self.__getattribute__(item)
411+
409412
try:
410-
if (
411-
"_" in item
412-
and item not in ("_is_frozen", "tuple_prior_tuples")
413-
and not item.startswith("_")
414-
):
413+
if "_" in item and not item.startswith("_"):
415414
return getattr(
416415
[v for k, v in self.tuple_prior_tuples if item.split("_")[0] == k][
417416
0
@@ -422,6 +421,16 @@ def __getattr__(self, item):
422421
except IndexError:
423422
pass
424423

424+
try:
425+
return getattr(
426+
self.instance_for_arguments(
427+
{prior: prior for prior in self.priors},
428+
),
429+
item,
430+
)
431+
except (AttributeError, TypeError):
432+
pass
433+
425434
self.__getattribute__(item)
426435

427436
@property

test_autofit/graphical/gaussian/test_optimizer.py

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -34,8 +34,9 @@ def test_default(factor_model, laplace):
3434
assert model.normalization.mean == pytest.approx(25, rel=0.1)
3535
assert model.sigma.mean == pytest.approx(10, rel=0.1)
3636

37-
@pytest.mark.filterwarnings('ignore::RuntimeWarning')
38-
def test_set_model_identifier(dynesty, prior_model, analysis):
37+
38+
@pytest.mark.filterwarnings("ignore::RuntimeWarning")
39+
def _test_set_model_identifier(dynesty, prior_model, analysis):
3940
dynesty.fit(prior_model, analysis)
4041

4142
identifier = dynesty.paths.identifier
@@ -48,13 +49,13 @@ def test_set_model_identifier(dynesty, prior_model, analysis):
4849

4950

5051
class TestDynesty:
51-
@pytest.mark.filterwarnings('ignore::RuntimeWarning')
52+
@pytest.mark.filterwarnings("ignore::RuntimeWarning")
5253
@output_path_for_test()
5354
def test_optimisation(self, factor_model, laplace, dynesty):
5455
factor_model.optimiser = dynesty
5556
factor_model.optimise(laplace)
5657

57-
@pytest.mark.filterwarnings('ignore::RuntimeWarning')
58+
@pytest.mark.filterwarnings("ignore::RuntimeWarning")
5859
def test_null_paths(self, factor_model):
5960
search = af.DynestyStatic(maxcall=10)
6061
result, status = search.optimise(
@@ -64,7 +65,7 @@ def test_null_paths(self, factor_model):
6465
assert isinstance(result, g.MeanField)
6566
assert isinstance(status, Status)
6667

67-
@pytest.mark.filterwarnings('ignore::RuntimeWarning')
68+
@pytest.mark.filterwarnings("ignore::RuntimeWarning")
6869
@output_path_for_test()
6970
def test_optimise(self, factor_model, dynesty):
7071
result, status = dynesty.optimise(
Lines changed: 64 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,64 @@
1+
import autofit as af
2+
from autofit.mock import MockAnalysis
3+
4+
5+
def test():
6+
model_1 = af.Model(af.Gaussian)
7+
analysis_factor_1 = af.AnalysisFactor(
8+
prior_model=model_1,
9+
analysis=MockAnalysis(),
10+
)
11+
12+
model_2 = af.Model(af.Gaussian)
13+
analysis_factor_2 = af.AnalysisFactor(
14+
prior_model=model_2,
15+
analysis=MockAnalysis(),
16+
)
17+
18+
model_3 = af.Collection(
19+
model_1.fwhm,
20+
model_2.fwhm,
21+
)
22+
analysis_factor_3 = af.AnalysisFactor(
23+
prior_model=model_3,
24+
analysis=MockAnalysis(),
25+
)
26+
27+
factor_graph = af.FactorGraphModel(
28+
analysis_factor_1,
29+
analysis_factor_2,
30+
analysis_factor_3,
31+
)
32+
33+
assert (
34+
factor_graph.info
35+
== """PriorFactors
36+
37+
PriorFactor0 (AnalysisFactor1.sigma, AnalysisFactor2.1.self) UniformPrior [5], lower_limit = 0.0, upper_limit = 1.0
38+
PriorFactor1 (AnalysisFactor1.normalization) UniformPrior [4], lower_limit = 0.0, upper_limit = 1.0
39+
PriorFactor2 (AnalysisFactor1.centre) UniformPrior [3], lower_limit = 0.0, upper_limit = 1.0
40+
PriorFactor3 (AnalysisFactor0.sigma, AnalysisFactor2.0.self) UniformPrior [2], lower_limit = 0.0, upper_limit = 1.0
41+
PriorFactor4 (AnalysisFactor0.normalization) UniformPrior [1], lower_limit = 0.0, upper_limit = 1.0
42+
PriorFactor5 (AnalysisFactor0.centre) UniformPrior [0], lower_limit = 0.0, upper_limit = 1.0
43+
44+
AnalysisFactors
45+
46+
AnalysisFactor0
47+
48+
centre (PriorFactor5) UniformPrior [0], lower_limit = 0.0, upper_limit = 1.0
49+
normalization (PriorFactor4) UniformPrior [1], lower_limit = 0.0, upper_limit = 1.0
50+
sigma (AnalysisFactor2.0.self, PriorFactor3) UniformPrior [2], lower_limit = 0.0, upper_limit = 1.0
51+
52+
AnalysisFactor1
53+
54+
centre (PriorFactor2) UniformPrior [3], lower_limit = 0.0, upper_limit = 1.0
55+
normalization (PriorFactor1) UniformPrior [4], lower_limit = 0.0, upper_limit = 1.0
56+
sigma (AnalysisFactor2.1.self, PriorFactor0) UniformPrior [5], lower_limit = 0.0, upper_limit = 1.0
57+
58+
AnalysisFactor2
59+
60+
0
61+
self (AnalysisFactor0.sigma, PriorFactor3) UniformPrior [2], lower_limit = 0.0, upper_limit = 1.0
62+
1
63+
self (AnalysisFactor1.sigma, PriorFactor0) UniformPrior [5], lower_limit = 0.0, upper_limit = 1.0"""
64+
)

0 commit comments

Comments
 (0)