diff --git a/openmc/tracks.py b/openmc/tracks.py index e9097e8ee..671aa65cf 100644 --- a/openmc/tracks.py +++ b/openmc/tracks.py @@ -6,6 +6,7 @@ import h5py from .checkvalue import check_filetype_version from .source import SourceParticle, ParticleType +from pathlib import Path ParticleTrack = namedtuple('ParticleTrack', ['particle', 'states']) ParticleTrack.__doc__ = """\ @@ -266,6 +267,57 @@ class Tracks(list): track.plot(ax) return ax + def write_tracks_to_vtk(self, filename=Path('tracks.vtp')): + """Creates a VTP file of the tracks + + Parameters + ---------- + filename : path-like + Name of the VTP file to write. + + Returns + ------- + vtk.vtkPolyData + the VTK vtkPolyData object produced + """ + + import vtk + + # Initialize data arrays and offset. + points = vtk.vtkPoints() + cells = vtk.vtkCellArray() + + point_offset = 0 + for particle in self: + for pt in particle.particle_tracks: + for state in pt.states: + points.InsertNextPoint(state['r']) + + # Create VTK line and assign points to line. + n = pt.states.size + line = vtk.vtkPolyLine() + line.GetPointIds().SetNumberOfIds(n) + for i in range(n): + line.GetPointIds().SetId(i, point_offset + i) + point_offset += n + + # Add line to cell array + cells.InsertNextCell(line) + + data = vtk.vtkPolyData() + data.SetPoints(points) + data.SetLines(cells) + + writer = vtk.vtkXMLPPolyDataWriter() + if vtk.vtkVersion.GetVTKMajorVersion() > 5: + writer.SetInputData(data) + else: + writer.SetInput(data) + writer.SetFileName(str(filename)) # SetFileName requires a string + writer.Write() + + return data + @staticmethod def combine(track_files, path='tracks.h5'): """Combine multiple track files into a single track file diff --git a/scripts/openmc-track-to-vtk b/scripts/openmc-track-to-vtk index 3b3507974..118e65827 100755 --- a/scripts/openmc-track-to-vtk +++ b/scripts/openmc-track-to-vtk @@ -14,8 +14,8 @@ def _parse_args(): # Create argument parser. parser = argparse.ArgumentParser( description='Convert particle track file(s) to a .pvtp file.') - parser.add_argument('input', metavar='IN', type=str, nargs='+', - help='Input particle track data filename(s).') + parser.add_argument('input', metavar='IN', type=str, + help='Input particle track data filename.') parser.add_argument('-o', '--out', metavar='OUT', type=str, dest='out', help='Output VTK poly data filename.') @@ -33,41 +33,9 @@ def main(): elif not args.out.endswith('.pvtp'): args.out += '.pvtp' - # Initialize data arrays and offset. - points = vtk.vtkPoints() - cells = vtk.vtkCellArray() - point_offset = 0 - for fname in args.input: - # Write coordinate values to points array. - track_file = openmc.Tracks(fname) - for track in track_file: - for particle in track: - for state in particle.states: - points.InsertNextPoint(state['r']) - - # Create VTK line and assign points to line. - n = particle.states.size - line = vtk.vtkPolyLine() - line.GetPointIds().SetNumberOfIds(n) - for i in range(n): - line.GetPointIds().SetId(i, point_offset + i) - point_offset += n - - # Add line to cell array - cells.InsertNextCell(line) - - data = vtk.vtkPolyData() - data.SetPoints(points) - data.SetLines(cells) - - writer = vtk.vtkXMLPPolyDataWriter() - if vtk.vtkVersion.GetVTKMajorVersion() > 5: - writer.SetInputData(data) - else: - writer.SetInput(data) - writer.SetFileName(args.out) - writer.Write() - + # Write coordinate values to points array. + track_file = openmc.Tracks(args.input) + track_file.write_tracks_to_vtk(args.out) if __name__ == '__main__': main() diff --git a/tests/unit_tests/test_tracks.py b/tests/unit_tests/test_tracks.py index e18015d81..5f3940fad 100644 --- a/tests/unit_tests/test_tracks.py +++ b/tests/unit_tests/test_tracks.py @@ -141,3 +141,19 @@ def test_filter(sphere_model, run_in_tmpdir): assert matches == tracks matches = tracks.filter(particle='bunnytron') assert matches == [] + + +def test_write_tracks_to_vtk(sphere_model): + vtk = pytest.importorskip('vtk') + # 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.Tracks('tracks.h5') + polydata = tracks.write_tracks_to_vtk('tracks.vtp') + + assert isinstance(polydata, vtk.vtkPolyData) + assert Path('tracks.vtp').is_file()