From 2e00ce70c559c68e4e653eac2c8709bbb3a59191 Mon Sep 17 00:00:00 2001 From: Richard Date: Mon, 11 Jul 2022 16:23:47 +0100 Subject: [PATCH 1/4] attempt to unify laplace and other optimiser post optimisation projection logic --- .../factor_optimiser.py | 18 +++++++++++++-- autofit/non_linear/abstract_search.py | 23 ++++--------------- 2 files changed, 20 insertions(+), 21 deletions(-) diff --git a/autofit/graphical/expectation_propagation/factor_optimiser.py b/autofit/graphical/expectation_propagation/factor_optimiser.py index 9b8d447a0..b9ec18980 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,7 +31,20 @@ 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 + + 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() + }) + new_approx, status = model_approx.project_mean_field( new_model_dist, factor_approx, diff --git a/autofit/non_linear/abstract_search.py b/autofit/non_linear/abstract_search.py index cc9fb392c..86d845416 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,8 +204,6 @@ def __init__( self.optimisation_counter = Counter() - self.dynamic_delta = kwargs.get("dynamic_delta", True) - __identifier_fields__ = tuple() def optimise( @@ -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): From 9e2900a479a6d2010120af18f4b2d6143a3c6a80 Mon Sep 17 00:00:00 2001 From: Richard Date: Mon, 11 Jul 2022 16:25:35 +0100 Subject: [PATCH 2/4] default status is Status not None --- autofit/non_linear/abstract_search.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/autofit/non_linear/abstract_search.py b/autofit/non_linear/abstract_search.py index 86d845416..c8d6a73a6 100644 --- a/autofit/non_linear/abstract_search.py +++ b/autofit/non_linear/abstract_search.py @@ -210,7 +210,7 @@ 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 From 24cb76380dbc410cafbe23ea455157051f163cb7 Mon Sep 17 00:00:00 2001 From: Richard Date: Mon, 11 Jul 2022 16:26:58 +0100 Subject: [PATCH 3/4] fix --- test_autofit/graphical/regression/test_static.py | 1 + 1 file changed, 1 insertion(+) 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, From 74ab906da27001c3dca6a28b9b1ff619dc9edeb5 Mon Sep 17 00:00:00 2001 From: Richard Date: Mon, 11 Jul 2022 16:34:30 +0100 Subject: [PATCH 4/4] format --- autofit/graphical/expectation_propagation/factor_optimiser.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/autofit/graphical/expectation_propagation/factor_optimiser.py b/autofit/graphical/expectation_propagation/factor_optimiser.py index b9ec18980..61613430f 100644 --- a/autofit/graphical/expectation_propagation/factor_optimiser.py +++ b/autofit/graphical/expectation_propagation/factor_optimiser.py @@ -45,13 +45,12 @@ def update_model_approx( for variable, message_count in variable_message_count.items() }) - new_approx, status = model_approx.project_mean_field( + return model_approx.project_mean_field( new_model_dist, factor_approx, delta=delta, status=status, ) - return new_approx, status @abstractmethod def optimise(