Skip to content

Commit cbd9732

Browse files
Adding an iterator for xarray particleset
1 parent e67eb82 commit cbd9732

3 files changed

Lines changed: 18 additions & 3 deletions

File tree

parcels/kernel.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -316,8 +316,7 @@ def execute(self, pset, endtime, dt):
316316
self.add_positionupdate_kernels()
317317
self._positionupdate_kernels_added = True
318318

319-
for i in pset.trajectory.values:
320-
p = pset[i]
319+
for p in pset:
321320
self.evaluate_particle(p, endtime)
322321
if p.state == StatusCode.StopAllExecution:
323322
return StatusCode.StopAllExecution

parcels/particleset.py

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -174,7 +174,15 @@ def __del__(self):
174174
self._data = None
175175

176176
def __iter__(self):
177-
return iter(self._data) # TODO write an iter that iterates over particles (instead of variables)
177+
self._index = 0
178+
return self
179+
180+
def __next__(self):
181+
if self._index < len(self):
182+
p = self._data.sel(trajectory=self._index)
183+
self._index += 1
184+
return p
185+
raise StopIteration
178186

179187
def __getattr__(self, name):
180188
"""

tests/v4/test_particleset.py

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -116,6 +116,14 @@ def test_pset_add_implicit(fieldset):
116116
assert np.allclose(np.diff(pset._data.trajectory), np.ones(6), atol=1e-12)
117117

118118

119+
def test_pset_iterator(fieldset):
120+
npart = 10
121+
pset = ParticleSet(fieldset, lon=np.zeros(npart), lat=np.ones(npart))
122+
for i, particle in enumerate(pset):
123+
assert particle.trajectory == i
124+
assert i == npart - 1
125+
126+
119127
@pytest.mark.parametrize("verbose_progress", [True, False])
120128
def test_uxstommelgyre_pset_execute(verbose_progress):
121129
ds = datasets_unstructured["stommel_gyre_delaunay"]

0 commit comments

Comments
 (0)