Implement Tallies.from_xml classmethod

This commit is contained in:
Paul Romano 2022-01-20 17:51:16 -06:00
parent 5dfc880d91
commit 6bd33602cc
6 changed files with 281 additions and 48 deletions

View file

@ -16,6 +16,7 @@ from .material import Material
from .mixin import IDManagerMixin
from .surface import Surface
from .universe import UniverseBase
from ._xml import get_text
_FILTER_TYPES = (
@ -231,9 +232,43 @@ class Filter(IDManagerMixin, metaclass=FilterMeta):
subelement = ET.SubElement(element, 'bins')
subelement.text = ' '.join(str(b) for b in self.bins)
return element
@classmethod
def from_xml_element(cls, elem, **kwargs):
"""Generate a filter from an XML element
Parameters
----------
elem : xml.etree.ElementTree.Element
XML element
**kwargs
Keyword arguments (e.g., mesh information)
Returns
-------
openmc.Filter
Filter object
"""
filter_type = elem.get('type')
# If the filter type matches this class's short_name, then
# there is no overriden from_xml_element method.
if filter_type == cls.short_name.lower():
# Get bins from element -- the default here works for any filters
# that just store a list of bins that can be represented as integers
filter_id = int(elem.get('id'))
bins = [int(x) for x in get_text(elem, 'bins').split()]
return cls(bins, filter_id=filter_id)
# Search through all subclasses and find the one matching the HDF5
# 'type'. Call that class's from_hdf5 method.
for subclass in cls._recursive_subclasses():
if filter_type == subclass.short_name.lower():
return subclass.from_xml_element(elem, **kwargs)
def can_merge(self, other):
"""Determine if filter can be merged with another.
@ -622,6 +657,13 @@ class CellInstanceFilter(Filter):
subelement.text = ' '.join(str(i) for i in self.bins.ravel())
return element
@classmethod
def from_xml_element(cls, elem, **kwargs):
filter_id = int(elem.get('id'))
bins = [int(x) for x in get_text(elem, 'bins').split()]
cell_instances = list(zip(bins[::2], bins[1::2]))
return cls(cell_instances, filter_id=filter_id)
class SurfaceFilter(WithIDFilter):
"""Filters particles by surface crossing
@ -661,8 +703,8 @@ class ParticleFilter(Filter):
Attributes
----------
bins : Iterable of Integral
The Particles to tally
bins : iterable of str
The particles to tally
id : int
Unique identifier for the filter
num_bins : Integral
@ -698,6 +740,12 @@ class ParticleFilter(Filter):
filter_id = int(group.name.split('/')[-1].lstrip('filter '))
return cls(particles, filter_id=filter_id)
@classmethod
def from_xml_element(cls, elem, **kwargs):
filter_id = int(elem.get('id'))
bins = get_text(elem, 'bins').split()
return cls(bins, filter_id=filter_id)
class MeshFilter(Filter):
"""Bins tally event locations onto a regular, rectangular mesh.
@ -877,6 +925,18 @@ class MeshFilter(Filter):
element.set('translation', ' '.join(map(str, self.translation)))
return element
@classmethod
def from_xml_element(cls, elem, **kwargs):
mesh_id = int(get_text(elem, 'bins'))
mesh_obj = kwargs['meshes'][mesh_id]
filter_id = int(elem.get('id'))
out = cls(mesh_obj, filter_id=filter_id)
translation = elem.get('translation')
if translation:
out.translation = [float(x) for x in translation.split()]
return out
class MeshSurfaceFilter(MeshFilter):
"""Filter events by surface crossings on a regular, rectangular mesh.
@ -1019,35 +1079,12 @@ class CollisionFilter(Filter):
self.bins = np.asarray(bins)
self.id = filter_id
def __repr__(self):
string = type(self).__name__ + '\n'
string += '{: <16}=\t{}\n'.format('\tValues', self.bins)
string += '{: <16}=\t{}\n'.format('\tID', self.id)
return string
@Filter.bins.setter
def bins(self, bins):
Filter.bins.__set__(self, np.asarray(bins))
def check_bins(self, bins):
for x in bins:
# Values should be integers
cv.check_type('filter value', x, Integral)
cv.check_greater_than('filter value', x, 0, equality=True)
def to_xml_element(self):
"""Return XML Element representing the Filter.
Returns
-------
element : xml.etree.ElementTree.Element
XML element containing filter data
"""
element = super().to_xml_element()
element[0].text = ' '.join(str(x) for x in self.bins)
return element
class RealFilter(Filter):
"""Tally modifier that describes phase-space and other characteristics
@ -1236,6 +1273,12 @@ class RealFilter(Filter):
element[0].text = ' '.join(str(x) for x in self.values)
return element
@classmethod
def from_xml_element(cls, elem, **kwargs):
filter_id = int(elem.get('id'))
bins = [float(x) for x in get_text(elem, 'bins').split()]
return cls(bins, filter_id=filter_id)
class EnergyFilter(RealFilter):
"""Bins tally events based on incident particle energy.
@ -1969,6 +2012,13 @@ class EnergyFunctionFilter(Filter):
return element
@classmethod
def from_xml_element(cls, elem, **kwargs):
filter_id = int(elem.get('id'))
energy = [float(x) for x in get_text(elem, 'energy').split()]
y = [float(x) for x in get_text(elem, 'y').split()]
return cls(energy, y, filter_id=filter_id)
def can_merge(self, other):
return False

View file

@ -46,6 +46,12 @@ class ExpansionFilter(Filter):
return element
@classmethod
def from_xml_element(cls, elem, **kwargs):
filter_id = int(elem.get('id'))
order = int(elem.find('order').text)
return cls(order, filter_id=filter_id)
class LegendreFilter(ExpansionFilter):
r"""Score Legendre expansion moments up to specified order.
@ -226,6 +232,15 @@ class SpatialLegendreFilter(ExpansionFilter):
return element
@classmethod
def from_xml_element(cls, elem, **kwargs):
filter_id = int(elem.get('id'))
order = int(elem.find('order').text)
axis = elem.find('axis').text
minimum = float(elem.find('min').text)
maximum = float(elem.find('max').text)
return cls(order, axis, minimum, maximum, filter_id=filter_id)
class SphericalHarmonicsFilter(ExpansionFilter):
r"""Score spherical harmonic expansion moments up to specified order.
@ -316,6 +331,14 @@ class SphericalHarmonicsFilter(ExpansionFilter):
element.set('cosine', self.cosine)
return element
@classmethod
def from_xml_element(cls, elem, **kwargs):
filter_id = int(elem.get('id'))
order = int(elem.find('order').text)
filter = cls(order, filter_id=filter_id)
filter.cosine = elem.get('cosine')
return filter
class ZernikeFilter(ExpansionFilter):
r"""Score Zernike expansion moments in space up to specified order.
@ -358,7 +381,7 @@ class ZernikeFilter(ExpansionFilter):
x-coordinate of center of circle for normalization
y : float
y-coordinate of center of circle for normalization
r : int or None
r : float
Radius of circle for normalization
Attributes
@ -369,7 +392,7 @@ class ZernikeFilter(ExpansionFilter):
x-coordinate of center of circle for normalization
y : float
y-coordinate of center of circle for normalization
r : int or None
r : float
Radius of circle for normalization
id : int
Unique identifier for the filter
@ -464,6 +487,15 @@ class ZernikeFilter(ExpansionFilter):
return element
@classmethod
def from_xml_element(cls, elem, **kwargs):
filter_id = int(elem.get('id'))
order = int(elem.find('order').text)
x = float(elem.find('x').text)
y = float(elem.find('y').text)
r = float(elem.find('r').text)
return cls(order, x, y, r, filter_id=filter_id)
class ZernikeRadialFilter(ZernikeFilter):
r"""Score the :math:`m = 0` (radial variation only) Zernike moments up to
@ -499,7 +531,7 @@ class ZernikeRadialFilter(ZernikeFilter):
x-coordinate of center of circle for normalization
y : float
y-coordinate of center of circle for normalization
r : int or None
r : float
Radius of circle for normalization
Attributes
@ -510,7 +542,7 @@ class ZernikeRadialFilter(ZernikeFilter):
x-coordinate of center of circle for normalization
y : float
y-coordinate of center of circle for normalization
r : int or None
r : float
Radius of circle for normalization
id : int
Unique identifier for the filter

View file

@ -204,10 +204,10 @@ class Model:
@classmethod
def from_xml(cls, geometry='geometry.xml', materials='materials.xml',
settings='settings.xml'):
settings='settings.xml', tallies='tallies.xml'):
"""Create model from existing XML files
When initializing this way, the user must manually load plots and
tallies.
When initializing this way, the user must manually load plots.
Parameters
----------
@ -217,6 +217,10 @@ class Model:
Path to materials.xml file
settings : str
Path to settings.xml file
tallies : str
Path to tallies.xml file
.. versionadded:: 0.13.0
Returns
-------
@ -227,7 +231,8 @@ class Model:
materials = openmc.Materials.from_xml(materials)
geometry = openmc.Geometry.from_xml(geometry, materials)
settings = openmc.Settings.from_xml(settings)
return cls(geometry, materials, settings)
tallies = openmc.Tallies.from_xml(tallies)
return cls(geometry, materials, settings, tallies)
def init_lib(self, threads=None, geometry_debug=False, restart_file=None,
tracks=False, output=True, event_based=None, intracomm=None):

View file

@ -10,12 +10,14 @@ from xml.etree import ElementTree as ET
import h5py
import numpy as np
import pandas as pd
from scipy.misc import derivative
import scipy.sparse as sps
import openmc
import openmc.checkvalue as cv
from ._xml import clean_indentation, reorder_attributes
from ._xml import clean_indentation, reorder_attributes, get_text
from .mixin import IDManagerMixin
from .mesh import MeshBase
# The tally arithmetic product types. The tensor product performs the full
@ -844,13 +846,8 @@ class Tally(IDManagerMixin):
'not contain any scores'
raise ValueError(msg)
else:
scores = ''
for score in self.scores:
scores += f'{score} '
subelement = ET.SubElement(element, "scores")
subelement.text = scores.rstrip(' ')
subelement = ET.SubElement(element, "scores")
subelement.text = ' '.join(str(x) for x in self.scores)
# Tally estimator type
if self.estimator is not None:
@ -859,7 +856,7 @@ class Tally(IDManagerMixin):
# Optional Triggers
for trigger in self.triggers:
trigger.get_trigger_xml(element)
element.append(trigger.to_xml_element())
# Optional derivatives
if self.derivative is not None:
@ -868,6 +865,60 @@ class Tally(IDManagerMixin):
return element
@classmethod
def from_xml_element(cls, elem, **kwargs):
"""Generate tally object from an XML element
Parameters
----------
elem : xml.etree.ElementTree.Element
XML element
Returns
-------
openmc.Tally
Tally object
"""
tally_id = int(elem.get('id'))
name = elem.get('name', '')
tally = cls(tally_id=tally_id, name=name)
# Read filters
filters_elem = elem.find('filters')
if filters_elem is not None:
filter_ids = [int(x) for x in filters_elem.text.split()]
tally.filters = [kwargs['filters'][uid] for uid in filter_ids]
# Read nuclides
nuclides_elem = elem.find('nuclides')
if nuclides_elem is not None:
tally.nuclides = nuclides_elem.text.split()
# Read scores
scores_elem = elem.find('scores')
if scores_elem is not None:
tally.scores = scores_elem.text.split()
# Set estimator
estimator_elem = elem.find('estimator')
if estimator_elem is not None:
tally.estimator = estimator_elem.text
# Read triggers
tally.triggers = [
openmc.Trigger.from_xml_element(trigger_elem)
for trigger_elem in elem.findall('trigger')
]
# Read tally derivative
deriv_elem = elem.find('derivative')
if deriv_elem is not None:
deriv_id = int(deriv_elem.text)
tally.derivative = kwargs['derivatives'][deriv_id]
return tally
def contains_filter(self, filter_type):
"""Looks for a filter in the tally that matches a specified type
@ -3143,3 +3194,49 @@ class Tallies(cv.CheckedList):
reorder_attributes(root_element) # TODO: Remove when support is Python 3.8+
tree = ET.ElementTree(root_element)
tree.write(str(p), xml_declaration=True, encoding='utf-8')
@classmethod
def from_xml(cls, path='tallies.xml'):
"""Generate tallies from XML file
Parameters
----------
path : str, optional
Path to tallies XML file
Returns
-------
openmc.Tallies
Tallies object
"""
tree = ET.parse(path)
root = tree.getroot()
# Read mesh elements
meshes = {}
for elem in root.findall('mesh'):
mesh = MeshBase.from_xml_element(elem)
meshes[mesh.id] = mesh
# Read filter elements
filters = {}
for elem in root.findall('filter'):
filter = openmc.Filter.from_xml_element(elem, meshes=meshes)
filters[filter.id] = filter
# Read derivative elements
derivatives = {}
for elem in root.findall('derivative'):
deriv = openmc.TallyDerivative.from_xml_element(elem)
derivatives[deriv.id] = deriv
# Read tally elements
tallies = []
for elem in root.findall('tally'):
tally = openmc.Tally.from_xml_element(
elem, filters=filters, derivatives=derivatives
)
tallies.append(tally)
return cls(tallies)

