diff --git a/openmc/_xml.py b/openmc/_xml.py index 05412128d..758d80525 100644 --- a/openmc/_xml.py +++ b/openmc/_xml.py @@ -64,8 +64,8 @@ def get_text(elem, name, default=None): -def get_elem_tuple(elem, name, dtype=int): - """Helper function to get a tuple of values from an elem +def get_elem_list(elem, name, dtype=int): + """Helper function to get a list of values from an elem Parameters ---------- @@ -78,9 +78,9 @@ def get_elem_tuple(elem, name, dtype=int): Returns ------- - tuple of dtype - Data read from the tuple + list of dtype + Data read from the list """ - subelem = elem.find(name) - if subelem is not None: - return tuple([dtype(x) for x in subelem.text.split()]) + text = get_text(elem, name) + if text is not None: + return [dtype(x) for x in text.split()] diff --git a/openmc/cell.py b/openmc/cell.py index cd0573e8b..672afe095 100644 --- a/openmc/cell.py +++ b/openmc/cell.py @@ -8,7 +8,7 @@ from uncertainties import UFloat import openmc import openmc.checkvalue as cv -from ._xml import get_text +from ._xml import get_elem_list, get_text from .mixin import IDManagerMixin from .plots import add_plot_params from .region import Region, Complement @@ -689,9 +689,8 @@ class Cell(IDManagerMixin): c = cls(cell_id, name) # Assign material/distributed materials or fill - mat_text = get_text(elem, 'material') - if mat_text is not None: - mat_ids = mat_text.split() + mat_ids = get_elem_list(elem, 'material', str) + if mat_ids is not None: if len(mat_ids) > 1: c.fill = [materials[i] for i in mat_ids] else: @@ -706,19 +705,18 @@ class Cell(IDManagerMixin): c.region = Region.from_expression(region, surfaces) # Check for other attributes - t = get_text(elem, 'temperature') - if t is not None: - if ' ' in t: - c.temperature = [float(t_i) for t_i in t.split()] + temperature = get_elem_list(elem, 'temperature', float) + if temperature is not None: + if len(temperature) > 1: + c.temperature = temperature else: - c.temperature = float(t) + c.temperature = temperature[0] v = get_text(elem, 'volume') if v is not None: c.volume = float(v) for key in ('temperature', 'rotation', 'translation'): - value = get_text(elem, key) - if value is not None: - values = [float(x) for x in value.split()] + values = get_elem_list(elem, key, float) + if values is not None: if key == 'rotation' and len(values) == 9: values = np.array(values).reshape(3, 3) setattr(c, key, values) diff --git a/openmc/dagmc.py b/openmc/dagmc.py index 2486f7516..d1265be26 100644 --- a/openmc/dagmc.py +++ b/openmc/dagmc.py @@ -8,7 +8,7 @@ import warnings import openmc import openmc.checkvalue as cv -from ._xml import get_text +from ._xml import get_elem_list, get_text from .checkvalue import check_type, check_value from .surface import _BOUNDARY_TYPES from .bounding_box import BoundingBox @@ -468,8 +468,8 @@ class DAGMCUniverse(openmc.UniverseBase): if name is not None: out.name = name - out.auto_geom_ids = bool(elem.get('auto_geom_ids')) - out.auto_mat_ids = bool(elem.get('auto_mat_ids')) + out.auto_geom_ids = bool(get_text(elem, "auto_geom_ids")) + out.auto_mat_ids = bool(get_text(elem, "auto_mat_ids")) el_mat_override = elem.find('material_overrides') if el_mat_override is not None: @@ -480,7 +480,7 @@ class DAGMCUniverse(openmc.UniverseBase): out._material_overrides = {} for elem in el_mat_override.findall('cell_override'): cell_id = int(get_text(elem, 'id')) - mat_ids = get_text(elem, 'material_ids').split(' ') + mat_ids = get_elem_list(elem, "material_ids", str) or [] mat_objs = [mats[mat_id] for mat_id in mat_ids] out._material_overrides[cell_id] = mat_objs diff --git a/openmc/data/library.py b/openmc/data/library.py index bec538c06..b49757b0d 100644 --- a/openmc/data/library.py +++ b/openmc/data/library.py @@ -5,7 +5,7 @@ import h5py import lxml.etree as ET import openmc -from openmc._xml import clean_indentation +from openmc._xml import get_elem_list, get_text, clean_indentation class DataLibrary(list): @@ -172,9 +172,9 @@ class DataLibrary(list): directory = os.path.dirname(path) for lib_element in root.findall('library'): - filename = os.path.join(directory, lib_element.attrib['path']) - filetype = lib_element.attrib['type'] - materials = lib_element.attrib['materials'].split() + filename = os.path.join(directory, get_text(lib_element, "path")) + filetype = get_text(lib_element, "type") + materials = get_elem_list(lib_element, "materials", str) or [] library = {'path': filename, 'type': filetype, 'materials': materials} data.libraries.append(library) @@ -182,7 +182,7 @@ class DataLibrary(list): # get depletion chain data dep_node = root.find("depletion_chain") if dep_node is not None: - filename = os.path.join(directory, dep_node.attrib['path']) + filename = os.path.join(directory, get_text(dep_node, "path")) library = {'path': filename, 'type': 'depletion_chain', 'materials': []} data.libraries.append(library) diff --git a/openmc/deplete/chain.py b/openmc/deplete/chain.py index 8e24f716b..f1a23317f 100644 --- a/openmc/deplete/chain.py +++ b/openmc/deplete/chain.py @@ -22,6 +22,7 @@ from openmc.checkvalue import check_type, check_greater_than, PathLike from openmc.data import gnds_name, zam from openmc.exceptions import DataError from .nuclide import FissionYieldDistribution, Nuclide +from .._xml import get_text import openmc.data @@ -553,7 +554,7 @@ class Chain: root = ET.parse(str(filename)) for i, nuclide_elem in enumerate(root.findall('nuclide')): - this_q = fission_q.get(nuclide_elem.get("name")) + this_q = fission_q.get(get_text(nuclide_elem, "name")) nuc = Nuclide.from_xml(nuclide_elem, root, this_q) chain.add_nuclide(nuc) diff --git a/openmc/deplete/nuclide.py b/openmc/deplete/nuclide.py index 60e3e5317..958814834 100644 --- a/openmc/deplete/nuclide.py +++ b/openmc/deplete/nuclide.py @@ -14,6 +14,7 @@ import numpy as np from openmc.checkvalue import check_type from openmc.stats import Univariate +from .._xml import get_elem_list, get_text __all__ = [ "DecayTuple", "ReactionTuple", "Nuclide", "FissionYield", @@ -225,38 +226,39 @@ class Nuclide: """ nuc = cls() - nuc.name = element.get('name') + nuc.name = get_text(element, "name") # Check for half-life - if 'half_life' in element.attrib: - nuc.half_life = float(element.get('half_life')) - nuc.decay_energy = float(element.get('decay_energy', '0')) + half_life = get_text(element, "half_life") + if half_life is not None: + nuc.half_life = float(half_life) + nuc.decay_energy = float(get_text(element, "decay_energy", 0.0)) # Check for decay paths for decay_elem in element.iter('decay'): - d_type = decay_elem.get('type') - target = decay_elem.get('target') + d_type = get_text(decay_elem, "type") + target = get_text(decay_elem, "target") if target is not None and target.lower() == "nothing": target = None - branching_ratio = float(decay_elem.get('branching_ratio')) + branching_ratio = float(get_text(decay_elem, "branching_ratio")) nuc.decay_modes.append(DecayTuple(d_type, target, branching_ratio)) # Check for sources for src_elem in element.iter('source'): - particle = src_elem.get('particle') + particle = get_text(src_elem, "particle") distribution = Univariate.from_xml_element(src_elem) nuc.sources[particle] = distribution # Check for reaction paths for reaction_elem in element.iter('reaction'): - r_type = reaction_elem.get('type') - Q = float(reaction_elem.get('Q', '0')) - branching_ratio = float(reaction_elem.get('branching_ratio', '1')) + r_type = get_text(reaction_elem, "type") + Q = float(get_text(reaction_elem, "Q", 0.0)) + branching_ratio = float(get_text(reaction_elem, "branching_ratio", 1.0)) # If the type is not fission, get target and Q value, otherwise # just set null values if r_type != 'fission': - target = reaction_elem.get('target') + target = get_text(reaction_elem, "target") if target is not None and target.lower() == "nothing": target = None else: @@ -271,7 +273,7 @@ class Nuclide: fpy_elem = element.find('neutron_fission_yields') if fpy_elem is not None: # Check for use of FPY from other nuclide - parent = fpy_elem.get('parent') + parent = get_text(fpy_elem, "parent") if parent is not None: assert root is not None fpy_elem = root.find( @@ -529,9 +531,9 @@ class FissionYieldDistribution(Mapping): """ all_yields = {} for yield_elem in element.iter("fission_yields"): - energy = float(yield_elem.get("energy")) - products = yield_elem.find("products").text.split() - yields = map(float, yield_elem.find("data").text.split()) + energy = float(get_text(yield_elem, "energy")) + products = get_elem_list(yield_elem, "products", str) or [] + yields = get_elem_list(yield_elem, "data", float) or [] # Get a map of products to their corresponding yield all_yields[energy] = dict(zip(products, yields)) diff --git a/openmc/filter.py b/openmc/filter.py index 53fec898f..6a666d2a0 100644 --- a/openmc/filter.py +++ b/openmc/filter.py @@ -17,7 +17,7 @@ from .material import Material from .mixin import IDManagerMixin from .surface import Surface from .universe import UniverseBase -from ._xml import get_text +from ._xml import get_elem_list, get_text _FILTER_TYPES = ( @@ -259,17 +259,15 @@ class Filter(IDManagerMixin, metaclass=FilterMeta): Filter object """ - filter_type = elem.get('type') - if filter_type is None: - filter_type = elem.find('type').text + filter_type = get_text(elem, "type") # If the filter type matches this class's short_name, then # there is no overridden 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()] + filter_id = int(get_text(elem, "id")) + bins = get_elem_list(elem, "bins", int) or [] return cls(bins, filter_id=filter_id) # Search through all subclasses and find the one matching the HDF5 @@ -701,8 +699,8 @@ class CellInstanceFilter(Filter): @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()] + filter_id = int(get_text(elem, "id")) + bins = get_elem_list(elem, "bins", int) or [] cell_instances = list(zip(bins[::2], bins[1::2])) return cls(cell_instances, filter_id=filter_id) @@ -784,8 +782,8 @@ class ParticleFilter(Filter): @classmethod def from_xml_element(cls, elem, **kwargs): - filter_id = int(elem.get('id')) - bins = get_text(elem, 'bins').split() + filter_id = int(get_text(elem, "id")) + bins = get_elem_list(elem, "bins", str) or [] return cls(bins, filter_id=filter_id) @@ -1004,12 +1002,12 @@ class MeshFilter(Filter): def from_xml_element(cls, elem: ET.Element, **kwargs) -> MeshFilter: mesh_id = int(get_text(elem, 'bins')) mesh_obj = kwargs['meshes'][mesh_id] - filter_id = int(elem.get('id')) + filter_id = int(get_text(elem, "id")) out = cls(mesh_obj, filter_id=filter_id) - translation = elem.get('translation') + translation = get_elem_list(elem, "translation", float) or [] if translation: - out.translation = [float(x) for x in translation.split()] + out.translation = translation return out @@ -1149,16 +1147,16 @@ class MeshMaterialFilter(MeshFilter): @classmethod def from_xml_element(cls, elem: ET.Element, **kwargs) -> MeshMaterialFilter: - filter_id = int(elem.get('id')) - mesh_id = int(elem.get('mesh')) + filter_id = int(get_text(elem, "id")) + mesh_id = int(get_text(elem, "mesh")) mesh_obj = kwargs['meshes'][mesh_id] - bins = [int(x) for x in get_text(elem, 'bins').split()] + bins = get_elem_list(elem, "bins", int) or [] bins = list(zip(bins[::2], bins[1::2])) out = cls(mesh_obj, bins, filter_id=filter_id) - translation = elem.get('translation') + translation = get_elem_list(elem, "translation", float) or [] if translation: - out.translation = [float(x) for x in translation.split()] + out.translation = translation return out @classmethod @@ -1557,8 +1555,8 @@ class RealFilter(Filter): @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()] + filter_id = int(get_text(elem, "id")) + bins = get_elem_list(elem, "bins", float) or [] return cls(bins, filter_id=filter_id) @@ -2447,12 +2445,13 @@ class EnergyFunctionFilter(Filter): @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()] + filter_id = int(get_text(elem, "id")) + energy = get_elem_list(elem, "energy", float) or [] + y = get_elem_list(elem, "y", float) or [] out = cls(energy, y, filter_id=filter_id) - if elem.find('interpolation') is not None: - out.interpolation = elem.find('interpolation').text + interpolation = get_text(elem, "interpolation") + if interpolation is not None: + out.interpolation = interpolation return out def can_merge(self, other): diff --git a/openmc/filter_expansion.py b/openmc/filter_expansion.py index cdb2f20e0..b79c8fc79 100644 --- a/openmc/filter_expansion.py +++ b/openmc/filter_expansion.py @@ -4,6 +4,7 @@ import lxml.etree as ET import openmc.checkvalue as cv from .filter import Filter +from ._xml import get_text class ExpansionFilter(Filter): @@ -49,8 +50,8 @@ class ExpansionFilter(Filter): @classmethod def from_xml_element(cls, elem, **kwargs): - filter_id = int(elem.get('id')) - order = int(elem.find('order').text) + filter_id = int(get_text(elem, "id")) + order = int(get_text(elem, "order")) return cls(order, filter_id=filter_id) def merge(self, other): @@ -263,11 +264,11 @@ class SpatialLegendreFilter(ExpansionFilter): @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) + filter_id = int(get_text(elem, "id")) + order = int(get_text(elem, "order")) + axis = get_text(elem, "axis") + minimum = float(get_text(elem, "min")) + maximum = float(get_text(elem, "max")) return cls(order, axis, minimum, maximum, filter_id=filter_id) @@ -362,10 +363,10 @@ class SphericalHarmonicsFilter(ExpansionFilter): @classmethod def from_xml_element(cls, elem, **kwargs): - filter_id = int(elem.get('id')) - order = int(elem.find('order').text) + filter_id = int(get_text(elem, "id")) + order = int(get_text(elem, "order")) filter = cls(order, filter_id=filter_id) - filter.cosine = elem.get('cosine') + filter.cosine = get_text(elem, "cosine") return filter @@ -518,11 +519,11 @@ class ZernikeFilter(ExpansionFilter): @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) + filter_id = int(get_text(elem, "id")) + order = int(get_text(elem, "order")) + x = float(get_text(elem, "x")) + y = float(get_text(elem, "y")) + r = float(get_text(elem, "r")) return cls(order, x, y, r, filter_id=filter_id) diff --git a/openmc/lattice.py b/openmc/lattice.py index 518068560..f0e8b40b0 100644 --- a/openmc/lattice.py +++ b/openmc/lattice.py @@ -10,7 +10,7 @@ import numpy as np import openmc import openmc.checkvalue as cv -from ._xml import get_text +from ._xml import get_elem_list, get_text from .mixin import IDManagerMixin @@ -959,18 +959,17 @@ class RectLattice(Lattice): lat_id = int(get_text(elem, 'id')) name = get_text(elem, 'name') lat = cls(lat_id, name) - lat.lower_left = [float(i) - for i in get_text(elem, 'lower_left').split()] - lat.pitch = [float(i) for i in get_text(elem, 'pitch').split()] + lat.lower_left = get_elem_list(elem, "lower_left", float) + lat.pitch = get_elem_list(elem, "pitch", float) outer = get_text(elem, 'outer') if outer is not None: lat.outer = get_universe(int(outer)) # Get array of universes - dimension = get_text(elem, 'dimension').split() + dimension = get_elem_list(elem, 'dimension', int) shape = np.array(dimension, dtype=int)[::-1] - uarray = np.array([get_universe(int(i)) for i in - get_text(elem, 'universes').split()]) + universes = get_elem_list(elem, 'universes', int) + uarray = np.array([get_universe(u) for u in universes]) uarray.shape = shape lat.universes = uarray return lat @@ -1530,8 +1529,8 @@ class HexLattice(Lattice): lat_id = int(get_text(elem, 'id')) name = get_text(elem, 'name') lat = cls(lat_id, name) - lat.center = [float(i) for i in get_text(elem, 'center').split()] - lat.pitch = [float(i) for i in get_text(elem, 'pitch').split()] + lat.center = get_elem_list(elem, "center", float) + lat.pitch = get_elem_list(elem, "pitch", float) lat.orientation = get_text(elem, 'orientation', 'y') outer = get_text(elem, 'outer') if outer is not None: @@ -1548,8 +1547,8 @@ class HexLattice(Lattice): univs = [deepcopy(univs) for i in range(n_axial)] # Get flat array of universes - uarray = np.array([get_universe(int(i)) for i in - get_text(elem, 'universes').split()]) + universes = get_elem_list(elem, "universes", int) + uarray = np.array([get_universe(u) for u in universes]) # Fill nested lists j = 0 diff --git a/openmc/material.py b/openmc/material.py index db67b709c..0afe5b670 100644 --- a/openmc/material.py +++ b/openmc/material.py @@ -17,7 +17,7 @@ import h5py import openmc import openmc.data import openmc.checkvalue as cv -from ._xml import clean_indentation +from ._xml import clean_indentation, get_elem_list, get_text from .mixin import IDManagerMixin from .utility_funcs import input_path from . import waste @@ -1672,50 +1672,55 @@ class Material(IDManagerMixin): Material generated from XML element """ - mat_id = int(elem.get('id')) + mat_id = int(get_text(elem, 'id')) + # Add NCrystal material from cfg string - if "cfg" in elem.attrib: - cfg = elem.get("cfg") + cfg = get_text(elem, "cfg") + if cfg is not None: return Material.from_ncrystal(cfg, material_id=mat_id) mat = cls(mat_id) - mat.name = elem.get('name') + mat.name = get_text(elem, 'name') - if "temperature" in elem.attrib: - mat.temperature = float(elem.get("temperature")) + temperature = get_text(elem, "temperature") + if temperature is not None: + mat.temperature = float(temperature) - if 'volume' in elem.attrib: - mat.volume = float(elem.get('volume')) + volume = get_text(elem, "volume") + if volume is not None: + mat.volume = float(volume) # Get each nuclide for nuclide in elem.findall('nuclide'): - name = nuclide.attrib['name'] + name = get_text(nuclide, "name") if 'ao' in nuclide.attrib: mat.add_nuclide(name, float(nuclide.attrib['ao'])) elif 'wo' in nuclide.attrib: mat.add_nuclide(name, float(nuclide.attrib['wo']), 'wo') # Get depletable attribute - mat.depletable = elem.get('depletable') in ('true', '1') + depletable = get_text(elem, "depletable") + mat.depletable = depletable in ('true', '1') # Get each S(a,b) table for sab in elem.findall('sab'): - fraction = float(sab.get('fraction', 1.0)) - mat.add_s_alpha_beta(sab.get('name'), fraction) + fraction = float(get_text(sab, "fraction", 1.0)) + name = get_text(sab, "name") + mat.add_s_alpha_beta(name, fraction) # Get total material density density = elem.find('density') - units = density.get('units') + units = get_text(density, "units") if units == 'sum': mat.set_density(units) else: - value = float(density.get('value')) + value = float(get_text(density, 'value')) mat.set_density(units, value) # Check for isotropic scattering nuclides - isotropic = elem.find('isotropic') + isotropic = get_elem_list(elem, "isotropic", str) if isotropic is not None: - mat.isotropic = isotropic.text.split() + mat.isotropic = isotropic return mat @@ -1982,9 +1987,9 @@ class Materials(cv.CheckedList): materials.append(Material.from_xml_element(material)) # Check for cross sections settings - xs = elem.find('cross_sections') + xs = get_text(elem, "cross_sections") if xs is not None: - materials.cross_sections = xs.text + materials.cross_sections = xs return materials diff --git a/openmc/mesh.py b/openmc/mesh.py index 339702884..2e9abd1b6 100644 --- a/openmc/mesh.py +++ b/openmc/mesh.py @@ -16,7 +16,7 @@ import openmc import openmc.checkvalue as cv from openmc.checkvalue import PathLike from openmc.utility_funcs import change_directory -from ._xml import get_text +from ._xml import get_elem_list, get_text from .mixin import IDManagerMixin from .surface import _BOUNDARY_TYPES from .utility_funcs import input_path @@ -1187,21 +1187,21 @@ class RegularMesh(StructuredMesh): mesh_id = int(get_text(elem, 'id')) mesh = cls(mesh_id=mesh_id) - dimension = get_text(elem, 'dimension') + dimension = get_elem_list(elem, "dimension", int) if dimension is not None: - mesh.dimension = [int(x) for x in dimension.split()] + mesh.dimension = dimension - lower_left = get_text(elem, 'lower_left') + lower_left = get_elem_list(elem, "lower_left", float) if lower_left is not None: - mesh.lower_left = [float(x) for x in lower_left.split()] + mesh.lower_left = lower_left - upper_right = get_text(elem, 'upper_right') + upper_right = get_elem_list(elem, "upper_right", float) if upper_right is not None: - mesh.upper_right = [float(x) for x in upper_right.split()] + mesh.upper_right = upper_right - width = get_text(elem, 'width') + width = get_elem_list(elem, "width", float) if width is not None: - mesh.width = [float(x) for x in width.split()] + mesh.width = width return mesh @@ -1507,9 +1507,9 @@ class RectilinearMesh(StructuredMesh): """ mesh_id = int(get_text(elem, 'id')) mesh = cls(mesh_id=mesh_id) - mesh.x_grid = [float(x) for x in get_text(elem, 'x_grid').split()] - mesh.y_grid = [float(y) for y in get_text(elem, 'y_grid').split()] - mesh.z_grid = [float(z) for z in get_text(elem, 'z_grid').split()] + mesh.x_grid = get_elem_list(elem, "x_grid", float) + mesh.y_grid = get_elem_list(elem, "y_grid", float) + mesh.z_grid = get_elem_list(elem, "z_grid", float) return mesh @@ -1923,10 +1923,10 @@ class CylindricalMesh(StructuredMesh): mesh_id = int(get_text(elem, 'id')) mesh = cls( - r_grid = [float(x) for x in get_text(elem, "r_grid").split()], - phi_grid = [float(x) for x in get_text(elem, "phi_grid").split()], - z_grid = [float(x) for x in get_text(elem, "z_grid").split()], - origin = [float(x) for x in get_text(elem, "origin", default=[0., 0., 0.]).split()], + r_grid = get_elem_list(elem, "r_grid", float), + phi_grid = get_elem_list(elem, "phi_grid", float), + z_grid = get_elem_list(elem, "z_grid", float), + origin = get_elem_list(elem, "origin", float) or [0., 0., 0.], mesh_id=mesh_id, ) @@ -2296,10 +2296,10 @@ class SphericalMesh(StructuredMesh): mesh_id = int(get_text(elem, 'id')) mesh = cls( mesh_id=mesh_id, - r_grid = [float(x) for x in get_text(elem, "r_grid").split()], - theta_grid = [float(x) for x in get_text(elem, "theta_grid").split()], - phi_grid = [float(x) for x in get_text(elem, "phi_grid").split()], - origin = [float(x) for x in get_text(elem, "origin", default=[0., 0., 0.]).split()], + r_grid = get_elem_list(elem, "r_grid", float), + theta_grid = get_elem_list(elem, "theta_grid", float), + phi_grid = get_elem_list(elem, "phi_grid", float), + origin = get_elem_list(elem, "origin", float) or [0., 0., 0.], ) return mesh @@ -2842,7 +2842,7 @@ class UnstructuredMesh(MeshBase): filename = get_text(elem, 'filename') library = get_text(elem, 'library') length_multiplier = float(get_text(elem, 'length_multiplier', 1.0)) - options = elem.get('options') + options = get_text(elem, "options") return cls(filename, library, mesh_id, '', length_multiplier, options) diff --git a/openmc/plots.py b/openmc/plots.py index 34dde84e4..072a9a319 100644 --- a/openmc/plots.py +++ b/openmc/plots.py @@ -10,7 +10,7 @@ import openmc import openmc.checkvalue as cv from openmc.checkvalue import PathLike -from ._xml import clean_indentation, get_elem_tuple, get_text +from ._xml import clean_indentation, get_elem_list, get_text from .mixin import IDManagerMixin _BASES = {'xy', 'xz', 'yz'} @@ -944,64 +944,61 @@ class Plot(PlotBase): Plot object """ - plot_id = int(elem.get("id")) + plot_id = int(get_text(elem, "id")) name = get_text(elem, 'name', '') plot = cls(plot_id, name) if "filename" in elem.keys(): - plot.filename = elem.get("filename") - plot.color_by = elem.get("color_by") - plot.type = elem.get("type") + plot.filename = get_text(elem, "filename") + plot.color_by = get_text(elem, "color_by") + plot.type = get_text(elem, "type") if plot.type == 'slice': - plot.basis = elem.get("basis") + plot.basis = get_text(elem, "basis") - plot.origin = get_elem_tuple(elem, "origin", float) - plot.width = get_elem_tuple(elem, "width", float) - plot.pixels = get_elem_tuple(elem, "pixels") - plot._background = get_elem_tuple(elem, "background") + plot.origin = tuple(get_elem_list(elem, "origin", float)) + plot.width = tuple(get_elem_list(elem, "width", float)) + plot.pixels = tuple(get_elem_list(elem, "pixels")) + background = get_elem_list(elem, "background") + if background is not None: + plot._background = tuple(background) # Set plot colors colors = {} for color_elem in elem.findall("color"): - uid = int(color_elem.get("id")) - colors[uid] = tuple([int(x) - for x in color_elem.get("rgb").split()]) + uid = int(get_text(color_elem, "id")) + colors[uid] = tuple(get_elem_list(color_elem, "rgb", int)) plot.colors = colors # Set masking information mask_elem = elem.find("mask") if mask_elem is not None: - plot.mask_components = [ - int(x) for x in mask_elem.get("components").split()] - background = mask_elem.get("background") + plot.mask_components = get_elem_list(mask_elem, "components", int) + background = get_elem_list(mask_elem, "background", int) if background is not None: - plot.mask_background = tuple( - [int(x) for x in background.split()]) + plot.mask_background = tuple(background) # show overlaps - overlap_elem = elem.find("show_overlaps") - if overlap_elem is not None: - plot.show_overlaps = (overlap_elem.text in ('true', '1')) - overlap_color = get_elem_tuple(elem, "overlap_color") + overlap = get_text(elem, "show_overlaps") + if overlap is not None: + plot.show_overlaps = (overlap in ('true', '1')) + overlap_color = get_elem_list(elem, "overlap_color", int) if overlap_color is not None: - plot.overlap_color = overlap_color + plot.overlap_color = tuple(overlap_color) # Set universe level - level = elem.find("level") + level = get_text(elem, "level") if level is not None: - plot.level = int(level.text) + plot.level = int(level) # Set meshlines mesh_elem = elem.find("meshlines") if mesh_elem is not None: - meshlines = {'type': mesh_elem.get('meshtype')} + meshlines = {'type': get_text(mesh_elem, "meshtype")} if 'id' in mesh_elem.keys(): - meshlines['id'] = int(mesh_elem.get('id')) + meshlines['id'] = int(get_text(mesh_elem, "id")) if 'linewidth' in mesh_elem.keys(): - meshlines['linewidth'] = int(mesh_elem.get('linewidth')) + meshlines['linewidth'] = int(get_text(mesh_elem, "linewidth")) if 'color' in mesh_elem.keys(): - meshlines['color'] = tuple( - [int(x) for x in mesh_elem.get('color').split()] - ) + meshlines['color'] = tuple(get_elem_list(mesh_elem, "color", int)) plot.meshlines = meshlines return plot @@ -1259,38 +1256,39 @@ class RayTracePlot(PlotBase): None """ - if "filename" in elem.keys(): - self.filename = elem.get("filename") - self.color_by = elem.get("color_by") + filename = get_text(elem, "filename") + if filename is not None: + self.filename = filename + self.color_by = get_text(elem, "color_by") - horizontal_fov = elem.find("horizontal_field_of_view") + horizontal_fov = get_text(elem, "horizontal_field_of_view") if horizontal_fov is not None: - self.horizontal_field_of_view = float(horizontal_fov.text) + self.horizontal_field_of_view = float(horizontal_fov) - if (tmp := elem.find("orthographic_width")) is not None: - self.orthographic_width = float(tmp) + orthographic_width = get_text(elem, "orthographic_width") + if orthographic_width is not None: + self.orthographic_width = float(orthographic_width) - self.pixels = get_elem_tuple(elem, "pixels") - self.camera_position = get_elem_tuple(elem, "camera_position", float) - self.look_at = get_elem_tuple(elem, "look_at", float) + self.pixels = tuple(get_elem_list(elem, "pixels", int)) + self.camera_position = tuple(get_elem_list(elem, "camera_position", float)) + self.look_at = tuple(get_elem_list(elem, "look_at", float)) - if elem.find("background") is not None: - self.background = get_elem_tuple(elem, "background") + background = get_elem_list(elem, "background", int) + if background is not None: + self.background = tuple(background) # Set masking information if (mask_elem := elem.find("mask")) is not None: - mask_components = [int(x) - for x in mask_elem.get("components").split()] + mask_components = get_elem_list(mask_elem, "components", int) # TODO: set mask components(needs geometry information) - background = mask_elem.get("background") + background = get_elem_list(mask_elem, "background", int) if background is not None: - self.mask_background = tuple( - [int(x) for x in background.split()]) + self.mask_background = tuple(background) # Set universe level - level = elem.find("level") + level = get_text(elem, "level") if level is not None: - self.level = int(level.text) + self.level = int(level) class WireframeRayTracePlot(RayTracePlot): @@ -1515,7 +1513,7 @@ class WireframeRayTracePlot(RayTracePlot): """ - plot_id = int(elem.get("id")) + plot_id = int(get_text(elem, "id")) plot_name = get_text(elem, 'name', '') plot = cls(plot_id, plot_name) plot.type = "wireframe_raytrace" @@ -1523,19 +1521,18 @@ class WireframeRayTracePlot(RayTracePlot): plot._read_xml_attributes(elem) # Attempt to get wireframe thickness.May not be present - wireframe_thickness = elem.find("wireframe_thickness") + wireframe_thickness = get_text(elem, "wireframe_thickness") if wireframe_thickness is not None: - plot.wireframe_thickness = int(wireframe_thickness.text) - wireframe_color = elem.get("wireframe_color") + plot.wireframe_thickness = int(wireframe_thickness) + wireframe_color = get_elem_list(elem, "wireframe_color", int) if wireframe_color: - plot.wireframe_color = [int(item) for item in wireframe_color] + plot.wireframe_color = wireframe_color # Set plot colors for color_elem in elem.findall("color"): - uid = int(color_elem.get("id")) - plot.colors[uid] = tuple(int(i) - for i in get_text(color_elem, 'rgb').split()) - plot.xs[uid] = float(color_elem.get("xs")) + uid = int(get_text(color_elem, "id")) + plot.colors[uid] = tuple(get_elem_list(color_elem, "rgb", int)) + plot.xs[uid] = float(get_text(color_elem, "xs")) return plot @@ -1693,15 +1690,17 @@ class SolidRayTracePlot(RayTracePlot): def _read_phong_attributes(self, elem): """Read attributes specific to the Phong plot from an XML element""" - if elem.find('light_position') is not None: - self.light_position = get_elem_tuple(elem, 'light_position', float) + light_position = get_elem_list(elem, 'light_position', float) + if light_position is not None: + self.light_position = tuple(light_position) - diffuse_fraction = elem.find('diffuse_fraction') + diffuse_fraction = get_text(elem, "diffuse_fraction") if diffuse_fraction is not None: - self.diffuse_fraction = float(diffuse_fraction.text) + self.diffuse_fraction = float(diffuse_fraction) - if elem.find('opaque_ids') is not None: - self.opaque_domains = list(get_elem_tuple(elem, 'opaque_ids', int)) + opaque_domains = get_elem_list(elem, 'opaque_ids', int) + if opaque_domains is not None: + self.opaque_domains = opaque_domains @classmethod def from_xml_element(cls, elem): @@ -1719,7 +1718,7 @@ class SolidRayTracePlot(RayTracePlot): """ - plot_id = int(elem.get("id")) + plot_id = int(get_text(elem, "id")) plot_name = get_text(elem, 'name', '') plot = cls(plot_id, plot_name) plot.type = "solid_raytrace" @@ -1729,8 +1728,8 @@ class SolidRayTracePlot(RayTracePlot): # Set plot colors for color_elem in elem.findall("color"): - uid = color_elem.get("id") - plot.colors[uid] = get_elem_tuple(color_elem, "rgb") + uid = get_text(color_elem, "id") + plot.colors[uid] = tuple(get_elem_list(color_elem, "rgb", int)) return plot @@ -1897,7 +1896,7 @@ class Plots(cv.CheckedList): # Generate each plot plots = cls() for e in elem.findall('plot'): - plot_type = e.get('type') + plot_type = get_text(e, "type") if plot_type == 'wireframe_raytrace': plots.append(WireframeRayTracePlot.from_xml_element(e)) elif plot_type == 'solid_raytrace': diff --git a/openmc/settings.py b/openmc/settings.py index 327faf544..ce3743ff5 100644 --- a/openmc/settings.py +++ b/openmc/settings.py @@ -11,7 +11,7 @@ import openmc import openmc.checkvalue as cv from openmc.checkvalue import PathLike from openmc.stats.multivariate import MeshSpatial -from ._xml import clean_indentation, get_text +from ._xml import clean_indentation, get_elem_list, get_text from .mesh import _read_meshes, RegularMesh, MeshBase from .source import SourceBase, MeshSource, IndependentSource from .utility_funcs import input_path @@ -1791,23 +1791,21 @@ class Settings: def _statepoint_from_xml_element(self, root): elem = root.find('state_point') if elem is not None: - text = get_text(elem, 'batches') - if text is not None: - self.statepoint['batches'] = [int(x) for x in text.split()] + batches = get_elem_list(elem, "batches", int) + if batches is not None: + self.statepoint['batches'] = batches def _sourcepoint_from_xml_element(self, root): elem = root.find('source_point') if elem is not None: for key in ('separate', 'write', 'overwrite_latest', 'batches', 'mcpl'): - value = get_text(elem, key) - if value is not None: - if key in ('separate', 'write', 'mcpl'): - value = value in ('true', '1') - elif key == 'overwrite_latest': - value = value in ('true', '1') + if key in ('separate', 'write', 'mcpl', 'overwrite_latest'): + value = get_text(elem, key) in ('true', '1') + if key == 'overwrite_latest': key = 'overwrite' - else: - value = [int(x) for x in value.split()] + else: + value = get_elem_list(elem, key, int) + if value is not None: self.sourcepoint[key] = value def _surf_source_read_from_xml_element(self, root): @@ -1824,11 +1822,12 @@ class Settings: if elem is None: return for key in ('surface_ids', 'max_particles', 'max_source_files', 'mcpl', 'cell', 'cellto', 'cellfrom'): - value = get_text(elem, key) + if key == 'surface_ids': + value = get_elem_list(elem, key, int) + else: + value = get_text(elem, key) if value is not None: - if key == 'surface_ids': - value = [int(x) for x in value.split()] - elif key == 'mcpl': + if key == 'mcpl': value = value in ('true', '1') elif key in ('max_particles', 'max_source_files', 'cell', 'cellfrom', 'cellto'): value = int(value) @@ -1958,22 +1957,21 @@ class Settings: text = get_text(root, 'temperature_method') if text is not None: self.temperature['method'] = text - text = get_text(root, 'temperature_range') + text = get_elem_list(root, "temperature_range", float) if text is not None: - self.temperature['range'] = [float(x) for x in text.split()] + self.temperature['range'] = text text = get_text(root, 'temperature_multipole') if text is not None: self.temperature['multipole'] = text in ('true', '1') def _trace_from_xml_element(self, root): - text = get_text(root, 'trace') + text = get_elem_list(root, "trace", int) if text is not None: - self.trace = [int(x) for x in text.split()] + self.trace = text def _track_from_xml_element(self, root): - text = get_text(root, 'track') - if text is not None: - values = [int(x) for x in text.split()] + values = get_elem_list(root, "track", int) + if values is not None: self.track = list(zip(values[::3], values[1::3], values[2::3])) def _ufs_mesh_from_xml_element(self, root, meshes): @@ -1990,14 +1988,15 @@ class Settings: if elem is not None: keys = ('enable', 'method', 'energy_min', 'energy_max', 'nuclides') for key in keys: - value = get_text(elem, key) + if key == 'nuclides': + value = get_elem_list(elem, key, str) + else: + value = get_text(elem, key) if value is not None: if key == 'enable': value = value in ('true', '1') elif key in ('energy_min', 'energy_max'): value = float(value) - elif key == 'nuclides': - value = value.split() self.resonance_scattering[key] = value def _create_fission_neutrons_from_xml_element(self, root): @@ -2109,8 +2108,8 @@ class Settings: mesh = MeshBase.from_xml_element(mesh_elem) domains = [] for domain_elem in mesh_elem.findall('domain'): - domain_id = int(domain_elem.get('id')) - domain_type = domain_elem.get('type') + domain_id = int(get_text(domain_elem, "id")) + domain_type = get_text(domain_elem, "type") if domain_type == 'material': domain = openmc.Material(domain_id) elif domain_type == 'cell': diff --git a/openmc/source.py b/openmc/source.py index 87e734e9a..c463ccb27 100644 --- a/openmc/source.py +++ b/openmc/source.py @@ -18,7 +18,7 @@ import openmc.checkvalue as cv from openmc.checkvalue import PathLike from openmc.stats.multivariate import UnitSphere, Spatial from openmc.stats.univariate import Univariate -from ._xml import get_text +from ._xml import get_elem_list, get_text from .mesh import MeshBase, StructuredMesh, UnstructuredMesh from .utility_funcs import input_path @@ -210,7 +210,7 @@ class SourceBase(ABC): constraints = {} domain_type = get_text(elem, "domain_type") if domain_type is not None: - domain_ids = [int(x) for x in get_text(elem, "domain_ids").split()] + domain_ids = get_elem_list(elem, "domain_ids", int) # Instantiate some throw-away domains that are used by the # constructor to assign IDs @@ -224,13 +224,13 @@ class SourceBase(ABC): domains = [openmc.Universe(uid) for uid in domain_ids] constraints['domains'] = domains - time_bounds = get_text(elem, "time_bounds") + time_bounds = get_elem_list(elem, "time_bounds", float) if time_bounds is not None: - constraints['time_bounds'] = [float(x) for x in time_bounds.split()] + constraints['time_bounds'] = time_bounds - energy_bounds = get_text(elem, "energy_bounds") + energy_bounds = get_elem_list(elem, "energy_bounds", float) if energy_bounds is not None: - constraints['energy_bounds'] = [float(x) for x in energy_bounds.split()] + constraints['energy_bounds'] = energy_bounds fissionable = get_text(elem, "fissionable") if fissionable is not None: diff --git a/openmc/stats/multivariate.py b/openmc/stats/multivariate.py index cd474fa9b..222d2d18a 100644 --- a/openmc/stats/multivariate.py +++ b/openmc/stats/multivariate.py @@ -10,7 +10,7 @@ import numpy as np import openmc import openmc.checkvalue as cv -from .._xml import get_text +from .._xml import get_elem_list, get_text from ..mesh import MeshBase from .univariate import PowerLaw, Uniform, Univariate @@ -152,9 +152,9 @@ class PolarAzimuthal(UnitSphere): """ mu_phi = cls() - uvw = get_text(elem, 'reference_uvw') + uvw = get_elem_list(elem, "reference_uvw", float) if uvw is not None: - mu_phi.reference_uvw = [float(x) for x in uvw.split()] + mu_phi.reference_uvw = uvw mu_phi.mu = Univariate.from_xml_element(elem.find('mu')) mu_phi.phi = Univariate.from_xml_element(elem.find('phi')) return mu_phi @@ -246,9 +246,9 @@ class Monodirectional(UnitSphere): """ monodirectional = cls() - uvw = get_text(elem, 'reference_uvw') + uvw = get_elem_list(elem, "reference_uvw", float) if uvw is not None: - monodirectional.reference_uvw = [float(x) for x in uvw.split()] + monodirectional.reference_uvw = uvw return monodirectional @@ -504,7 +504,7 @@ class SphericalIndependent(Spatial): r = Univariate.from_xml_element(elem.find('r')) cos_theta = Univariate.from_xml_element(elem.find('cos_theta')) phi = Univariate.from_xml_element(elem.find('phi')) - origin = [float(x) for x in elem.get('origin').split()] + origin = get_elem_list(elem, "origin", float) return cls(r, cos_theta, phi, origin=origin) @@ -626,7 +626,7 @@ class CylindricalIndependent(Spatial): r = Univariate.from_xml_element(elem.find('r')) phi = Univariate.from_xml_element(elem.find('phi')) z = Univariate.from_xml_element(elem.find('z')) - origin = [float(x) for x in elem.get('origin').split()] + origin = get_elem_list(elem, "origin", float) return cls(r, phi, z, origin=origin) @@ -743,18 +743,14 @@ class MeshSpatial(Spatial): """ - mesh_id = int(elem.get('mesh_id')) + mesh_id = int(get_text(elem, "mesh_id")) # check if this mesh has been read in from another location already if mesh_id not in meshes: raise ValueError(f'Could not locate mesh with ID "{mesh_id}"') - volume_normalized = elem.get("volume_normalized") volume_normalized = get_text(elem, 'volume_normalized').lower() == 'true' - strengths = get_text(elem, 'strengths') - if strengths is not None: - strengths = [float(b) for b in get_text(elem, 'strengths').split()] - + strengths = get_elem_list(elem, 'strengths', float) return cls(meshes[mesh_id], strengths, volume_normalized) @@ -860,12 +856,10 @@ class PointCloud(Spatial): """ - coord_data = get_text(elem, 'coords') - positions = np.array([float(b) for b in coord_data.split()]).reshape((-1, 3)) + coord_data = get_elem_list(elem, 'coords', float) + positions = np.array(coord_data).reshape((-1, 3)) - strengths = get_text(elem, 'strengths') - if strengths is not None: - strengths = [float(b) for b in strengths.split()] + strengths = get_elem_list(elem, 'strengths', float) return cls(positions, strengths) @@ -979,7 +973,7 @@ class Box(Spatial): """ only_fissionable = get_text(elem, 'type') == 'fission' - params = [float(x) for x in get_text(elem, 'parameters').split()] + params = get_elem_list(elem, "parameters", float) lower_left = params[:len(params)//2] upper_right = params[len(params)//2:] return cls(lower_left, upper_right, only_fissionable) @@ -1046,7 +1040,7 @@ class Point(Spatial): Point distribution generated from XML element """ - xyz = [float(x) for x in get_text(elem, 'parameters').split()] + xyz = get_elem_list(elem, "parameters", float) return cls(xyz) diff --git a/openmc/stats/univariate.py b/openmc/stats/univariate.py index e0475bf78..28d5e87ef 100644 --- a/openmc/stats/univariate.py +++ b/openmc/stats/univariate.py @@ -12,7 +12,7 @@ import numpy as np from scipy.integrate import trapezoid import openmc.checkvalue as cv -from .._xml import get_text +from .._xml import get_elem_list, get_text from ..mixin import EqualityMixin _INTERPOLATION_SCHEMES = { @@ -57,8 +57,7 @@ class Univariate(EqualityMixin, ABC): return Normal.from_xml_element(elem) elif distribution == 'muir': # Support older files where Muir had its own class - params = [float(x) for x in get_text(elem, 'parameters').split()] - return muir(*params) + return muir(*get_elem_list(elem, "parameters", float)) elif distribution == 'tabular': return Tabular.from_xml_element(elem) elif distribution == 'legendre': @@ -240,7 +239,7 @@ class Discrete(Univariate): Discrete distribution generated from XML element """ - params = [float(x) for x in get_text(elem, 'parameters').split()] + params = get_elem_list(elem, "parameters", float) x = params[:len(params)//2] p = params[len(params)//2:] return cls(x, p) @@ -448,8 +447,8 @@ class Uniform(Univariate): Uniform distribution generated from XML element """ - params = get_text(elem, 'parameters').split() - return cls(*map(float, params)) + params = get_elem_list(elem, "parameters", float) + return cls(*params) class PowerLaw(Univariate): @@ -557,8 +556,8 @@ class PowerLaw(Univariate): Distribution generated from XML element """ - params = get_text(elem, 'parameters').split() - return cls(*map(float, params)) + params = get_elem_list(elem, "parameters", float) + return cls(*params) class Maxwell(Univariate): @@ -737,8 +736,8 @@ class Watt(Univariate): Watt distribution generated from XML element """ - params = get_text(elem, 'parameters').split() - return cls(*map(float, params)) + params = get_elem_list(elem, "parameters", float) + return cls(*params) class Normal(Univariate): @@ -827,8 +826,8 @@ class Normal(Univariate): Normal distribution generated from XML element """ - params = get_text(elem, 'parameters').split() - return cls(*map(float, params)) + params = get_elem_list(elem, "parameters", float) + return cls(*params) def muir(e0: float, m_rat: float, kt: float): @@ -1115,7 +1114,7 @@ class Tabular(Univariate): """ interpolation = get_text(elem, 'interpolation') - params = [float(x) for x in get_text(elem, 'parameters').split()] + params = get_elem_list(elem, "parameters", float) m = (len(params) + 1)//2 # +1 for when len(params) is odd x = params[:m] p = params[m:] diff --git a/openmc/surface.py b/openmc/surface.py index 840c3125d..4839783ff 100644 --- a/openmc/surface.py +++ b/openmc/surface.py @@ -13,6 +13,7 @@ from .checkvalue import check_type, check_value, check_length, check_greater_tha from .mixin import IDManagerMixin, IDWarning from .region import Region, Intersection, Union from .bounding_box import BoundingBox +from ._xml import get_elem_list, get_text _BOUNDARY_TYPES = {'transmission', 'vacuum', 'reflective', 'periodic', 'white'} @@ -451,17 +452,17 @@ class Surface(IDManagerMixin, ABC): """ # Determine appropriate class - surf_type = elem.get('type') + surf_type = get_text(elem, "type") cls = _SURFACE_CLASSES[surf_type] # Determine ID, boundary type, boundary albedo, coefficients kwargs = {} - kwargs['surface_id'] = int(elem.get('id')) - kwargs['boundary_type'] = elem.get('boundary', 'transmission') + kwargs['surface_id'] = int(get_text(elem, "id")) + kwargs['boundary_type'] = get_text(elem, "boundary", "transmission") if kwargs['boundary_type'] in _ALBEDO_BOUNDARIES: - kwargs['albedo'] = float(elem.get('albedo', 1.0)) - kwargs['name'] = elem.get('name') - coeffs = [float(x) for x in elem.get('coeffs').split()] + kwargs['albedo'] = float(get_text(elem, "albedo", 1.0)) + kwargs['name'] = get_text(elem, "name") + coeffs = get_elem_list(elem, "coeffs", float) kwargs.update(dict(zip(cls._coeff_keys, coeffs))) return cls(**kwargs) diff --git a/openmc/tallies.py b/openmc/tallies.py index 3a5b42519..075b1e991 100644 --- a/openmc/tallies.py +++ b/openmc/tallies.py @@ -15,7 +15,7 @@ import scipy.sparse as sps import openmc import openmc.checkvalue as cv -from ._xml import clean_indentation, get_text +from ._xml import clean_indentation, get_elem_list, get_text from .mixin import IDManagerMixin from .mesh import MeshBase @@ -1006,8 +1006,8 @@ class Tally(IDManagerMixin): Tally object """ - tally_id = int(elem.get('id')) - name = elem.get('name', '') + tally_id = int(get_text(elem, "id")) + name = get_text(elem, "name", "") tally = cls(tally_id=tally_id, name=name) text = get_text(elem, 'multiply_density') @@ -1015,25 +1015,24 @@ class Tally(IDManagerMixin): tally.multiply_density = text in ('true', '1') # Read filters - filters_elem = elem.find('filters') - if filters_elem is not None: - filter_ids = [int(x) for x in filters_elem.text.split()] + filter_ids = get_elem_list(elem, "filters", int) + if filter_ids is not None: 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() + nuclides = get_elem_list(elem, "nuclides", str) + if nuclides is not None: + tally.nuclides = nuclides # Read scores - scores_elem = elem.find('scores') - if scores_elem is not None: - tally.scores = scores_elem.text.split() + scores = get_elem_list(elem, "scores", str) + if scores is not None: + tally.scores = scores # Set estimator - estimator_elem = elem.find('estimator') - if estimator_elem is not None: - tally.estimator = estimator_elem.text + estimator = get_text(elem, "estimator") + if estimator is not None: + tally.estimator = estimator # Read triggers tally.triggers = [ @@ -1042,9 +1041,9 @@ class Tally(IDManagerMixin): ] # Read tally derivative - deriv_elem = elem.find('derivative') - if deriv_elem is not None: - deriv_id = int(deriv_elem.text) + deriv = get_text(elem, "derivative") + if deriv is not None: + deriv_id = int(deriv) tally.derivative = kwargs['derivatives'][deriv_id] return tally diff --git a/openmc/tally_derivative.py b/openmc/tally_derivative.py index ff918c1b3..f7ba5dce5 100644 --- a/openmc/tally_derivative.py +++ b/openmc/tally_derivative.py @@ -4,6 +4,7 @@ import lxml.etree as ET import openmc.checkvalue as cv from .mixin import EqualityMixin, IDManagerMixin +from ._xml import get_text class TallyDerivative(EqualityMixin, IDManagerMixin): @@ -122,8 +123,8 @@ class TallyDerivative(EqualityMixin, IDManagerMixin): 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 + derivative_id = int(get_text(elem, "id")) + variable = get_text(elem, "variable") + material = int(get_text(elem, "material")) + nuclide = get_text(elem, "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 be2537e3a..70b6b7a03 100644 --- a/openmc/trigger.py +++ b/openmc/trigger.py @@ -5,6 +5,7 @@ import lxml.etree as ET import openmc.checkvalue as cv from .mixin import EqualityMixin +from ._xml import get_elem_list, get_text class Trigger(EqualityMixin): @@ -129,16 +130,16 @@ class Trigger(EqualityMixin): """ # Generate trigger object - trigger_type = elem.get("type") - threshold = float(elem.get("threshold")) - ignore_zeros = str(elem.get("ignore_zeros", "false")).lower() + trigger_type = get_text(elem, "type") + threshold = float(get_text(elem, "threshold")) + ignore_zeros = str(get_text(elem, "ignore_zeros", "false")).lower() # Try to convert to bool. Let Trigger error out on instantiation. ignore_zeros = ignore_zeros in ('true', '1') trigger = cls(trigger_type, threshold, ignore_zeros) # Add scores if present - scores = elem.get("scores") + scores = get_elem_list(elem, "scores", str) if scores is not None: - trigger.scores = scores.split() + trigger.scores = scores return trigger diff --git a/openmc/volume.py b/openmc/volume.py index df19def1e..c44adf98a 100644 --- a/openmc/volume.py +++ b/openmc/volume.py @@ -10,7 +10,7 @@ from uncertainties import ufloat import openmc import openmc.checkvalue as cv -from openmc._xml import get_text +from openmc._xml import get_elem_list, get_text _VERSION_VOLUME = 1 @@ -375,13 +375,10 @@ class VolumeCalculation: """ domain_type = get_text(elem, "domain_type") - domain_ids = get_text(elem, "domain_ids").split() - ids = [int(x) for x in domain_ids] + ids = get_elem_list(elem, "domain_ids", int) samples = int(get_text(elem, "samples")) - lower_left = get_text(elem, "lower_left").split() - lower_left = tuple([float(x) for x in lower_left]) - upper_right = get_text(elem, "upper_right").split() - upper_right = tuple([float(x) for x in upper_right]) + lower_left = tuple(get_elem_list(elem, "lower_left", float)) + upper_right = tuple(get_elem_list(elem, "upper_right", float)) # Instantiate some throw-away domains that are used by the constructor # to assign IDs diff --git a/openmc/weight_windows.py b/openmc/weight_windows.py index 90856ebab..5d52a579a 100644 --- a/openmc/weight_windows.py +++ b/openmc/weight_windows.py @@ -14,7 +14,7 @@ from openmc.filter import _PARTICLES from openmc.mesh import MeshBase, RectilinearMesh, CylindricalMesh, SphericalMesh, UnstructuredMesh import openmc.checkvalue as cv from openmc.checkvalue import PathLike -from ._xml import get_text, clean_indentation +from ._xml import get_elem_list, get_text, clean_indentation from .mixin import IDManagerMixin from .utility_funcs import change_directory @@ -379,9 +379,9 @@ class WeightWindows(IDManagerMixin): mesh = meshes[mesh_id] # Read all other parameters - lower_ww_bounds = [float(l) for l in get_text(elem, 'lower_ww_bounds').split()] - upper_ww_bounds = [float(u) for u in get_text(elem, 'upper_ww_bounds').split()] - e_bounds = [float(b) for b in get_text(elem, 'energy_bounds').split()] + lower_ww_bounds = get_elem_list(elem, "lower_ww_bounds", float) + upper_ww_bounds = get_elem_list(elem, "upper_ww_bounds", float) + e_bounds = get_elem_list(elem, "energy_bounds", float) particle_type = get_text(elem, 'particle_type') survival_ratio = float(get_text(elem, 'survival_ratio')) @@ -730,11 +730,8 @@ class WeightWindowGenerator: mesh_id = int(get_text(elem, 'mesh')) mesh = meshes[mesh_id] - - if (energy_bounds := get_text(elem, 'energy_bounds')) is not None: - energy_bounds = [float(x) for x in energy_bounds.split()] - else: - energy_bounds = None + + energy_bounds = get_elem_list(elem, "energy_bounds, float") particle_type = get_text(elem, 'particle_type') wwg = cls(mesh, energy_bounds, particle_type)