review suggestions from @paulromano

This commit is contained in:
Jonathan Shimwell 2022-08-15 13:39:47 +01:00
parent 2263549b9b
commit 21b874f0be
2 changed files with 13 additions and 45 deletions

View file

@ -6,7 +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__ = """\
Particle track information
@ -266,18 +266,18 @@ class Tracks(list):
track.plot(ax)
return ax
def write_tracks_to_vtk(self, filename='tracks.vtk'):
"""Creates a VTK file of the tracks
def write_tracks_to_vtk(self, filename=Path('tracks.vtk')):
"""Creates a VTP file of the tracks
Parameters
----------
filename : str
filename : path-like
Name of the VTK file to write.
Returns
-------
vtk.vtkStructuredGrid
the VTK object
vtk.vtkPolyData
the VTK vtkPolyData object produced
"""
import vtk
@ -312,10 +312,10 @@ class Tracks(list):
writer.SetInputData(data)
else:
writer.SetInput(data)
writer.SetFileName(filename)
writer.SetFileName(str(filename)) # SetFileName requires a string
writer.Write()
return filename
return data
@staticmethod
def combine(track_files, path='tracks.h5'):

View file

@ -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()