diff --git a/openmc/stats/univariate.py b/openmc/stats/univariate.py index 272367f6f..94258c833 100644 --- a/openmc/stats/univariate.py +++ b/openmc/stats/univariate.py @@ -397,7 +397,7 @@ class Discrete(Univariate): def merge( cls, dists: Sequence[Discrete], - probs: Sequence[int] + probs: Sequence[float] ): """Merge multiple discrete distributions into a single distribution @@ -1897,7 +1897,7 @@ class Mixture(Univariate): def combine_distributions( - dists: Sequence[Univariate], + dists: Sequence[Discrete | Tabular], probs: Sequence[float] ): """Combine distributions with specified probabilities @@ -1912,41 +1912,40 @@ def combine_distributions( Parameters ---------- - dists : iterable of openmc.stats.Univariate + dists : sequence of openmc.stats.Discrete or openmc.stats.Tabular Distributions to combine - probs : iterable of float + probs : sequence of float Probability (or intensity) of each distribution """ - # Get copy of distribution list so as not to modify the argument - dist_list = deepcopy(dists) + for i, dist in enumerate(dists): + cv.check_type(f'dists[{i}]', dist, (Discrete, Tabular)) + cv.check_type(f'probs[{i}]', probs[i], Real) + cv.check_greater_than(f'probs[{i}]', probs[i], 0.0) # Get list of discrete/continuous distribution indices - discrete_index = [i for i, d in enumerate(dist_list) if isinstance(d, Discrete)] - cont_index = [i for i, d in enumerate(dist_list) if isinstance(d, Tabular)] + discrete_index = [i for i, d in enumerate(dists) if isinstance(d, Discrete)] + cont_index = [i for i, d in enumerate(dists) if isinstance(d, Tabular)] - # Apply probabilites to continuous distributions - for i in cont_index: - dist = dist_list[i] - dist._p *= probs[i] + cont_dists = [dists[i] for i in cont_index] + cont_probs = [probs[i] for i in cont_index] if discrete_index: # Create combined discrete distribution - dist_discrete = [dist_list[i] for i in discrete_index] + dist_discrete = [dists[i] for i in discrete_index] discrete_probs = [probs[i] for i in discrete_index] combined_dist = Discrete.merge(dist_discrete, discrete_probs) - - # Replace multiple discrete distributions with merged - for idx in reversed(discrete_index): - dist_list.pop(idx) - dist_list.append(combined_dist) - - # Combine discrete and continuous if present - if len(dist_list) > 1: - probs = [1.0]*len(dist_list) - dist_list[:] = [Mixture(probs, dist_list.copy())] - - return dist_list[0] + if cont_index: + return Mixture(cont_probs + [1.0], cont_dists + [combined_dist]) + else: + return combined_dist + else: + if len(cont_dists) == 1: + dist = cont_dists[0] + return Tabular(dist.x, dist.p * cont_probs[0], + dist.interpolation, bias=dist.bias) + else: + return Mixture(cont_probs, cont_dists) def check_bias_support(parent: Univariate, bias: Univariate | None): diff --git a/tests/unit_tests/test_stats.py b/tests/unit_tests/test_stats.py index e92e9e135..507e85743 100644 --- a/tests/unit_tests/test_stats.py +++ b/tests/unit_tests/test_stats.py @@ -633,12 +633,27 @@ def test_combine_distributions(): assert len(mixed.distribution) == 2 assert len(mixed.probability) == 2 + # Single tabular returns a tabular distribution with scaled probabilities + t_single = openmc.stats.Tabular([0.0, 1.0], [2.0, 0.0]) + scaled = openmc.stats.combine_distributions([t_single], [0.25]) + assert isinstance(scaled, openmc.stats.Tabular) + assert scaled.p == pytest.approx([0.5, 0.0]) + + # Mixture with biased tabular should preserve unbiased mean via weights + bias = openmc.stats.Tabular([0.0, 1.0], [2.0, 0.0]) + t_biased = openmc.stats.Tabular([0.0, 1.0], [1.0, 1.0], bias=bias) + d1 = openmc.stats.delta_function(0.0) + mixed = openmc.stats.combine_distributions([t_biased, d1], [0.5, 0.5]) + assert isinstance(mixed, openmc.stats.Mixture) + samples, weights = mixed.sample(10_000) + assert_sample_mean(samples*weights, 0.25) + # Combine 1 discrete and 2 tabular -- the tabular distributions should # combine to produce a uniform distribution with mean 0.5. The combined # distribution should have a mean of 0.25. t1 = openmc.stats.Tabular([0., 1.], [2.0, 0.0]) t2 = openmc.stats.Tabular([0., 1.], [0.0, 2.0]) - d1 = openmc.stats.Discrete([0.0], [1.0]) + d1 = openmc.stats.delta_function(0.0) combined = openmc.stats.combine_distributions([t1, t2, d1], [0.25, 0.25, 0.5]) assert combined.integral() == pytest.approx(1.0)