From 6bd33602cc1f81564227497759a71d939edcfde7 Mon Sep 17 00:00:00 2001 From: Paul Romano Date: Thu, 20 Jan 2022 17:51:16 -0600 Subject: [PATCH] Implement Tallies.from_xml classmethod --- openmc/filter.py | 102 +++++++++++++++++++++++--------- openmc/filter_expansion.py | 40 +++++++++++-- openmc/model/model.py | 13 +++-- openmc/tallies.py | 115 ++++++++++++++++++++++++++++++++++--- openmc/tally_derivative.py | 21 +++++++ openmc/trigger.py | 38 ++++++++++-- 6 files changed, 281 insertions(+), 48 deletions(-) diff --git a/openmc/filter.py b/openmc/filter.py index 699fa1aef9..fbbc855c28 100644 --- a/openmc/filter.py +++ b/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 diff --git a/openmc/filter_expansion.py b/openmc/filter_expansion.py index 9a915d18fa..1c07f58b71 100644 --- a/openmc/filter_expansion.py +++ b/openmc/filter_expansion.py @@ -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 diff --git a/openmc/model/model.py b/openmc/model/model.py index c03c151c6f..5c9fefda7d 100644 --- a/openmc/model/model.py +++ b/openmc/model/model.py @@ -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): diff --git a/openmc/tallies.py b/openmc/tallies.py index d99eb41ff1..6755e805d0 100644 --- a/openmc/tallies.py +++ b/openmc/tallies.py @@ -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) diff --git a/openmc/tally_derivative.py b/openmc/tally_derivative.py index 125946197b..05a27681d5 100644 --- a/openmc/tally_derivative.py +++ b/openmc/tally_derivative.py @@ -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) diff --git a/openmc/trigger.py b/openmc/trigger.py index 3b00ccd204..550f7ee312 100644 --- a/openmc/trigger.py +++ b/openmc/trigger.py @@ -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