diff --git a/openmc/deplete/nuclide.py b/openmc/deplete/nuclide.py index a989bf7029..2ec1d3a66d 100644 --- a/openmc/deplete/nuclide.py +++ b/openmc/deplete/nuclide.py @@ -3,9 +3,11 @@ Contains the per-nuclide components of a depletion chain. """ +import bisect from collections.abc import Mapping from collections import namedtuple, defaultdict from warnings import warn +from numbers import Real try: import lxml.etree as ET except ImportError: @@ -339,6 +341,7 @@ class Nuclide(object): return valid + class FissionYieldDistribution(Mapping): """Energy-dependent fission product yields for a single nuclide @@ -504,10 +507,15 @@ class FissionYield(Mapping): self.products = products self.yields = yields + def __contains__(self, product): + ix = bisect.bisect_left(self.products, product) + return ix != len(self.products) and self.products[ix] == product + def __getitem__(self, product): - if product not in self.products: + ix = bisect.bisect_left(self.products, product) + if ix == len(self.products) or self.products[ix] != product: raise KeyError(product) - return self.yields[self.products.index(product)] + return self.yields[ix] def __len__(self): return len(self.products) @@ -516,35 +524,42 @@ class FissionYield(Mapping): return iter(self.products) def items(self): + """Return pairs of product, yield""" return zip(self.products, self.yields) def __add__(self, other): - new = self.copy() + if not isinstance(other, FissionYield): + return NotImplemented + new = FissionYield(self.products, self.yields.copy()) new += other return new def __iadd__(self, other): """Increment value from other fission yield""" + if not isinstance(other, FissionYield): + return NotImplemented self.yields += other.yields return self + def __radd__(self, other): + return self + other + def __imul__(self, scalar): + if not isinstance(scalar, Real): + return NotImplemented self.yields *= scalar return self def __mul__(self, scalar): - new = self.copy() + if not isinstance(scalar, Real): + return NotImplemented + new = FissionYield(self.products, self.yields.copy()) new *= scalar return new def __rmul__(self, scalar): - new = self.copy() - new *= scalar - return new + return self * scalar def __repr__(self): - return "<{}: {}>".format(self.__class__.__name__, repr(dict(self))) - - def copy(self): - """Return an identical yield object, with unique yields""" - return FissionYield(self.products, self.yields.copy()) + return "<{} containing {} products and yields>".format( + self.__class__.__name__, len(self)) diff --git a/tests/unit_tests/test_deplete_nuclide.py b/tests/unit_tests/test_deplete_nuclide.py index 6f5356b129..7d68b3a3d8 100644 --- a/tests/unit_tests/test_deplete_nuclide.py +++ b/tests/unit_tests/test_deplete_nuclide.py @@ -134,25 +134,42 @@ def test_fission_yield_distribution(): act_dist = yield_dict[exp_ene] for exp_prod, exp_yield in exp_dist.items(): assert act_dist[exp_prod] == exp_yield - exp_yield_matrix = numpy.array([ + exp_yield = numpy.array([ [4.08e-12, 1.71e-12, 7.85e-4], [1.32e-12, 0.0, 1.12e-3], [5.83e-8, 2.69e-8, 4.54e-3]]) - assert numpy.array_equal(yield_dist.yield_matrix, exp_yield_matrix) + assert numpy.array_equal(yield_dist.yield_matrix, exp_yield) # Test the operations / special methods for fission yield - orig_yield_obj = yield_dist[0.0253] + orig_yields = yield_dist[0.0253] + assert len(orig_yields) == len(yield_dict[0.0253]) + for key, value in yield_dict[0.0253].items(): + assert key in orig_yields + assert orig_yields[key] == value # __getitem__ return yields as a view into yield matrix - assert orig_yield_obj.yields.base is yield_dist.yield_matrix - copied_yield = orig_yield_obj.copy() - # copied yields own their own memory -> not a view - assert copied_yield.yields.base is None + assert orig_yields.yields.base is yield_dist.yield_matrix # Fission yield feature uses scaled and incremented - mod_yields = orig_yield_obj * 2 - assert numpy.array_equal(orig_yield_obj.yields * 2, mod_yields.yields) - mod_yields += orig_yield_obj - assert numpy.array_equal(orig_yield_obj.yields * 3, mod_yields.yields) + mod_yields = orig_yields * 2 + assert numpy.array_equal(orig_yields.yields * 2, mod_yields.yields) + mod_yields += orig_yields + assert numpy.array_equal(orig_yields.yields * 3, mod_yields.yields) + + # Failure modes for adding, multiplying yields + similar = numpy.empty_like(orig_yields.yields) + with pytest.raises(TypeError): + orig_yields + similar + with pytest.raises(TypeError): + similar + orig_yields + with pytest.raises(TypeError): + orig_yields += similar + with pytest.raises(TypeError): + orig_yields * similar + with pytest.raises(TypeError): + similar * orig_yields + with pytest.raises(TypeError): + orig_yields *= similar + def test_validate():