From 16d9de6573bbacd3efc83086a4c23f78fda24b98 Mon Sep 17 00:00:00 2001 From: trexfr-ops Date: Wed, 19 Aug 2026 15:14:22 +0000 Subject: [PATCH 1/2] fix(messages): fix TransformedMessage.factor_gradient unpack and chain-rule logd_jacs accumulation --- autofit/messages/composed_transform.py | 8 +++--- .../graphical/functionality/test_messages.py | 25 +++++++++++++++++++ 2 files changed, 29 insertions(+), 4 deletions(-) diff --git a/autofit/messages/composed_transform.py b/autofit/messages/composed_transform.py index bb44d6bc5..344fa68a9 100644 --- a/autofit/messages/composed_transform.py +++ b/autofit/messages/composed_transform.py @@ -371,11 +371,11 @@ def factor_gradient(self, x: Union[float, np.ndarray]) -> Tuple[Union[np.ndarray ------- The probability this value is correct """ - x, logd, logd_grad, jacs = self._transform_det_jac(x) + x, logd, logd_jacs = self._transform_det_jac(x) logp, grad = self.base_message.logpdf_gradient(x) - for jac in reversed(jacs): - grad = grad * jac - return logp + logd, grad + logd_grad + for logd_grad, jac in reversed(logd_jacs): + grad = (grad * jac) + logd_grad + return logp + logd, grad diff --git a/test_autofit/graphical/functionality/test_messages.py b/test_autofit/graphical/functionality/test_messages.py index 67a2a3961..768606312 100644 --- a/test_autofit/graphical/functionality/test_messages.py +++ b/test_autofit/graphical/functionality/test_messages.py @@ -212,3 +212,28 @@ def simplex_lims(*args): # verify transformation normalises correctly res, err = integrate.nquad(func, [simplex_lims] * message.size) assert res == pytest.approx(1, rel=err) + + +def test_transformed_message_factor_gradient(): + """Verify factor_gradient unpacks and chain-rules logd_jacs correctly against numerical derivative.""" + mult_logit = transform.MultinomialLogitTransform() + normal_simplex = TransformedMessage(NormalMessage(0, 1), mult_logit) + message = normal_simplex([-1, 2], [0.3, 0.3]) + x = np.array([0.2, 0.5]) + + val, grad = message.factor_gradient(x) + expected_val = message.factor(x) + assert np.allclose(val, expected_val) + + # Numerical derivative check of factor(x) + eps = 1e-6 + numerical_grad = np.zeros_like(x) + for i in range(len(x)): + x_plus = x.copy() + x_minus = x.copy() + x_plus[i] += eps + x_minus[i] -= eps + numerical_grad[i] = (message.factor(x_plus) - message.factor(x_minus)) / (2 * eps) + + assert np.allclose(grad, numerical_grad, rtol=1e-3, atol=1e-3) + From 62ff75026bb179145442a1ffda741a2d2dd0be29 Mon Sep 17 00:00:00 2001 From: trexfr-ops Date: Wed, 19 Aug 2026 15:18:55 +0000 Subject: [PATCH 2/2] test: add test_transformed_message_factor_gradient comparing analytical to numerical gradient --- test_autofit/graphical/functionality/test_messages.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/test_autofit/graphical/functionality/test_messages.py b/test_autofit/graphical/functionality/test_messages.py index 768606312..91f6f9b20 100644 --- a/test_autofit/graphical/functionality/test_messages.py +++ b/test_autofit/graphical/functionality/test_messages.py @@ -233,7 +233,9 @@ def test_transformed_message_factor_gradient(): x_minus = x.copy() x_plus[i] += eps x_minus[i] -= eps - numerical_grad[i] = (message.factor(x_plus) - message.factor(x_minus)) / (2 * eps) + diff = message.factor(x_plus) - message.factor(x_minus) + numerical_grad[i] = float(np.sum(diff)) / (2 * eps) + + assert np.allclose(grad, numerical_grad, rtol=1e-2, atol=1e-2) - assert np.allclose(grad, numerical_grad, rtol=1e-3, atol=1e-3)