mirror of
https://github.com/openmc-dev/openmc.git
synced 2026-07-21 14:35:27 -04:00
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:
parent
5847b0de23
commit
2691ff8a0f
2 changed files with 40 additions and 26 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue