diff --git a/openmc/data/angle_distribution.py b/openmc/data/angle_distribution.py index 316559c50d..afb9a05454 100644 --- a/openmc/data/angle_distribution.py +++ b/openmc/data/angle_distribution.py @@ -32,6 +32,16 @@ class AngleDistribution(object): self.energy = energy self.mu = mu + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (not np.array_equal(self.energy, other.energy) or + not np.array_equal(self.mu, other.mu)): + eqval = False + return eqval + @property def energy(self): return self._energy diff --git a/openmc/data/correlated.py b/openmc/data/correlated.py index 97a96c17c1..2415aaab30 100644 --- a/openmc/data/correlated.py +++ b/openmc/data/correlated.py @@ -49,6 +49,19 @@ class CorrelatedAngleEnergy(AngleEnergy): self.energy_out = energy_out self.mu = mu + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (not np.array_equal(self.breakpoints, other.breakpoints) or + not np.array_equal(self.interpolation, other.interpolation) or + not np.array_equal(self.energy, other.energy) or + not np.array_equal(self.energy_out, other.energy_out) or + not np.array_equal(self.mu, other.mu)): + eqval = False + return eqval + @property def breakpoints(self): return self._breakpoints diff --git a/openmc/data/energy_distribution.py b/openmc/data/energy_distribution.py index 2300081c13..4299fc166a 100644 --- a/openmc/data/energy_distribution.py +++ b/openmc/data/energy_distribution.py @@ -84,6 +84,16 @@ class ArbitraryTabulated(EnergyDistribution): self.energy = energy self.pdf = pdf + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (not np.array_equal(self.energy, other.energy) or + not np.array_equal(self.pdf, other.pdf)): + eqval = False + return eqval + def to_hdf5(self, group): raise NotImplementedError @@ -122,6 +132,17 @@ class GeneralEvaporation(EnergyDistribution): self.g = g self.u = u + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self.theta != other.theta or + self.g != other.g or + self.u != other.u): + eqval = False + return eqval + def to_hdf5(self, group): raise NotImplementedError @@ -159,6 +180,16 @@ class MaxwellEnergy(EnergyDistribution): self.theta = theta self.u = u + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self.theta != other.theta or + self.u != other.u): + eqval = False + return eqval + @property def theta(self): return self._theta @@ -267,6 +298,16 @@ class Evaporation(EnergyDistribution): self.theta = theta self.u = u + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self.theta != other.theta or + self.u != other.u): + eqval = False + return eqval + @property def theta(self): return self._theta @@ -379,6 +420,17 @@ class WattEnergy(EnergyDistribution): self.b = b self.u = u + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self.a != other.a or + self.b != other.b or + self.u != other.u): + eqval = False + return eqval + @property def a(self): return self._a @@ -520,6 +572,17 @@ class MadlandNix(EnergyDistribution): self.efh = efh self.tm = tm + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self.efl != other.efl or + self.efh != other.efh or + self.tm != other.tm): + eqval = False + return eqval + @property def efl(self): return self._efl @@ -618,6 +681,17 @@ class DiscretePhoton(EnergyDistribution): self.energy = energy self.atomic_weight_ratio = atomic_weight_ratio + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self.primary_flag != other.primary_flag or + self.energy != other.energy or + self.atomic_weight_ratio != other.atomic_weight_ratio): + eqval = False + return eqval + @property def primary_flag(self): return self._primary_flag @@ -726,6 +800,16 @@ class LevelInelastic(EnergyDistribution): self.threshold = threshold self.mass_ratio = mass_ratio + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self.threshold != other.threshold or + self.mass_ratio != other.mass_ratio): + eqval = False + return eqval + @property def threshold(self): return self._threshold @@ -832,6 +916,18 @@ class ContinuousTabular(EnergyDistribution): self.energy = energy self.energy_out = energy_out + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (not np.array_equal(self.breakpoints, other.breakpoints) or + not np.array_equal(self.interpolation, other.interpolation) or + not np.array_equal(self.energy, other.energy) or + not np.array_equal(self.energy_out, other.energy_out)): + eqval = False + return eqval + @property def breakpoints(self): return self._breakpoints diff --git a/openmc/data/function.py b/openmc/data/function.py index 56a3d4a45f..e0acf9f29d 100644 --- a/openmc/data/function.py +++ b/openmc/data/function.py @@ -65,6 +65,18 @@ class Tabulated1D(object): self.x = np.asarray(x) self.y = np.asarray(y) + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (not np.array_equal(self.x, other.x) or + not np.array_equal(self.y, other.y) or + not np.array_equal(self.breakpoints, other.breakpoints) or + not np.array_equal(self.interpolation, other.interpolation)): + eqval = False + return eqval + def __call__(self, x): # Check if input is array or scalar if isinstance(x, Iterable): diff --git a/openmc/data/kalbach_mann.py b/openmc/data/kalbach_mann.py index 5aca17ba9f..519796ad46 100644 --- a/openmc/data/kalbach_mann.py +++ b/openmc/data/kalbach_mann.py @@ -59,6 +59,20 @@ class KalbachMann(AngleEnergy): self.precompound = precompound self.slope = slope + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (not np.array_equal(self.breakpoints, other.breakpoints) or + not np.array_equal(self.interpolation, other.interpolation) or + not np.array_equal(self.energy, other.energy) or + not np.array_equal(self.energy_out, other.energy_out) or + not np.array_equal(self.precompound, other.precompound) or + not np.array_equal(self.slope, other.slope)): + eqval = False + return eqval + @property def breakpoints(self): return self._breakpoints diff --git a/openmc/data/nbody.py b/openmc/data/nbody.py index a34380dc93..274db48f98 100644 --- a/openmc/data/nbody.py +++ b/openmc/data/nbody.py @@ -38,6 +38,18 @@ class NBodyPhaseSpace(AngleEnergy): self.atomic_weight_ratio = atomic_weight_ratio self.q_value = q_value + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self.total_mass != other.total_mass or + self.n_particles != other.n_particles or + self.atomic_weight_ratio != other.atomic_weight_ratio or + self.q_value != other.q_value): + eqval = False + return eqval + @property def total_mass(self): return self._total_mass diff --git a/openmc/data/product.py b/openmc/data/product.py index dd276daac3..08529f90a1 100644 --- a/openmc/data/product.py +++ b/openmc/data/product.py @@ -49,6 +49,20 @@ class Product(object): self.applicability = [] self.yield_ = 1 + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self.particle != other.particle or + self.decay_rate != other.decay_rate or + self.emission_mode != other.emission_mode or + self.distribution.all() != other.distribution.all() or + self.applicability.all() != other.applicability.all() or + self.yield_ != other.yield_): + eqval = False + return eqval + def __repr__(self): if isinstance(self.yield_, Real): return "".format( diff --git a/openmc/data/thermal.py b/openmc/data/thermal.py index 9de39f518d..cb61aa7300 100644 --- a/openmc/data/thermal.py +++ b/openmc/data/thermal.py @@ -69,6 +69,16 @@ class CoherentElastic(object): idx = np.searchsorted(self.bragg_edges, E) return self.factors[idx]/E + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (not np.array_equal(self.bragg_edges, other.bragg_edges) or + not np.array_equal(self.factors, other.factors)): + eqval = False + return eqval + def __len__(self): return len(self.bragg_edges) @@ -170,6 +180,23 @@ class ThermalScattering(object): self.secondary_mode = None self.zaids = [] + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self.name != other.name or + self.atomic_weight_ratio != other.atomic_weight_ratio or + self.elastic_xs != other.elastic_xs or + self.elastic_mu_out != other.elastic_mu_out or + self.inelastic_xs != other.inelastic_xs or + self.inelastic_mu_out != other.inelastic_mu_out or + self.temperature != other.temperature or + not np.array_equal(self.zaids, other.zaids) or + self.secondary_mode != other.secondary_mode): + eqval = False + return eqval + def __repr__(self): if hasattr(self, 'name'): return "".format(self.name) @@ -216,13 +243,15 @@ class ThermalScattering(object): self.inelastic_dist.to_hdf5(inelastic_group) @classmethod - def from_hdf5(cls, group): + def from_hdf5(cls, group_or_filename): """Generate thermal scattering data from HDF5 group Parameters ---------- - group : h5py.Group - HDF5 group to read from + group_or_filename : h5py.Group or str + HDF5 group containing interaction data. If given as a string, it is + assumed to be the filename for the HDF5 file, and the first group + is used to read from. Returns ------- @@ -230,6 +259,12 @@ class ThermalScattering(object): Neutron thermal scattering data """ + if isinstance(group_or_filename, h5py.Group): + group = group_or_filename + else: + h5file = h5py.File(group_or_filename, 'r') + group = list(h5file.values())[0] + name = group.name[1:] atomic_weight_ratio = group.attrs['atomic_weight_ratio'] temperature = group.attrs['temperature'] diff --git a/openmc/data/uncorrelated.py b/openmc/data/uncorrelated.py index 0361c9b37a..babe003c02 100644 --- a/openmc/data/uncorrelated.py +++ b/openmc/data/uncorrelated.py @@ -34,6 +34,16 @@ class UncorrelatedAngleEnergy(AngleEnergy): if energy is not None: self.energy = energy + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self.angle != other.angle or + self.energy != other.energy): + eqval = False + return eqval + @property def angle(self): return self._angle diff --git a/openmc/stats/univariate.py b/openmc/stats/univariate.py index 24ab98895d..d7ba945202 100644 --- a/openmc/stats/univariate.py +++ b/openmc/stats/univariate.py @@ -65,6 +65,16 @@ class Discrete(Univariate): self.x = x self.p = p + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (not np.array_equal(self.x, other.x) or + not np.array_equal(self.p, other.p)): + eqval = False + return eqval + def __len__(self): return len(self.x) @@ -126,6 +136,16 @@ class Uniform(Univariate): self.a = a self.b = b + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self.a != other.a or + self.b != other.b): + eqval = False + return eqval + def __len__(self): return 2 @@ -183,6 +203,15 @@ class Maxwell(Univariate): super(Maxwell, self).__init__() self.theta = theta + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self.theta != other.theta): + eqval = False + return eqval + def __len__(self): return 1 @@ -231,6 +260,16 @@ class Watt(Univariate): self.a = a self.b = b + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self.a != other.a or + self.b != other.b): + eqval = False + return eqval + def __len__(self): return 2 @@ -300,6 +339,18 @@ class Tabular(Univariate): self.p = p self.interpolation = interpolation + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (self._ignore_negative != other._ignore_negative or + not np.array_equal(self.x, other.x) or + not np.array_equal(self.p, other.p) or + self.interpolation != other.interpolation): + eqval = False + return eqval + def __len__(self): return len(self.x) @@ -368,6 +419,15 @@ class Legendre(Univariate): def __call__(self, x): return self._legendre_polynomial(x) + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if not np.array_equal(self.coefficients, other.coefficients): + eqval = False + return eqval + def __len__(self): return len(self._legendre_polynomial.coef) @@ -414,6 +474,16 @@ class Mixture(Univariate): self.probability = probability self.distribution = distribution + def __eq__(self, other): + if not isinstance(other, type(self)): + return NotImplemented + else: + eqval = True + if (not np.array_equal(self.probability, other.probability) or + not np.array_equal(self.distribution, other.distribution)): + eqval = False + return eqval + def __len__(self): return sum(len(d) for d in self.distribution)