Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 23 additions & 11 deletions autofit/mapper/prior_model/abstract.py
Original file line number Diff line number Diff line change
Expand Up @@ -1005,24 +1005,32 @@ def instance_from_vector(self, vector, ignore_assertions: bool = False, xp=np):
model_instance : autofit.mapper.model.ModelInstance
An object containing reconstructed model_mapper instances
"""
if len(vector) != self.prior_count:
priors = self._vector_priors()
if len(vector) != len(priors):
raise AssertionError(
f"Vector length {len(vector)} != prior count {self.prior_count}"
)
arguments = dict(
map(
lambda prior_tuple, physical_unit: (prior_tuple.prior, physical_unit),
self.prior_tuples_ordered_by_id,
vector,
f"Vector length {len(vector)} != prior count {len(priors)}"
)
)
arguments = dict(zip(priors, vector))

return self.instance_for_arguments(
arguments,
ignore_assertions=ignore_assertions,
xp=xp
)

@frozen_cache
def _vector_priors(self) -> tuple:
"""
The unique priors of this model in the canonical (id-sorted) parameter
order that ``instance_from_vector`` pairs with the entries of a vector.

Cached while the model is frozen (a search freezes its model for the whole
fit), so the per-call rebuild of the ``prior_tuples_ordered_by_id``
name/value wrappers is skipped on the likelihood hot path. Unfrozen models
recompute it on every call.
"""
return tuple(prior_tuple.prior for prior_tuple in self.prior_tuples_ordered_by_id)

def constrained_model_tuples(self):
"""
Every component in this model whose class declares a model constraint.
Expand Down Expand Up @@ -1835,8 +1843,12 @@ def instance_for_arguments(
-------
An instance of the class
"""
if not (
conf.instance["general"]["test"]["exception_override"] or ignore_assertions
# A model with no assertions has nothing to check, so the config lookup
# (a hot-path cost on every likelihood call) is skipped for it.
if (
getattr(self, "_assertions", None)
and not ignore_assertions
and not conf.instance["general"]["test"]["exception_override"]
):
self.check_assertions(arguments, xp=xp)

Expand Down
132 changes: 131 additions & 1 deletion autofit/mapper/prior_model/prior_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

from autonerves.class_path import get_class_path
from autonerves.exc import ConfigException
from autofit.mapper.model import ModelInstance, assert_not_frozen
from autofit.mapper.model import ModelInstance, assert_not_frozen, frozen_cache
from autofit.mapper.model_object import ModelObject
from autofit.mapper.prior.abstract import Prior
from autofit.mapper.prior.constant import Constant
Expand Down Expand Up @@ -498,6 +498,11 @@ def _instance_for_arguments(
-------
An instance of the class
"""
if getattr(self, "_is_frozen", False):
return self._instance_for_arguments_frozen(
arguments, ignore_assertions=ignore_assertions, xp=xp
)

model_arguments = dict()
attribute_arguments = {
key: value
Expand Down Expand Up @@ -584,6 +589,131 @@ def _instance_for_arguments(

return result

@frozen_cache
def _instance_plan(self):
"""
The parts of ``_instance_for_arguments`` that depend only on the model's
structure, never on the argument values: which attributes are constructor
arguments, which are tuple priors / child models / priors, whether the
model is deferred, how the class is constructed, and which attributes are
candidates for being set on the instance after construction.

Only ever built for a frozen model (whose ``__dict__`` cannot change) and
cached until ``unfreeze()``; see ``_instance_for_arguments_frozen``.
"""
cls = self.cls
is_class = inspect.isclass(cls)
excluded = type(self)._cached_property_names()
constructor_argument_names = self.constructor_argument_names

attribute_arguments = {
key: value
for key, value in self.__dict__.items()
if key in constructor_argument_names
}
tuple_priors = tuple(
(name, tuple_prior)
for name, tuple_prior in self.direct_tuples_with_type(TuplePrior)
)
child_models = tuple(
(name, prior_model)
for name, prior_model in self.direct_tuples_with_type(AbstractPriorModel)
)
priors = tuple(
(name, prior) for name, prior in self.direct_tuples_with_type(Prior)
)
# ``(key, value)`` pairs that pass every value-independent condition of the
# post-construction loop; ``hasattr(result, key)`` is checked per call.
post_construction = tuple(
(key, value)
for key, value in self.__dict__.items()
if not isinstance(value, Prior)
and not key == "cls"
and not key.startswith("_")
and key not in excluded
)
return (
attribute_arguments,
tuple_priors,
child_models,
priors,
self.is_deferred_arguments,
is_class and issubclass(cls, Prior),
None if is_class else inspect._findclass(cls),
post_construction,
)

def _instance_for_arguments_frozen(
self,
arguments: {ModelObject: object},
ignore_assertions=False,
xp=np,
):
"""
``_instance_for_arguments`` for a frozen model: identical behaviour, with
the value-independent structure read from the cached ``_instance_plan``
instead of being rediscovered from ``__dict__`` on every call.
"""
(
attribute_arguments,
tuple_priors,
child_models,
priors,
is_deferred,
is_prior_class,
found_class,
post_construction,
) = self._instance_plan()

constructor_arguments = dict(attribute_arguments)

for name, tuple_prior in tuple_priors:
constructor_arguments[name] = tuple_prior.value_for_arguments(arguments)
for name, prior_model in child_models:
constructor_arguments[name] = prior_model.instance_for_arguments(
arguments, ignore_assertions=ignore_assertions, xp=xp
)
for name, prior in priors:
try:
constructor_arguments[name] = arguments[prior]
except KeyError as e:
raise KeyError(f"No argument given for prior {name}") from e

constructor_arguments = {
key: value.value if isinstance(value, Constant) else value
for key, value in constructor_arguments.items()
}

if is_deferred:
return DeferredInstance(self.cls, constructor_arguments)

if is_prior_class and any(
isinstance(value, tuple) for value in constructor_arguments.values()
):
# See ``_instance_for_arguments``: a bounds pair on a Prior model.
return ModelInstance(constructor_arguments)

if found_class is not None:
result = object.__new__(found_class)
self.cls(result, **constructor_arguments)
else:
result = self.cls(**constructor_arguments)

for key, value in post_construction:
if not hasattr(result, key):
if isinstance(value, Model):
value = value.instance_for_arguments(
arguments, ignore_assertions=ignore_assertions, xp=xp
)
elif isinstance(value, Constant):
value = value.value
try:
setattr(result, key, value)
except AttributeError:
pass

return result

def gaussian_prior_model_for_arguments(self, arguments):
"""
Returns a new instance of model mapper with a set of Gaussian priors based on \
Expand Down
165 changes: 165 additions & 0 deletions test_autofit/mapper/test_instance_from_vector_frozen.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,165 @@
"""
A frozen model builds instances through cached structure (``_vector_priors`` and
``Model._instance_plan``); an unfrozen model rediscovers that structure on every call.
Both must produce the same instance, attribute by attribute and type by type, and raise
the same exception when an assertion fails.
"""
import copy

import numpy as np
import pytest

import autofit as af
from autofit import exc
from autofit.example.model import PhysicalNFW


def _gaussian():
return af.Model(af.ex.Gaussian)


def _nested_collection():
return af.Collection(
a=af.Model(af.ex.Gaussian),
inner=af.Collection(
b=af.Model(af.ex.Exponential),
c=af.Model(af.ex.Gaussian),
),
)


def _shared_priors():
g0 = af.Model(af.ex.Gaussian)
g1 = af.Model(af.ex.Gaussian)
g1.centre = g0.centre
g1.sigma = g0.sigma
return af.Collection(g0=g0, g1=g1)


def _tuple_priors():
return af.Model(PhysicalNFW)


def _deferred():
model = af.Model(af.ex.Gaussian)
model.sigma = af.DeferredArgument()
return af.Collection(g=model, other=af.Model(af.ex.Exponential))


def _constants():
model = af.Model(af.ex.Gaussian)
model.normalization = af.Constant(2.0)
model.sigma = 3.0
return af.Collection(a=1.0, b=af.Constant(2.0), g=model)


def _child_model_attribute():
# Attributes that are not constructor arguments are set on the instance after
# construction by the post-construction loop.
model = af.Model(af.ex.Gaussian)
model.extra_constant = af.Constant(4.0)
model.extra_value = "a string"
return model


def _assertions():
model = af.Collection(
g0=af.Model(af.ex.Gaussian), g1=af.Model(af.ex.Gaussian)
)
model.add_assertion(model.g0.centre < model.g1.centre)
model.g1.add_assertion(model.g1.sigma > model.g1.normalization)
return model


MODELS = {
"gaussian": _gaussian,
"nested_collection": _nested_collection,
"shared_priors": _shared_priors,
"tuple_priors": _tuple_priors,
"deferred": _deferred,
"constants": _constants,
"non_constructor_attributes": _child_model_attribute,
"assertions": _assertions,
}


def assert_equal_instances(a, b, path="instance"):
assert type(a) is type(b), f"{path}: {type(a)} != {type(b)}"
if isinstance(a, (list, tuple)):
assert len(a) == len(b), path
for i, (x, y) in enumerate(zip(a, b)):
assert_equal_instances(x, y, f"{path}[{i}]")
elif isinstance(a, dict):
assert list(a) == list(b), f"{path}: keys {list(a)} != {list(b)}"
for key in a:
assert_equal_instances(a[key], b[key], f"{path}[{key!r}]")
elif isinstance(a, (int, float, str, bool, type(None), np.ndarray, np.generic)):
assert np.array_equal(a, b), f"{path}: {a} != {b}"
elif isinstance(a, type):
assert a is b, path
else:
assert list(vars(a)) == list(vars(b)), (
f"{path}: attributes {list(vars(a))} != {list(vars(b))}"
)
for key in vars(a):
assert_equal_instances(
getattr(a, key), getattr(b, key), f"{path}.{key}"
)


def _build(model, vector):
try:
return model.instance_from_vector(vector), None
except exc.FitException as e:
return None, str(e)


@pytest.mark.parametrize("name", list(MODELS))
def test_frozen_matches_unfrozen(name):
unfrozen = MODELS[name]()
frozen = copy.deepcopy(unfrozen)
frozen.freeze()

np.random.seed(1)
raised = 0
for _ in range(50):
vector = unfrozen.random_vector_from_priors
expected, expected_error = _build(unfrozen, vector)
# Twice, so the second frozen call runs off the populated cache.
for _ in range(2):
result, error = _build(frozen, vector)
assert error == expected_error
if expected_error is None:
assert_equal_instances(result, expected)
raised += expected_error is not None

if name == "assertions":
assert 0 < raised < 50


def test_frozen_cache_populated_and_reset_on_unfreeze():
model = _nested_collection()
model.freeze()
model.instance_from_vector(model.random_vector_from_priors)

assert any(key[0] == "_vector_priors" for key in model._frozen_cache)
child = model.a
assert any(key[0] == "_instance_plan" for key in child._frozen_cache)

model.unfreeze()
assert model._frozen_cache == {}
assert child._frozen_cache == {}

# An edit after unfreezing is seen by the next frozen build.
child.sigma = af.Constant(7.0)
model.freeze()
instance = model.instance_from_vector(model.random_vector_from_priors)
assert instance.a.sigma == 7.0
assert model.prior_count == len(model._vector_priors())


def test_wrong_vector_length_raises():
model = _gaussian()
model.freeze()
with pytest.raises(AssertionError):
model.instance_from_vector([1.0, 2.0])
Loading