diff --git a/src/utils/openmc/filter.py b/src/utils/openmc/filter.py index 8b3cb6ca0c..9ebb42e9f0 100644 --- a/src/utils/openmc/filter.py +++ b/src/utils/openmc/filter.py @@ -245,6 +245,47 @@ class Filter(object): self._stride = stride + def can_merge(self, filter): + + if not isinstance(filter, Filter): + return False + + elif self.type != filter.type: + return False + + elif self.type == 'distribcell': + return False + + elif self.type == 'mesh' and self.bins != filter.bins: + return False + + elif self.type == 'energy' and self.bins != filter.bins: + return False + + elif self.type == 'energyout' and self.bins != filter.bins: + return False + + else: + return True + + + def merge(self, filter): + + if not self.can_merge(filter): + msg = 'Unable to merge {0} with {1} filters'.format(self._type, filter._type) + raise ValueError(msg) + + # Create deep copy of filter to return as merged filter + merged_filter = copy.deepcopy(self) + + # Merge unique filter bins + merged_bins = list(set(self._bins + filter._bins)) + merged_filter.bins = merged_bins + merged_filter.num_bins = len(merged_bins) + + return merged_filter + + def get_bin_index(self, bin): try: diff --git a/src/utils/openmc/mesh.py b/src/utils/openmc/mesh.py index f2c2593f9f..0006bd804e 100644 --- a/src/utils/openmc/mesh.py +++ b/src/utils/openmc/mesh.py @@ -27,6 +27,29 @@ class Mesh(object): self._width = None + def __eq__(self, mesh2): + + # Check type + if self._type != mesh2._type: + return False + + # Check dimension + elif self._dimension != mesh2._dimension: + return False + + # Check width + elif self._width != mesh2._width: + return False + + # Check lower left / upper right + elif self._lower_left != mesh2._lower_left and \ + self._upper_right != mesh2._upper_right: + return False + + else: + return True + + def __deepcopy__(self, memo): existing = memo.get(id(self)) diff --git a/src/utils/openmc/tallies.py b/src/utils/openmc/tallies.py index 4013483cad..8ef5c87df1 100644 --- a/src/utils/openmc/tallies.py +++ b/src/utils/openmc/tallies.py @@ -437,8 +437,15 @@ class Tally(object): if len(self._filters) != len(tally._filters): return False - for filter in self._filters: - if not filter in tally._filters: + for filter1 in self._filters: + contains_filter = False + + for filter2 in tally._filters: + if filter1 == filter2 or filter1.can_merge(filter2): + contains_filter = True + break + + if not contains_filter: return False # Must have same nuclides @@ -453,9 +460,6 @@ class Tally(object): return True - # FIXME: Check if filters are mergeable - - def merge(self, tally): if not self.can_merge(tally): @@ -468,6 +472,14 @@ class Tally(object): # Differentiate Tally with a new auto-generated Tally ID merged_tally.id = None + # Merge filters + for i, filter1 in enumerate(merged_tally._filters): + for filter2 in tally._filters: + if filter1 != filter2 and filter1.can_merge(filter2): + merged_filter = filter1.merge(filter2) + merged_tally._filters[i] = merged_filter + break + # Add scores from second tally to merged tally for score in tally._scores: merged_tally.add_score(score) @@ -476,8 +488,6 @@ class Tally(object): for trigger in tally._triggers: merged_tally.add_trigger(trigger) - # FIXME: Check if filters are mergeable - return merged_tally