Fix type hinting and simplify implementation of combine_distributions (#3445)

Co-authored-by: Paul Romano <paul.k.romano@gmail.com>
This commit is contained in:
GuySten 2026-01-21 16:33:30 +02:00 committed by GitHub
parent 5847b0de23
commit 2691ff8a0f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 40 additions and 26 deletions

View file

@ -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):

View file

@ -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)