mirror of
https://github.com/openmc-dev/openmc.git
synced 2026-07-27 05:35:49 -04:00
Add filter method to TrackFile and Track classes
This commit is contained in:
parent
22bc340f9d
commit
7b0862c4dc
2 changed files with 115 additions and 0 deletions
|
|
@ -84,6 +84,62 @@ class Track(Sequence):
|
|||
def __len__(self):
|
||||
return len(self.particle_tracks)
|
||||
|
||||
def filter(self, particle=None, state_filter=None):
|
||||
"""Filter particle tracks by given criteria
|
||||
|
||||
Parameters
|
||||
----------
|
||||
particle : {'neutron', 'photon', 'electron', 'positron'}
|
||||
Matching particle type
|
||||
state_filter : function
|
||||
Function that takes a state (structured datatype) and returns a bool
|
||||
depending on some criteria.
|
||||
|
||||
Returns
|
||||
-------
|
||||
list
|
||||
List of :class:`openmc.ParticleTrack` objects
|
||||
|
||||
Examples
|
||||
--------
|
||||
Get all particle tracks for photons:
|
||||
|
||||
>>> track.filter(particle='photon')
|
||||
|
||||
Get all particle tracks that entered cell with ID=15:
|
||||
|
||||
>>> track.filter(state_filter=lambda s: s['cell_id'] == 15)
|
||||
|
||||
Get all particle tracks in entered material with ID=2:
|
||||
|
||||
>>> track.filter(state_filter=lambda s: s['material_id'] == 2)
|
||||
|
||||
See Also
|
||||
--------
|
||||
openmc.ParticleTrack
|
||||
|
||||
"""
|
||||
matching = []
|
||||
for t in self:
|
||||
# Check for matching particle
|
||||
if particle is not None:
|
||||
if t.particle.name.lower() != particle:
|
||||
continue
|
||||
|
||||
# Apply arbitrary state filter
|
||||
match = True
|
||||
if state_filter is not None:
|
||||
for state in t.states:
|
||||
if state_filter(state):
|
||||
break
|
||||
else:
|
||||
match = False
|
||||
|
||||
if match:
|
||||
matching.append(t)
|
||||
|
||||
return matching
|
||||
|
||||
def plot(self, axes=None):
|
||||
"""Produce a 3D plot of particle tracks
|
||||
|
||||
|
|
@ -156,6 +212,33 @@ class TrackFile(list):
|
|||
dset = fh[dset_name]
|
||||
self.append(Track(dset))
|
||||
|
||||
def filter(self, particle=None, state_filter=None):
|
||||
"""Filter tracks by given criteria
|
||||
|
||||
Parameters
|
||||
----------
|
||||
particle : {'neutron', 'photon', 'electron', 'positron'}
|
||||
Matching particle type
|
||||
state_filter : function
|
||||
Function that takes a state (structured datatype) and returns a bool
|
||||
depending on some criteria.
|
||||
|
||||
Returns
|
||||
-------
|
||||
list
|
||||
List of :class:`openmc.Track` objects
|
||||
|
||||
See Also
|
||||
--------
|
||||
openmc.Track.filter
|
||||
|
||||
"""
|
||||
matching = []
|
||||
for track in self:
|
||||
if track.filter(particle, state_filter):
|
||||
matching.append(track)
|
||||
return matching
|
||||
|
||||
def plot(self):
|
||||
"""Produce a 3D plot of particle tracks
|
||||
|
||||
|
|
|
|||
|
|
@ -107,3 +107,35 @@ def test_max_tracks(sphere_model, run_in_tmpdir):
|
|||
# Open track file and make sure we have correct number of tracks
|
||||
tracks = openmc.TrackFile('tracks.h5')
|
||||
assert len(tracks) == expected_num_tracks
|
||||
|
||||
|
||||
def test_filter(sphere_model, run_in_tmpdir):
|
||||
# Set maximum number of tracks per process to write
|
||||
sphere_model.settings.max_tracks = 25
|
||||
sphere_model.settings.photon_transport = True
|
||||
|
||||
# Run OpenMC to generate tracks.h5 file
|
||||
generate_track_file(sphere_model, tracks=True)
|
||||
|
||||
tracks = openmc.TrackFile('tracks.h5')
|
||||
for track in tracks:
|
||||
# Test filtering by particle
|
||||
matches = track.filter(particle='photon')
|
||||
for x in matches:
|
||||
assert x.particle == openmc.ParticleType.PHOTON
|
||||
|
||||
# Test general state filter
|
||||
matches = track.filter(state_filter=lambda s: s['cell_id'] == 1)
|
||||
assert matches == track.particle_tracks
|
||||
matches = track.filter(state_filter=lambda s: s['cell_id'] == 2)
|
||||
assert matches == []
|
||||
matches = track.filter(state_filter=lambda s: s['E'] < 0.0)
|
||||
assert matches == []
|
||||
|
||||
# Test filter method on TrackFile
|
||||
matches = tracks.filter(particle='neutron')
|
||||
assert matches == tracks
|
||||
matches = tracks.filter(state_filter=lambda s: s['E'] > 0.0)
|
||||
assert matches == tracks
|
||||
matches = tracks.filter(particle='bunnytron')
|
||||
assert matches == []
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue