diff --git a/autofit/graphical/expectation_propagation/factor_optimiser.py b/autofit/graphical/expectation_propagation/factor_optimiser.py index 9b8d447a0..61613430f 100644 --- a/autofit/graphical/expectation_propagation/factor_optimiser.py +++ b/autofit/graphical/expectation_propagation/factor_optimiser.py @@ -16,11 +16,12 @@ class AbstractFactorOptimiser(ABC): """ logger = logger.debug - def __init__(self, initial_values=None, deltas=None, inplace=False, delta=1): + def __init__(self, initial_values=None, deltas=None, inplace=False, delta=1, dynamic_delta=False): self.initial_values = initial_values or {} self.inplace = inplace self.delta = delta self.deltas = deltas or {} + self.dynamic_delta = dynamic_delta def update_model_approx( self, @@ -30,14 +31,26 @@ def update_model_approx( status: Optional[Status] = Status(), delta: Optional[float] = None, ) -> Tuple[EPMeanField, Status]: - delta = delta or self.deltas.get(factor_approx.factor) or self.delta - new_approx, status = model_approx.project_mean_field( + + variable_message_count = model_approx.variable_message_count + min_value = min(variable_message_count.values()) + + delta = delta or self.delta + + if factor_approx.factor in self.deltas: + delta = self.deltas[factor_approx.factor] + elif self.dynamic_delta: + delta = MeanField({ + variable: self.delta * (min_value / message_count) + for variable, message_count in variable_message_count.items() + }) + + return model_approx.project_mean_field( new_model_dist, factor_approx, delta=delta, status=status, ) - return new_approx, status @abstractmethod def optimise( diff --git a/autofit/non_linear/abstract_search.py b/autofit/non_linear/abstract_search.py index cc9fb392c..c8d6a73a6 100644 --- a/autofit/non_linear/abstract_search.py +++ b/autofit/non_linear/abstract_search.py @@ -101,7 +101,7 @@ def __init__( session An SQLAlchemy session instance so the results of the model-fit are written to an SQLite database. """ - super().__init__(delta=kwargs.get("delta", 1.0)) + super().__init__(delta=kwargs.get("delta", 1.0), dynamic_delta=kwargs.get("dynamic_delta", True)) from autofit.non_linear.paths.database import DatabasePaths @@ -204,15 +204,13 @@ def __init__( self.optimisation_counter = Counter() - self.dynamic_delta = kwargs.get("dynamic_delta", True) - __identifier_fields__ = tuple() def optimise( self, factor: Factor, model_approx: EPMeanField, - status: Optional[Status] = None + status: Status = Status(), ) -> Tuple[EPMeanField, Status]: """ Perform optimisation for expectation propagation. Currently only @@ -289,25 +287,12 @@ def optimise( result.projected_model.priors ) - variable_message_count = model_approx.variable_message_count - min_value = min(variable_message_count.values()) - - if self.dynamic_delta: - delta = MeanField({ - variable: self.delta * (min_value / message_count) - for variable, message_count in variable_message_count.items() - }) - else: - delta = self.delta - - projection, status = factor_approx.project( - new_model_dist, - delta=delta + model_approx, status = self.update_model_approx( + new_model_dist, factor_approx, model_approx, status ) - status.result = result - return model_approx.project(projection, status) + return model_approx, status @property def name(self): diff --git a/test_autofit/graphical/regression/test_static.py b/test_autofit/graphical/regression/test_static.py index 6a0b66094..82f1a1a52 100644 --- a/test_autofit/graphical/regression/test_static.py +++ b/test_autofit/graphical/regression/test_static.py @@ -11,6 +11,7 @@ def __init__(self): self._paths = af.DirectoryPaths() self.delta = 1.0 self.dynamic_delta = False + self.deltas = {} def fit( self,