Support UnstructuredMesh for IndependentSource (#2949)

Co-authored-by: Jonathan Shimwell <drshimwell@gmail.com>
Co-authored-by: Paul Romano <paul.k.romano@gmail.com>
This commit is contained in:
Patrick Shriwise 2024-04-12 12:14:13 -05:00 committed by GitHub
parent 4ba053ca47
commit e77a5247b6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 123 additions and 19 deletions

View file

@ -864,10 +864,10 @@ class MeshFilter(Filter):
cv.check_type('filter mesh', mesh, openmc.MeshBase)
self._mesh = mesh
if isinstance(mesh, openmc.UnstructuredMesh):
if mesh.volumes is None:
self.bins = []
else:
if mesh.has_statepoint_data:
self.bins = list(range(len(mesh.volumes)))
else:
self.bins = []
else:
self.bins = list(mesh.indices)
@ -982,7 +982,7 @@ class MeshFilter(Filter):
if translation:
out.translation = [float(x) for x in translation.split()]
return out
class MeshBornFilter(MeshFilter):
"""Filter events by the mesh cell a particle originated from.

View file

@ -3,6 +3,7 @@ import typing
import warnings
from abc import ABC, abstractmethod
from collections.abc import Iterable
from functools import wraps
from math import pi, sqrt, atan2
from numbers import Integral, Real
from pathlib import Path
@ -39,7 +40,8 @@ class MeshBase(IDManagerMixin, ABC):
bounding_box : openmc.BoundingBox
Axis-aligned bounding box of the mesh as defined by the upper-right and
lower-left coordinates.
indices : Iterable of tuple
An iterable of mesh indices for each mesh element, e.g. [(1, 1, 1), (2, 1, 1), ...]
"""
next_id = 1
@ -66,6 +68,11 @@ class MeshBase(IDManagerMixin, ABC):
def bounding_box(self) -> openmc.BoundingBox:
return openmc.BoundingBox(self.lower_left, self.upper_right)
@property
@abstractmethod
def indices(self):
pass
def __repr__(self):
string = type(self).__name__ + '\n'
string += '{0: <16}{1}{2}\n'.format('\tID', '=\t', self._id)
@ -1914,6 +1921,17 @@ class SphericalMesh(StructuredMesh):
return arr
def require_statepoint_data(func):
@wraps(func)
def wrapper(self: UnstructuredMesh, *args, **kwargs):
if not self._has_statepoint_data:
raise AttributeError(f'The "{func.__name__}" property requires '
'information about this mesh to be loaded '
'from a statepoint file.')
return func(self, *args, **kwargs)
return wrapper
class UnstructuredMesh(MeshBase):
"""A 3D unstructured mesh
@ -1990,6 +2008,7 @@ class UnstructuredMesh(MeshBase):
self.library = library
self._output = False
self.length_multiplier = length_multiplier
self._has_statepoint_data = False
@property
def filename(self):
@ -2010,6 +2029,7 @@ class UnstructuredMesh(MeshBase):
self._library = lib
@property
@require_statepoint_data
def size(self):
return self._size
@ -2028,6 +2048,7 @@ class UnstructuredMesh(MeshBase):
self._output = val
@property
@require_statepoint_data
def volumes(self):
"""Return Volumes for every mesh cell if
populated by a StatePoint file
@ -2046,26 +2067,32 @@ class UnstructuredMesh(MeshBase):
self._volumes = volumes
@property
@require_statepoint_data
def total_volume(self):
return np.sum(self.volumes)
@property
@require_statepoint_data
def vertices(self):
return self._vertices
@property
@require_statepoint_data
def connectivity(self):
return self._connectivity
@property
@require_statepoint_data
def element_types(self):
return self._element_types
@property
@require_statepoint_data
def centroids(self):
return np.array([self.centroid(i) for i in range(self.n_elements)])
@property
@require_statepoint_data
def n_elements(self):
if self._n_elements is None:
raise RuntimeError("No information about this mesh has "
@ -2096,6 +2123,15 @@ class UnstructuredMesh(MeshBase):
def n_dimension(self):
return 3
@property
@require_statepoint_data
def indices(self):
return [(i,) for i in range(self.n_elements)]
@property
def has_statepoint_data(self) -> bool:
return self._has_statepoint_data
def __repr__(self):
string = super().__repr__()
string += '{: <16}=\t{}\n'.format('\tFilename', self.filename)
@ -2106,13 +2142,16 @@ class UnstructuredMesh(MeshBase):
return string
@property
@require_statepoint_data
def lower_left(self):
return self.vertices.min(axis=0)
@property
@require_statepoint_data
def upper_right(self):
return self.vertices.max(axis=0)
@require_statepoint_data
def centroid(self, bin: int):
"""Return the vertex averaged centroid of an element
@ -2257,6 +2296,7 @@ class UnstructuredMesh(MeshBase):
library = group['library'][()].decode()
mesh = cls(filename=filename, library=library, mesh_id=mesh_id)
mesh._has_statepoint_data = True
vol_data = group['volumes'][()]
mesh.volumes = np.reshape(vol_data, (vol_data.shape[0],))
mesh.n_elements = mesh.volumes.size

View file

@ -500,9 +500,13 @@ class MeshSource(SourceBase):
elem.set("mesh", str(self.mesh.id))
# write in the order of mesh indices
for idx in self.mesh.indices:
idx = tuple(i - 1 for i in idx)
elem.append(self.sources[idx].to_xml_element())
if isinstance(self.mesh, openmc.UnstructuredMesh):
for s in self.sources:
elem.append(s.to_xml_element())
else:
for idx in self.mesh.indices:
idx = tuple(i - 1 for i in idx)
elem.append(self.sources[idx].to_xml_element())
@classmethod
def from_xml_element(cls, elem: ET.Element, meshes) -> openmc.MeshSource:

View file

@ -298,3 +298,22 @@ def test_CylindricalMesh_get_indices_at_coords():
assert mesh.get_indices_at_coords([98, 200.1, 299]) == (0, 1, 0) # second angle quadrant
assert mesh.get_indices_at_coords([98, 199.9, 299]) == (0, 2, 0) # third angle quadrant
assert mesh.get_indices_at_coords([102, 199.1, 299]) == (0, 3, 0) # forth angle quadrant
def test_umesh_roundtrip(run_in_tmpdir, request):
umesh = openmc.UnstructuredMesh(request.path.parent / 'test_mesh_tets.e', 'moab')
umesh.output = True
# create a tally using this mesh
mf = openmc.MeshFilter(umesh)
tally = openmc.Tally()
tally.filters = [mf]
tally.scores = ['flux']
tallies = openmc.Tallies([tally])
tallies.export_to_xml()
xml_tallies = openmc.Tallies.from_xml()
xml_tally = xml_tallies[0]
xml_mesh = xml_tally.filters[0].mesh
assert umesh.id == xml_mesh.id

View file

@ -208,23 +208,29 @@ def test_roundtrip(run_in_tmpdir, model, request):
###################
# MeshSource tests
###################
@pytest.mark.parametrize('mesh_type', ('rectangular', 'cylindrical'))
def test_mesh_source_independent(run_in_tmpdir, mesh_type):
@pytest.fixture
def void_model():
"""
A void model containing a single box
"""
min, max = -10, 10
box = openmc.model.RectangularParallelepiped(
min, max, min, max, min, max, boundary_type='vacuum')
model = openmc.Model()
geometry = openmc.Geometry([openmc.Cell(region=-box)])
box = openmc.model.RectangularParallelepiped(*[-10, 10]*3, boundary_type='vacuum')
model.geometry = openmc.Geometry([openmc.Cell(region=-box)])
settings = openmc.Settings()
settings.particles = 100
settings.batches = 10
settings.run_mode = 'fixed source'
model.settings.particles = 100
model.settings.batches = 10
model.settings.run_mode = 'fixed source'
model = openmc.Model(geometry=geometry, settings=settings)
return model
@pytest.mark.parametrize('mesh_type', ('rectangular', 'cylindrical'))
def test_mesh_source_independent(run_in_tmpdir, void_model, mesh_type):
"""
A void model containing a single box
"""
model = void_model
# define a 2 x 2 x 2 mesh
if mesh_type == 'rectangular':
@ -310,6 +316,41 @@ def test_mesh_source_independent(run_in_tmpdir, mesh_type):
assert mesh_source.strength == 1.0
@pytest.mark.parametrize("library", ('moab', 'libmesh'))
def test_umesh_source_independent(run_in_tmpdir, request, void_model, library):
import openmc.lib
# skip the test if the library is not enabled
if library == 'moab' and not openmc.lib._dagmc_enabled():
pytest.skip("DAGMC (and MOAB) mesh not enabled in this build.")
if library == 'libmesh' and not openmc.lib._libmesh_enabled():
pytest.skip("LibMesh is not enabled in this build.")
model = void_model
mesh_filename = Path(request.fspath).parent / "test_mesh_tets.e"
uscd_mesh = openmc.UnstructuredMesh(mesh_filename, library)
ind_source = openmc.IndependentSource()
n_elements = 12_000
model.settings.source = openmc.MeshSource(uscd_mesh, n_elements*[ind_source])
model.export_to_model_xml()
try:
openmc.lib.init()
openmc.lib.simulation_init()
sites = openmc.lib.sample_external_source(10)
openmc.lib.statepoint_write('statepoint.h5')
finally:
openmc.lib.finalize()
with openmc.StatePoint('statepoint.h5') as sp:
uscd_mesh = sp.meshes[uscd_mesh.id]
# ensure at least that all sites are inside the mesh
bounding_box = uscd_mesh.bounding_box
for site in sites:
assert site.r in bounding_box
def test_mesh_source_file(run_in_tmpdir):
# Creating a source file with a single particle
source_particle = openmc.SourceParticle(time=10.0)