diff --git a/preliz/distributions/mixture.py b/preliz/distributions/mixture.py index 76b4d2e8..6ff1c627 100644 --- a/preliz/distributions/mixture.py +++ b/preliz/distributions/mixture.py @@ -1,4 +1,5 @@ import numpy as np +from scipy.special import logsumexp from preliz.distributions.distributions import DistributionTransformer from preliz.internal.distribution_helper import all_not_none, num_kurtosis, num_skewness @@ -87,17 +88,21 @@ def ppf(self, q): return find_ppf(self, q) def logpdf(self, x): - return np.sum( - [dist.logpdf(x) * weight for dist, weight in zip(self.dist, self.weights)], axis=0 + log_terms = np.array( + [np.log(weight) + dist.logpdf(x) for dist, weight in zip(self.dist, self.weights)] ) + return logsumexp(log_terms, axis=0) def entropy(self): x_values = self.xvals("restricted") logpdf = self.logpdf(x_values) + with np.errstate(divide="ignore", invalid="ignore"): + weighted_logpdf = np.exp(logpdf) * logpdf + weighted_logpdf = np.where(np.isfinite(weighted_logpdf), weighted_logpdf, 0.0) if self.kind == "discrete": - return -np.sum(np.exp(logpdf) * logpdf) + return -np.sum(weighted_logpdf) else: - return -np.trapzoid(np.exp(logpdf) * logpdf, x_values) + return -np.trapezoid(weighted_logpdf, x_values) def mean(self): return np.sum(