View file

@ -105,3 +105,24 @@ class TallyDerivative(EqualityMixin, IDManagerMixin):
if self.variable == 'nuclide_density':
element.set("nuclide", self.nuclide)
return element
@classmethod
def from_xml_element(cls, elem):
"""Generate tally derivative from an XML element
Parameters
----------
elem : xml.etree.ElementTree.Element
XML element
Returns
-------
openmc.TallyDerivative
Tally derivative object
"""
derivative_id = int(elem.get("id"))
variable = elem.get("variable")
material = int(elem.get("material"))
nuclide = elem.get("nuclide") if variable == "nuclide_density" else None
return cls(derivative_id, variable, material, nuclide)

View file

@ -74,7 +74,7 @@ class Trigger(EqualityMixin):
if score not in self._scores:
self._scores.append(score)
def get_trigger_xml(self, element):
def to_xml_element(self):
"""Return XML representation of the trigger
Returns
@ -84,8 +84,36 @@ class Trigger(EqualityMixin):
"""
subelement = ET.SubElement(element, "trigger")
subelement.set("type", self._trigger_type)
subelement.set("threshold", str(self._threshold))
element = ET.Element("trigger")
element.set("type", self._trigger_type)
element.set("threshold", str(self._threshold))
if len(self._scores) != 0:
subelement.set("scores", ' '.join(map(str, self._scores)))
element.set("scores", ' '.join(map(str, self._scores)))
return element
@classmethod
def from_xml_element(cls, elem):
"""Generate trigger object from an XML element
Parameters
----------
elem : xml.etree.ElementTree.Element
XML element
Returns
-------
openmc.Trigger
Trigger object
"""
# Generate trigger object
trigger_type = elem.get("type")
threshold = float(elem.get("threshold"))
trigger = cls(trigger_type, threshold)
# Add scores if present
scores = elem.get("scores")
if scores is not None:
trigger.scores = scores.split()
return trigger