Refactored tally summation with new AggregateNuclide, AggregateScore and AggregateFilter classes

This commit is contained in:
wbinventor@gmail.com 2016-01-14 15:08:28 -05:00
parent 1aa38c45de
commit 6105f045e1
3 changed files with 569 additions and 95 deletions

404
openmc/aggregate.py Normal file
View file

@ -0,0 +1,404 @@
import sys
from numbers import Integral
import numpy as np
from openmc import Filter, Nuclide
from openmc.cross import CrossScore, CrossNuclide, CrossFilter
from openmc.filter import _FILTER_TYPES
import openmc.checkvalue as cv
if sys.version_info[0] >= 3:
basestring = str
# Acceptable tally aggregation operations
_TALLY_AGGREGATE_OPS = ['sum', 'mean']
class AggregateScore(object):
"""A special-purpose tally score used to encapsulate an aggregate of a
subset or all of tally's scores for tally aggregation.
Parameters
----------
scores : Iterable of str or CrossScore
The scores included in the aggregation
aggregate_op : str
The tally aggregation operator (e.g., 'sum', 'mean', etc.) used
to aggregate across a tally's scores with this AggregateScore
Attributes
----------
scores : Iterable of str or CrossScore
The scores included in the aggregation
aggregate_op : str
The tally aggregation operator (e.g., 'sum', 'mean', etc.) used
to aggregate across a tally's scores with this AggregateScore
"""
def __init__(self, scores=None, aggregate_op=None):
self._scores = None
self._aggregate_op = None
if scores is not None:
self.scores = scores
if aggregate_op is not None:
self.aggregate_op = aggregate_op
def __hash__(self):
return hash(repr(self))
def __eq__(self, other):
return str(other) == str(self)
def __ne__(self, other):
return not self == other
def __deepcopy__(self, memo):
existing = memo.get(id(self))
# If this is the first time we have tried to copy this object, create a copy
if existing is None:
clone = type(self).__new__(type(self))
clone._scores = self.scores
clone._aggregate_op = self.aggregate_op
memo[id(self)] = clone
return clone
# If this object has been copied before, return the first copy made
else:
return existing
def __repr__(self):
string = ''.join(map(str, self.scores))
string = '{0}({1})'.format(self.aggregate_op, string)
return string
@property
def scores(self):
return self._scores
@property
def aggregate_op(self):
return self._aggregate_op
@scores.setter
def scores(self, scores):
cv.check_iterable_type('scores', scores, basestring)
self._scores = scores
@aggregate_op.setter
def aggregate_op(self, aggregate_op):
cv.check_type('aggregate_op', aggregate_op, (basestring, CrossScore))
cv.check_value('aggregate_op', aggregate_op, _TALLY_AGGREGATE_OPS)
self._aggregate_op = aggregate_op
class AggregateNuclide(object):
"""A special-purpose tally nuclide used to encapsulate an aggregate of a
subset or all of tally's nuclides for tally aggregation.
Parameters
----------
nuclides : Iterable of Nuclide or CrossNuclide
The nuclides included in the aggregation
aggregate_op : str
The tally aggregation operator (e.g., 'sum', 'mean', etc.) used
to aggregate across a tally's nuclides with this AggregateNuclide
Attributes
----------
nuclides : Iterable of Nuclide or CrossNuclide
The nuclides included in the aggregation
aggregate_op : str
The tally aggregation operator (e.g., 'sum', 'mean', etc.) used
to aggregate across a tally's nuclides with this AggregateNuclide
"""
def __init__(self, nuclides=None, aggregate_op=None):
self._nuclides = None
self._aggregate_op = None
if nuclides is not None:
self.nuclides = nuclides
if aggregate_op is not None:
self.aggregate_op = aggregate_op
def __hash__(self):
return hash(repr(self))
def __eq__(self, other):
return str(other) == str(self)
def __ne__(self, other):
return not self == other
def __deepcopy__(self, memo):
existing = memo.get(id(self))
# If this is the first time we have tried to copy this object, create a copy
if existing is None:
clone = type(self).__new__(type(self))
clone._nuclides = self.nuclides
clone._aggregate_op = self._aggregate_op
memo[id(self)] = clone
return clone
# If this object has been copied before, return the first copy made
else:
return existing
def __repr__(self):
string = '{0}('.format(self.aggregate_op)
# Append each nuclide in the aggregate to the string
for nuclide in self.nuclides:
if isinstance(nuclide, Nuclide):
string += '{0}, '.format(nuclide.name)
else:
string += '{0}, '.format(str(nuclide))
string += ')'
return string
@property
def nuclides(self):
return self._nuclides
@property
def aggregate_op(self):
return self._aggregate_op
@nuclides.setter
def nuclides(self, nuclides):
cv.check_iterable_type('nuclides', nuclides, (Nuclide, CrossNuclide))
self._nuclides = nuclides
@aggregate_op.setter
def aggregate_op(self, aggregate_op):
cv.check_type('aggregate_op', aggregate_op, basestring)
cv.check_value('aggregate_op', aggregate_op, _TALLY_AGGREGATE_OPS)
self._aggregate_op = aggregate_op
class AggregateFilter(object):
"""A special-purpose tally filter used to encapsulate an aggregate of a
subset or all of a tally filter's bins for tally aggregation.
Parameters
----------
filter : Filter or CrossFilter
The filter included in the aggregation
filter_bins : Iterable of tuple
The filter bins included in the aggregation
aggregate_op : str
The tally aggregation operator (e.g., 'sum', 'mean', etc.) used
to aggregate across a tally filter's bins with this AggregateFilter
Attributes
----------
filter : filter
The filter included in the aggregation
filter_bins : Iterable of tuple
The filter bins included in the aggregation
aggregate_op : str
The tally aggregation operator (e.g., 'sum', 'mean', etc.) used
to aggregate across a tally filter's bins with this AggregateFilter
"""
def __init__(self, filter=None, bins=None, aggregate_op=None):
self._type = '{0}({1})'.format(aggregate_op, filter.type)
self._bins = None
self._stride = None
self._filter = None
self._aggregate_op = None
if filter is not None:
self.filter = filter
if bins is not None:
self.bins = bins
if aggregate_op is not None:
self.aggregate_op = aggregate_op
def __hash__(self):
return hash((self.type, self.bins, self.aggregate_op))
def __eq__(self, other):
return str(other) == str(self)
def __ne__(self, other):
return not self == other
def __repr__(self):
string = 'AggregateFilter\n'
string += '{0: <16}{1}{2}\n'.format('\tType', '=\t', self.type)
string += '{0: <16}{1}{2}\n'.format('\tBins', '=\t', self.bins)
return string
def __deepcopy__(self, memo):
existing = memo.get(id(self))
# If this is the first time we have tried to copy this object, create a copy
if existing is None:
clone = type(self).__new__(type(self))
clone._type = self.type
clone._filter = self.filter
clone._aggregate_op = self.aggregate_op
clone._bins = self._bins
clone._stride = self.stride
memo[id(self)] = clone
return clone
# If this object has been copied before, return the first copy made
else:
return existing
@property
def filter(self):
return self._filter
@property
def aggregate_op(self):
return self._aggregate_op
@property
def type(self):
return self._type
@property
def bins(self):
return self._bins
@property
def num_bins(self):
if self.filter:
return 1
else:
return 0
@property
def stride(self):
return self._stride
@type.setter
def type(self, filter_type):
if filter_type not in _FILTER_TYPES.values():
msg = 'Unable to set AggregateFilter type to "{0}" since it ' \
'is not one of the supported types'.format(filter_type)
raise ValueError(msg)
self._type = filter_type
@filter.setter
def filter(self, filter):
cv.check_type('filter', filter, (Filter, CrossFilter))
self._filter = filter
@bins.setter
def bins(self, bins):
cv.check_iterable_type('bins', bins, (Integral, tuple))
self._bins = bins
@aggregate_op.setter
def aggregate_op(self, aggregate_op):
cv.check_type('aggregate_op', aggregate_op, basestring)
cv.check_value('aggregate_op', aggregate_op, _TALLY_AGGREGATE_OPS)
self._aggregate_op = aggregate_op
@stride.setter
def stride(self, stride):
self._stride = stride
def get_bin_index(self, filter_bin):
"""Returns the index in the AggregateFilter for some bin.
Parameters
----------
filter_bin : Integral or tuple of Integral or tuple of Real
A tuple of value(s) corresponding to the bin of interest in
the aggregated filter. The bin is the integer ID for 'material',
'surface', 'cell', 'cellborn', and 'universe' Filters. The bin
is the integer cell instance ID for 'distribcell' Filters. The
bin is a 2-tuple of floats for 'energy' and 'energyout' filters
corresponding to the energy boundaries of the bin of interest.
The bin is a (x,y,z) 3-tuple for 'mesh' filters corresponding to
the mesh cell of interest.
Returns
-------
filter_index : Integral
The index in the Tally data array for this filter bin. For an
AggregateTally the filter bin index is always unity.
Raises
------
ValueError
When the filter_bin is not part of the aggregated filter's bins
"""
if filter_bin not in self.bins:
msg = 'Unable to get the bin index for AggregateFilter since ' \
'"{0}" is not one of the bins'.format(filter_bin)
raise ValueError(msg)
else:
return 0
def get_pandas_dataframe(self, datasize):
"""Builds a Pandas DataFrame for the AggregateFilter's bins.
This method constructs a Pandas DataFrame object for the AggregateFilter
with columns annotated by filter bin information. This is a helper
method for the Tally.get_pandas_dataframe(...) method.
Parameters
----------
datasize : Integral
The total number of bins in the tally corresponding to this filter
Returns
-------
pandas.DataFrame
A Pandas DataFrame with columns of strings that characterize the
aggregatefilter's bins. Each entry in the DataFrame will include
one or more aggregation operations used to construct the
aggregatefilter's bins. The number of rows in the DataFrame is the
same as the total number of bins in the corresponding tally, with
the filter bins appropriately tiled to map to the corresponding
tally bins.
See also
--------
Tally.get_pandas_dataframe(), Filter.get_pandas_dataframe(),
CrossFilter.get_pandas_dataframe()
"""
import pandas as pd
# Construct a sring representing the filter aggregation
aggregate_bin = '{0}('.format(self.aggregate_op)
aggregate_bin += ', '.join(self.bins) + ')'
# Construct NumPy array of bin repeated for each element in dataframe
aggregate_bin_array = np.array([aggregate_bin])
aggregate_bin_array = np.repeat(aggregate_bin_array, datasize)
# Construct Pandas DataFrame for the AggregateFilter
df = pd.DataFrame({self.type : aggregate_bin_array})
return df

