mirror of
https://github.com/openmc-dev/openmc.git
synced 2026-07-28 22:26:08 -04:00
Implement Tallies.from_xml classmethod
This commit is contained in:
parent
5dfc880d91
commit
6bd33602cc
6 changed files with 281 additions and 48 deletions
102
openmc/filter.py
102
openmc/filter.py
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue