From 50c3103ca386dba81bacfb75a360e9f6667f413e Mon Sep 17 00:00:00 2001 From: Jonathan Shimwell Date: Fri, 21 Jul 2023 18:08:02 +0100 Subject: [PATCH] Adding type hints to meshes (#2601) Co-authored-by: Paul Romano --- openmc/mesh.py | 118 +++++++++++++++++++++++++++---------------------- 1 file changed, 66 insertions(+), 52 deletions(-) diff --git a/openmc/mesh.py b/openmc/mesh.py index f9eaff6d5f..917aaa8a8c 100644 --- a/openmc/mesh.py +++ b/openmc/mesh.py @@ -1,17 +1,19 @@ from abc import ABC, abstractmethod from collections.abc import Iterable -from collections import OrderedDict from math import pi from numbers import Real, Integral from pathlib import Path +import typing import warnings import lxml.etree as ET +import h5py import numpy as np import openmc.checkvalue as cv import openmc from ._xml import get_text +from openmc.checkvalue import PathLike from .mixin import IDManagerMixin from .surface import _BOUNDARY_TYPES @@ -38,7 +40,7 @@ class MeshBase(IDManagerMixin, ABC): next_id = 1 used_ids = set() - def __init__(self, mesh_id=None, name=''): + def __init__(self, mesh_id: typing.Optional[int] = None, name: str = ''): # Initialize Mesh class attributes self.id = mesh_id self.name = name @@ -48,7 +50,7 @@ class MeshBase(IDManagerMixin, ABC): return self._name @name.setter - def name(self, name): + def name(self, name: str): if name is not None: cv.check_type(f'name for mesh ID="{self._id}"', name, str) self._name = name @@ -68,7 +70,7 @@ class MeshBase(IDManagerMixin, ABC): 'Volumes cannot be provided.') @classmethod - def from_hdf5(cls, group): + def from_hdf5(cls, group: h5py.Group): """Create mesh from HDF5 group Parameters @@ -98,7 +100,7 @@ class MeshBase(IDManagerMixin, ABC): raise ValueError('Unrecognized mesh type: "' + mesh_type + '"') @classmethod - def from_xml_element(cls, elem): + def from_xml_element(cls, elem: ET.Element): """Generates a mesh from an XML element Parameters @@ -261,10 +263,10 @@ class StructuredMesh(MeshBase): return np.prod(self.dimension) def write_data_to_vtk(self, - filename, - datasets=None, - volume_normalization=True, - curvilinear=False): + filename: PathLike, + datasets: typing.Optional[dict] = None, + volume_normalization: bool = True, + curvilinear: bool = False): """Creates a VTK object of the mesh Parameters @@ -393,7 +395,7 @@ class StructuredMesh(MeshBase): else: return result - ### Add all points to the unstructured grid, maintaining a flat list of IDs as we go ### + # Add all points to the unstructured grid, maintaining a flat list of IDs as we go ### # flat array storind point IDs for a given vertex # in the grid @@ -452,7 +454,7 @@ class StructuredMesh(MeshBase): return vtk_grid - def _check_vtk_datasets(self, datasets): + def _check_vtk_datasets(self, datasets: dict): """Perform some basic checks that the datasets are valid for this mesh Parameters @@ -510,7 +512,7 @@ class RegularMesh(StructuredMesh): """ - def __init__(self, mesh_id=None, name=''): + def __init__(self, mesh_id: typing.Optional[int] = None, name: str = ''): super().__init__(mesh_id, name) self._dimension = None @@ -523,7 +525,7 @@ class RegularMesh(StructuredMesh): return tuple(self._dimension) @dimension.setter - def dimension(self, dimension): + def dimension(self, dimension: typing.Iterable[int]): cv.check_type('mesh dimension', dimension, Iterable, Integral) cv.check_length('mesh dimension', dimension, 1, 3) self._dimension = dimension @@ -540,7 +542,7 @@ class RegularMesh(StructuredMesh): return self._lower_left @lower_left.setter - def lower_left(self, lower_left): + def lower_left(self, lower_left: typing.Iterable[Real]): cv.check_type('mesh lower_left', lower_left, Iterable, Real) cv.check_length('mesh lower_left', lower_left, 1, 3) self._lower_left = lower_left @@ -560,7 +562,7 @@ class RegularMesh(StructuredMesh): return [l + w * d for l, w, d in zip(ls, ws, dims)] @upper_right.setter - def upper_right(self, upper_right): + def upper_right(self, upper_right: typing.Iterable[Real]): cv.check_type('mesh upper_right', upper_right, Iterable, Real) cv.check_length('mesh upper_right', upper_right, 1, 3) self._upper_right = upper_right @@ -584,7 +586,7 @@ class RegularMesh(StructuredMesh): return [(u - l) / d for u, l, d in zip(us, ls, dims)] @width.setter - def width(self, width): + def width(self, width: typing.Iterable[Real]): cv.check_type('mesh width', width, Iterable, Real) cv.check_length('mesh width', width, 1, 3) self._width = width @@ -674,7 +676,7 @@ class RegularMesh(StructuredMesh): return string @classmethod - def from_hdf5(cls, group): + def from_hdf5(cls, group: h5py.Group): mesh_id = int(group.name.split('/')[-1].lstrip('mesh ')) # Read and assign mesh properties @@ -691,7 +693,13 @@ class RegularMesh(StructuredMesh): return mesh @classmethod - def from_rect_lattice(cls, lattice, division=1, mesh_id=None, name=''): + def from_rect_lattice( + cls, + lattice: 'openmc.RectLattice', + division: int = 1, + mesh_id: typing.Optional[int] = None, + name: str = '' + ): """Create mesh from an existing rectangular lattice Parameters @@ -727,10 +735,10 @@ class RegularMesh(StructuredMesh): @classmethod def from_domain( cls, - domain, - dimension=(10, 10, 10), - mesh_id=None, - name='' + domain: typing.Union['openmc.Cell', 'openmc.Region', 'openmc.Universe', 'openmc.Geometry'], + dimension: typing.Sequence[int] = (10, 10, 10), + mesh_id: typing.Optional[int] = None, + name: str = '' ): """Create mesh from an existing openmc cell, region, universe or geometry by making use of the objects bounding box property. @@ -797,7 +805,7 @@ class RegularMesh(StructuredMesh): return element @classmethod - def from_xml_element(cls, elem): + def from_xml_element(cls, elem: ET.Element): """Generate mesh from an XML element Parameters @@ -836,7 +844,7 @@ class RegularMesh(StructuredMesh): return mesh - def build_cells(self, bc=None): + def build_cells(self, bc: typing.Optional[str] = None): """Generates a lattice of universes with the same dimensionality as the mesh object. The individual cells/universes produced will not have material definitions applied and so downstream code @@ -961,6 +969,7 @@ class RegularMesh(StructuredMesh): return root_cell, cells + def Mesh(*args, **kwargs): warnings.warn("Mesh has been renamed RegularMesh. Future versions of " "OpenMC will not accept the name Mesh.") @@ -999,7 +1008,7 @@ class RectilinearMesh(StructuredMesh): """ - def __init__(self, mesh_id=None, name=''): + def __init__(self, mesh_id: int = None, name: str = ''): super().__init__(mesh_id, name) self._x_grid = None @@ -1106,7 +1115,7 @@ class RectilinearMesh(StructuredMesh): return string @classmethod - def from_hdf5(cls, group): + def from_hdf5(cls, group: h5py.Group): mesh_id = int(group.name.split('/')[-1].lstrip('mesh ')) # Read and assign mesh properties @@ -1118,7 +1127,7 @@ class RectilinearMesh(StructuredMesh): return mesh @classmethod - def from_xml_element(cls, elem): + def from_xml_element(cls, elem: ET.Element): """Generate a rectilinear mesh from an XML element Parameters @@ -1203,7 +1212,7 @@ class CylindricalMesh(StructuredMesh): """ - def __init__(self, mesh_id=None, name=''): + def __init__(self, mesh_id: int = None, name: str = ''): super().__init__(mesh_id, name) self._r_grid = None @@ -1252,7 +1261,7 @@ class CylindricalMesh(StructuredMesh): @property def z_grid(self): return self._z_grid - + @z_grid.setter def z_grid(self, grid): cv.check_type('mesh z_grid', grid, Iterable, Real) @@ -1295,7 +1304,7 @@ class CylindricalMesh(StructuredMesh): return string @classmethod - def from_hdf5(cls, group): + def from_hdf5(cls, group: h5py.Group): mesh_id = int(group.name.split('/')[-1].lstrip('mesh ')) # Read and assign mesh properties @@ -1311,10 +1320,10 @@ class CylindricalMesh(StructuredMesh): @classmethod def from_domain( cls, - domain, - dimension=(10, 10, 10), - mesh_id=None, - phi_grid_bounds=(0.0, 2*pi), + domain: typing.Union['openmc.Cell', 'openmc.Region', 'openmc.Universe', 'openmc.Geometry'], + dimension: typing.Sequence[int] = (10, 10, 10), + mesh_id: typing.Optional[int] = None, + phi_grid_bounds: typing.Sequence[float] = (0.0, 2*pi), name='' ): """Creates a regular CylindricalMesh from an existing openmc domain. @@ -1338,8 +1347,8 @@ class CylindricalMesh(StructuredMesh): Returns ------- - openmc.RegularMesh - RegularMesh instance + openmc.CylindricalMesh + CylindricalMesh instance """ cv.check_type( @@ -1407,7 +1416,7 @@ class CylindricalMesh(StructuredMesh): return element @classmethod - def from_xml_element(cls, elem): + def from_xml_element(cls, elem: ET.Element): """Generate a cylindrical mesh from an XML element Parameters @@ -1453,7 +1462,7 @@ class CylindricalMesh(StructuredMesh): return self._convert_to_cartesian(self.vertices, self.origin) @staticmethod - def _convert_to_cartesian(arr, origin): + def _convert_to_cartesian(arr, origin: typing.Sequence[float]): """Converts an array with xyz values in the first dimension (shape (3, ...)) to Cartesian coordinates. """ @@ -1595,7 +1604,7 @@ class SphericalMesh(StructuredMesh): return string @classmethod - def from_hdf5(cls, group): + def from_hdf5(cls, group: h5py.Group): mesh_id = int(group.name.split('/')[-1].lstrip('mesh ')) # Read and assign mesh properties @@ -1637,7 +1646,7 @@ class SphericalMesh(StructuredMesh): return element @classmethod - def from_xml_element(cls, elem): + def from_xml_element(cls, elem: ET.Element): """Generate a spherical mesh from an XML element Parameters @@ -1682,7 +1691,7 @@ class SphericalMesh(StructuredMesh): return self._convert_to_cartesian(self.vertices, self.origin) @staticmethod - def _convert_to_cartesian(arr, origin): + def _convert_to_cartesian(arr, origin: typing.Sequence[float]): """Converts an array with xyz values in the first dimension (shape (3, ...)) to Cartesian coordinates. """ @@ -1757,8 +1766,8 @@ class UnstructuredMesh(MeshBase): _LINEAR_TET = 0 _LINEAR_HEX = 1 - def __init__(self, filename, library, mesh_id=None, name='', - length_multiplier=1.0): + def __init__(self, filename: PathLike, library: str, mesh_id: typing.Optional[int] = None, + name: str = '', length_multiplier: float = 1.0): super().__init__(mesh_id, name) self.filename = filename self._volumes = None @@ -1783,7 +1792,7 @@ class UnstructuredMesh(MeshBase): return self._library @library.setter - def library(self, lib): + def library(self, lib: str): cv.check_value('Unstructured mesh library', lib, ('moab', 'libmesh')) self._library = lib @@ -1792,7 +1801,7 @@ class UnstructuredMesh(MeshBase): return self._size @size.setter - def size(self, size): + def size(self, size: int): cv.check_type("Unstructured mesh size", size, Integral) self._size = size @@ -1801,7 +1810,7 @@ class UnstructuredMesh(MeshBase): return self._output @output.setter - def output(self, val): + def output(self, val: bool): cv.check_type("Unstructured mesh output value", val, bool) self._output = val @@ -1819,7 +1828,7 @@ class UnstructuredMesh(MeshBase): return self._volumes @volumes.setter - def volumes(self, volumes): + def volumes(self, volumes: typing.Iterable[Real]): cv.check_type("Unstructured mesh volumes", volumes, Iterable, Real) self._volumes = volumes @@ -1851,7 +1860,7 @@ class UnstructuredMesh(MeshBase): return self._n_elements @n_elements.setter - def n_elements(self, val): + def n_elements(self, val: int): cv.check_type('Number of elements', val, Integral) self._n_elements = val @@ -1883,7 +1892,7 @@ class UnstructuredMesh(MeshBase): self.length_multiplier) return string - def centroid(self, bin): + def centroid(self, bin: int): """Return the vertex averaged centroid of an element Parameters @@ -1927,7 +1936,12 @@ class UnstructuredMesh(MeshBase): ) self.write_data_to_vtk(**kwargs) - def write_data_to_vtk(self, filename=None, datasets=None, volume_normalization=True): + def write_data_to_vtk( + self, + filename: typing.Optional[PathLike] = None, + datasets: typing.Optional[dict] = None, + volume_normalization: bool = True + ): """Map data to unstructured VTK mesh elements. Parameters @@ -2019,7 +2033,7 @@ class UnstructuredMesh(MeshBase): writer.Write() @classmethod - def from_hdf5(cls, group): + def from_hdf5(cls, group: h5py.Group): mesh_id = int(group.name.split('/')[-1].lstrip('mesh ')) filename = group['filename'][()].decode() library = group['library'][()].decode() @@ -2063,7 +2077,7 @@ class UnstructuredMesh(MeshBase): return element @classmethod - def from_xml_element(cls, elem): + def from_xml_element(cls, elem: ET.Element): """Generate unstructured mesh object from XML element Parameters