View file

@ -34,7 +34,7 @@ class CrossScore(object):
The right score in the outer product
binary_op : str
The tally arithmetic binary operator (e.g., '+', '-', etc.) used to
combine two tally's scores with this CrossNuclide
combine two tally's scores with this CrossScore
"""
@ -310,7 +310,6 @@ class CrossFilter(object):
clone._binary_op = self.binary_op
clone._type = self.type
clone._bins = self._bins
clone._num_bins = self.num_bins
clone._stride = self.stride
memo[id(self)] = clone
@ -355,8 +354,8 @@ class CrossFilter(object):
@type.setter
def type(self, filter_type):
if filter_type not in _FILTER_TYPES.values():
msg = 'Unable to set Filter type to "{0}" since it is not one ' \
'of the supported types'.format(filter_type)
msg = 'Unable to set CrossFilter type to "{0}" since it ' \
'is not one of the supported types'.format(filter_type)
raise ValueError(msg)
self._type = filter_type

View file

@ -13,6 +13,7 @@ import numpy as np
from openmc import Mesh, Filter, Trigger, Nuclide
from openmc.cross import CrossScore, CrossNuclide, CrossFilter
from openmc.aggregate import AggregateScore, AggregateNuclide, AggregateFilter
from openmc.filter import _FILTER_TYPES
import openmc.checkvalue as cv
from openmc.clean_xml import *
@ -259,6 +260,10 @@ class Tally(object):
def num_scores(self):
return len(self._scores)
@property
def num_filters(self):
return len(self.filters)
@property
def num_filter_bins(self):
num_bins = 1
@ -460,16 +465,28 @@ class Tally(object):
Parameters
----------
filter : openmc.filter.Filter
Filter to add
filter : Filter, CrossFilter or AggregateFilter
A filter to specify a discretization of the tally across some
dimension (e.g., 'energy', 'cell'). The filter should be a Filter
object when a user is adding filters to a Tally for input file
generation or when the Tally is created from a StatePoint. The
filter may be a CrossFilter or AggregateFilter for derived tallies
created by tally arithmetic.
"""
if not isinstance(filter, (Filter, CrossFilter)):
if not isinstance(filter, (Filter, CrossFilter, AggregateFilter)):
msg = 'Unable to add Filter "{0}" to Tally ID="{1}" since it is ' \
'not a Filter object'.format(filter, self.id)
raise ValueError(msg)
# If the filter is already in the Tally, raise an error
if filter in self.filters:
msg = 'Unable to add a duplicate filter "{0}" to Tally ID="{1}" ' \
'since duplicate filters are not supported in the OpenMC ' \
'Python API'.format(filter, self.id)
raise ValueError(msg)
self._filters.append(filter)
def add_nuclide(self, nuclide):
@ -477,11 +494,29 @@ class Tally(object):
Parameters
----------
nuclide : openmc.nuclide.Nuclide
Nuclide to add
nuclide : str, Nuclide, CrossNuclide or AggregateNuclide
Nuclide to add to the tally. The nuclide should be a Nuclide object
when a user is adding nuclides to a Tally for input file generation.
The nuclide is a str when a Tally is created from a StatePoint file
(e.g., 'H-1', 'U-235') unless a Summary has been linked with the
StatePoint. The nuclide may be a CrossNuclide or AggregateNuclide
for derived tallies created by tally arithmetic.
"""
if not isinstance(nuclide, (basestring, Nuclide,
CrossNuclide, AggregateNuclide)):
msg = 'Unable to add nuclide "{0}" to Tally ID="{1}" since it is ' \
'not a Nuclide object'.format(nuclide)
raise ValueError(msg)
# If the nuclide is already in the Tally, raise an error
if nuclide in self.nuclides:
msg = 'Unable to add a duplicate nuclide "{0}" to Tally ID="{1}" ' \
'since duplicate nuclides are not supported in the OpenMC ' \
'Python API'.format(nuclide, self.id)
raise ValueError(msg)
self._nuclides.append(nuclide)
def add_score(self, score):
@ -489,19 +524,23 @@ class Tally(object):
Parameters
----------
score : str
Score to be accumulated, e.g. 'flux'
score : str, CrossScore or AggregateScore
A score to be accumulated (e.g., 'flux', 'nu-fission'). The score
should be a str when a user is adding scores to a Tally for input
file generation or when the Tally is created from a StatePoint. The
score may be a CrossScore or AggregateScore for derived tallies
created by tally arithmetic.
"""
if not isinstance(score, (basestring, CrossScore)):
if not isinstance(score, (basestring, CrossScore, AggregateScore)):
msg = 'Unable to add score "{0}" to Tally ID="{1}" since it is ' \
'not a string'.format(score, self.id)
raise ValueError(msg)
# If the score is already in the Tally, raise an error
if score in self.scores:
msg = 'Unable to add a duplicate score {0} to Tally ID="{1}" ' \
msg = 'Unable to add a duplicate score "{0}" to Tally ID="{1}" ' \
'since duplicate scores are not supported in the OpenMC ' \
'Python API'.format(score, self.id)
raise ValueError(msg)
@ -509,7 +548,7 @@ class Tally(object):
# Normal score strings
if isinstance(score, basestring):
self._scores.append(score.strip())
# CrossScores
# CrossScores and AggrgateScore
else:
self._scores.append(score)
@ -1599,8 +1638,12 @@ class Tally(object):
raise ValueError(msg)
new_tally = Tally()
new_tally.with_batch_statistics = True
new_tally._derived = True
new_tally.with_batch_statistics = True
new_tally._num_realizations = self.num_realizations
new_tally._estimator = self.estimator
new_tally._with_summary = self.with_summary
new_tally._sp_filename = self._sp_filename
# Construct a combined derived name from the two tally operands
if self.name != '' and other.name != '':
@ -1698,14 +1741,21 @@ class Tally(object):
new_score = CrossScore(self_score, other_score, binary_op)
new_tally.add_score(new_score)
# Correct each Filter's stride
stride = new_tally.num_nuclides * new_tally.num_scores
for filter in reversed(new_tally.filters):
filter.stride = stride
stride *= filter.num_bins
# Update the new tally's filter strides
new_tally._update_filter_strides()
return new_tally
def _update_filter_strides(self):
"""Update each filter's stride based on the tally's nuclides and scores
for derived tallies created by tally arithmetic.
"""
stride = self.num_nuclides * self.num_scores
for filter in reversed(self.filters):
filter.stride = stride
stride *= filter.num_bins
def _align_tally_data(self, other, filter_product, nuclide_product,
score_product):
"""Aligns data from two tallies for tally arithmetic.
@ -1855,17 +1905,9 @@ class Tally(object):
if other_index != i:
other._swap_scores(score, other.scores[i])
# Correct the stride for other filters
stride = other.num_nuclides * other.num_scores
for filter in reversed(other.filters):
filter.stride = stride
stride *= filter.num_bins
# Correct the stride for self filters
stride = self.num_nuclides * self.num_scores
for filter in reversed(self.filters):
filter.stride = stride
stride *= filter.num_bins
# Update the tallies' filter strides
other._update_filter_strides()
self._update_filter_strides()
data = {}
data['self'] = {}
@ -1927,11 +1969,8 @@ class Tally(object):
self.filters[filter1_index] = filter2
self.filters[filter2_index] = filter1
# Update the strides for each of the filters
stride = self.num_nuclides * self.num_scores
for filter in reversed(self.filters):
filter.stride = stride
stride *= filter.num_bins
# Update the tally's filter strides
self._update_filter_strides()
# Construct lists of tuples for the bins in each of the two filters
filters = [filter1.type, filter2.type]
@ -2449,8 +2488,8 @@ class Tally(object):
if isinstance(power, Tally):
new_tally = self.hybrid_product(power, binary_op='^')
# If original tally operands were sparse, sparsify the new tally
if self.sparse and other.sparse:
# If original tally operand was sparse, sparsify the new tally
if self.sparse:
new_tally.sparse = True
elif isinstance(power, Real):
@ -2708,11 +2747,8 @@ class Tally(object):
filter.bins = filter.bins[bin_indices]
filter.num_bins = num_bins
# Correct each Filter's stride
stride = new_tally.num_nuclides * new_tally.num_scores
for filter in reversed(new_tally.filters):
filter.stride = stride
stride *= filter.num_bins
# Update the new tally's filter strides
new_tally._update_filter_strides()
# If original tally was sparse, sparsify the sliced tally
new_tally.sparse = self.sparse
@ -2758,69 +2794,107 @@ class Tally(object):
A new tally which encapsulates the sum of data requested.
"""
# If user did not specify any scores, do not sum across scores
if len(scores) == 0:
scores = [[]]
# Sum across any scores specified by the user
else:
scores = [[score] for score in scores]
# Create new derived Tally for summation
tally_sum = Tally()
tally_sum._derived = True
tally_sum._estimator = self.estimator
tally_sum._num_realizations = self.num_realizations
tally_sum.with_batch_statistics = self.with_batch_statistics
tally_sum._with_summary = self.with_summary
tally_sum._sp_filename = self._sp_filename
tally_sum._results_read = self._results_read
# If user did not specify any nuclides, do not sum across nuclides
if len(nuclides) == 0:
nuclides = [[]]
# Sum across any nuclides specified by the user
else:
nuclides = [[nuclide] for nuclide in nuclides]
# Get tally data arrays reshaped with one dimension per filter
mean = self.get_reshaped_data(value='mean')
std_dev = self.get_reshaped_data(value='std_dev')
# Sum across any filter bins specified by the user
if filter_type in _FILTER_TYPES:
filter = self.find_filter(filter_type)
# If user did not specify filter bins, sum across all bins
if len(filter_bins) == 0:
filter = self.find_filter(filter_type)
if filter.type == 'distribcell':
filter_bins = [[(i,)] for i in range(filter.num_bins)]
bin_indices = np.arange(filter.num_bins)
if filter_type == 'distribcell':
filter_bins = np.arange(filter.num_bins)
else:
filter_bins = \
[[(filter.get_bin(i),)] for i in range(filter.num_bins)]
[(filter.get_bin(i)) for i in range(filter.num_bins)]
# Only sum across bins specified by the user
else:
filter_bins = [[(filter_bin,)] for filter_bin in filter_bins]
bin_indices = \
[filter.get_bin_index(i) for i in range(filter_bins)]
filters = [[filter_type]]
# If user did not specify a filter type, do not sum across filter bins
# Sum across the bins in the user-specified filter
for i, filter in enumerate(self.filters):
if filter.type == filter_type:
mean = np.take(mean, indices=bin_indices, axis=i)
std_dev = np.take(std_dev, indices=bin_indices, axis=i)
mean = np.sum(mean, axis=i, keepdims=True)
std_dev = np.sum(std_dev**2, axis=i, keepdims=True)
std_dev = np.sqrt(std_dev)
# Add AggregateFilter to the tally sum
if not remove_filter:
filter_sum = AggregateFilter(filter, filter_bins, 'sum')
tally_sum.add_filter(filter_sum)
# Add a copy of each filter not summed across to the tally sum
else:
tally_sum.add_filter(copy.deepcopy(filter))
# Add a copy of this tally's filters to the tally sum
else:
filter_bins = [[]]
filters = [[]]
tally_sum._filters = copy.deepcopy(self.filters)
# Initialize Tally sum
tally_sum = 0
# Sum across any nuclides specified by the user
if len(nuclides) != 0:
nuclide_bins = [self.get_nuclide_index(nuclide) for nuclide in nuclides]
axis_index = len(self.filters)
mean = np.take(mean, indices=nuclide_bins, axis=axis_index)
std_dev = np.take(std_dev, indices=nuclide_bins, axis=axis_index)
mean = np.sum(mean, axis=axis_index, keepdims=True)
std_dev = np.sum(std_dev**2, axis=axis_index, keepdims=True)
std_dev = np.sqrt(std_dev)
# Iterate over all Tally slice operands in summation
prod = [scores, filters, filter_bins, nuclides]
summed_filters = defaultdict(list)
for scores, filters, filter_bins, nuclides in itertools.product(*prod):
tally_slice = self.get_slice(scores, filters, filter_bins, nuclides)
# Add AggregateNuclide to the tally sum
nuclide_sum = AggregateNuclide(nuclides, 'sum')
tally_sum.add_nuclide(nuclide_sum)
# Remove filters summed across to avoid bulky CrossFilters
if filter_type:
filter = tally_slice.find_filter(filter_type)
tally_slice.remove_filter(filter)
summed_filters[filter_type].append(filter)
# Accumulate this Tally slice into the Tally sum
tally_sum += tally_slice
# Add back the filter(s) which were summed across to derived tally,
# if filter bins were input; otherwise, leave out summed filter(s)
if remove_filter and filter_type is not None:
# Rename tally sum indicating a summation over a particular filter
tally_sum.name = 'sum({0}, {1})'.format(self.name, filter_type)
# Add a copy of this tally's nuclides to the tally sum
else:
for summed_filter_type in summed_filters:
filters = summed_filters[summed_filter_type]
for i in range(1, len(filters)):
filters[i] = CrossFilter(filters[i-1], filters[i], '+')
tally_sum.add_filter(filters[-1])
tally_sum._nuclides = copy.deepcopy(self.nuclides)
# Sum across any scores specified by the user
if len(scores) != 0:
score_bins = [self.get_score_index(score) for score in scores]
axis_index = self.num_nuclides + len(self.filters)
mean = np.take(mean, indices=score_bins, axis=axis_index)
std_dev = np.take(std_dev, indices=score_bins, axis=axis_index)
mean = np.sum(mean, axis=axis_index, keepdims=True)
std_dev = np.sum(std_dev**2, axis=axis_index, keepdims=True)
std_dev = np.sqrt(std_dev)
# Add AggregateScore to the tally sum
score_sum = AggregateScore(scores, 'sum')
tally_sum.add_score(score_sum)
# Add a copy of this tally's scores to the tally sum
else:
tally_sum._scores = copy.deepcopy(self.scores)
# Update the tally sum's filter strides
tally_sum._update_filter_strides()
# Reshape condensed data arrays with one dimension for all filters
mean = np.reshape(mean, tally_sum.shape)
std_dev = np.reshape(std_dev, tally_sum.shape)
# Assign tally sum's data with the new arrays
tally_sum._mean = mean
tally_sum._std_dev = std_dev
# If original tally was sparse, sparsify the tally summation
tally_sum.sparse = self.sparse
@ -2887,11 +2961,8 @@ class Tally(object):
new_tally._std_dev = np.zeros(new_tally.shape, dtype=np.float64)
new_tally._std_dev[diag_indices, :, :] = self.std_dev
# Correct each Filter's stride
stride = new_tally.num_nuclides * new_tally.num_scores
for filter in reversed(new_tally.filters):
filter.stride = stride
stride *= filter.num_bins
# Update the new tally's filter strides
new_tally._update_filter_strides()
# If original tally was sparse, sparsify the diagonalized tally
new_tally.sparse = self.sparse