Made filters mergeable in Python API

This commit is contained in:
Will Boyd 2015-05-11 18:59:14 -04:00
parent d5ba1a3273
commit 0605658ff1
3 changed files with 81 additions and 7 deletions

View file

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

View file

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

View file

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