Skip to content

Commit 0e4eb65

Browse files
Implement getattr for Particles
Using temporary TestParticle class for now
1 parent cbd9732 commit 0e4eb65

3 files changed

Lines changed: 47 additions & 32 deletions

File tree

parcels/kernel.py

Lines changed: 11 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -186,16 +186,16 @@ def Setcoords(particle, fieldset, time): # pragma: no cover
186186
particle_dlon = 0 # noqa
187187
particle_dlat = 0 # noqa
188188
particle_ddepth = 0 # noqa
189-
particle["lon"][:] = particle.lon_nextloop
190-
particle["lat"][:] = particle.lat_nextloop
191-
particle["depth"][:] = particle.depth_nextloop
192-
particle["time"][:] = particle.time_nextloop
189+
particle.lon = particle.lon_nextloop
190+
particle.lat = particle.lat_nextloop
191+
particle.depth = particle.depth_nextloop
192+
particle.time = particle.time_nextloop
193193

194194
def Updatecoords(particle, fieldset, time): # pragma: no cover
195-
particle["lon_nextloop"][:] = particle.lon + particle_dlon # type: ignore[name-defined] # noqa
196-
particle["lat_nextloop"][:] = particle.lat + particle_dlat # type: ignore[name-defined] # noqa
197-
particle["depth_nextloop"][:] = particle.depth + particle_ddepth # type: ignore[name-defined] # noqa
198-
particle["time_nextloop"][:] = particle.time + particle.dt
195+
particle.lon_nextloop = particle.lon + particle_dlon # type: ignore[name-defined] # noqa
196+
particle.lat_nextloop = particle.lat + particle_dlat # type: ignore[name-defined] # noqa
197+
particle.depth_nextloop = particle.depth + particle_ddepth # type: ignore[name-defined] # noqa
198+
particle.time_nextloop = particle.time + particle.dt
199199

200200
self._pyfunc = (Setcoords + self + Updatecoords)._pyfunc
201201

@@ -377,9 +377,9 @@ def evaluate_particle(self, p, endtime):
377377
computational integration timestep
378378
"""
379379
while p.state in [StatusCode.Evaluate, StatusCode.Repeat]:
380-
pre_dt = p["dt"]
380+
pre_dt = p.dt
381381

382-
sign_dt = np.sign(p.dt.values).astype(int)
382+
sign_dt = np.sign(p.dt).astype(int)
383383
if sign_dt * (p.time_nextloop - endtime) > np.timedelta64(0, "ns"):
384384
return p
385385

@@ -399,5 +399,5 @@ def evaluate_particle(self, p, endtime):
399399
else:
400400
p.state = res
401401

402-
p["dt"][:] = pre_dt
402+
p.dt = pre_dt
403403
return p

parcels/particleset.py

Lines changed: 34 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,28 @@
2323
__all__ = ["ParticleSet"]
2424

2525

26+
class TestParticle:
27+
# Temporary class to allow for testing of ParticleSet without needing to change v3-Particle class. TODO update the Particle class
28+
def __init__(self, data, index=None):
29+
self._data = data
30+
self._index = index
31+
32+
def __getattr__(self, name):
33+
if name in ["_data", "_index"]:
34+
return object.__getattribute__(self, name)
35+
_data = object.__getattribute__(self, "_data")
36+
if name in _data:
37+
return _data[name].values[self._index]
38+
else:
39+
return False
40+
41+
def __setattr__(self, name, value):
42+
if name in ["_data", "_index"]:
43+
object.__setattr__(self, name, value)
44+
else:
45+
self._data[name][self._index] = value
46+
47+
2648
class ParticleSet:
2749
"""Class for storing particle and executing kernel over them.
2850
@@ -138,25 +160,18 @@ def __init__(
138160

139161
self._data = xr.Dataset(
140162
{
141-
"lon": (
142-
["trajectory", "obs"],
143-
np.array(lon[:, np.newaxis], dtype=lonlatdepth_dtype),
144-
), # TODO check if newaxis is needed
145-
"lat": (["trajectory", "obs"], np.array(lat[:, np.newaxis], dtype=lonlatdepth_dtype)),
146-
"depth": (["trajectory", "obs"], np.array(depth[:, np.newaxis], dtype=lonlatdepth_dtype)),
147-
"time": (["trajectory", "obs"], np.array(time[:, np.newaxis])),
148-
"dt": (["trajectory", "obs"], np.timedelta64(1, "ns") * np.ones((len(pid_orig), 1))),
149-
"state": (["trajectory", "obs"], np.zeros((len(pid_orig), 1), dtype=np.int32)),
150-
"lon_nextloop": (
151-
["trajectory", "obs"],
152-
np.array(lon[:, np.newaxis], dtype=lonlatdepth_dtype),
153-
), # TODO check if newaxis is needed
154-
"lat_nextloop": (["trajectory", "obs"], np.array(lat[:, np.newaxis], dtype=lonlatdepth_dtype)),
155-
"depth_nextloop": (["trajectory", "obs"], np.array(depth[:, np.newaxis], dtype=lonlatdepth_dtype)),
156-
"time_nextloop": (["trajectory", "obs"], np.array(time[:, np.newaxis])),
163+
"lon": (["trajectory"], lon),
164+
"lat": (["trajectory"], lat),
165+
"depth": (["trajectory"], depth),
166+
"time": (["trajectory"], time),
167+
"dt": (["trajectory"], np.timedelta64(1, "ns") * np.ones(len(pid_orig))),
168+
"state": (["trajectory"], np.zeros((len(pid_orig)), dtype=np.int32)),
169+
"lon_nextloop": (["trajectory"], lon),
170+
"lat_nextloop": (["trajectory"], lat),
171+
"depth_nextloop": (["trajectory"], depth),
172+
"time_nextloop": (["trajectory"], time),
157173
},
158174
coords={
159-
"obs": ("obs", [0]),
160175
"trajectory": ("trajectory", pid_orig),
161176
},
162177
attrs={
@@ -179,7 +194,7 @@ def __iter__(self):
179194

180195
def __next__(self):
181196
if self._index < len(self):
182-
p = self._data.sel(trajectory=self._index)
197+
p = self.__getitem__(self._index)
183198
self._index += 1
184199
return p
185200
raise StopIteration
@@ -202,7 +217,7 @@ def __getattr__(self, name):
202217

203218
def __getitem__(self, index):
204219
"""Get a single particle by index."""
205-
return self._data.sel(trajectory=index)
220+
return TestParticle(self._data, index=index)
206221

207222
@staticmethod
208223
def lonlatdepth_dtype_from_field_interp_method(field):

tests/v4/test_particleset.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -104,8 +104,8 @@ def test_pset_add_explicit(fieldset):
104104
particle = ParticleSet(pclass=Particle, lon=lon[i], lat=lat[i], fieldset=fieldset)
105105
pset.add(particle)
106106
assert len(pset) == npart
107-
assert np.allclose(pset._data["lon"][:, 0], lon, atol=1e-12)
108-
assert np.allclose(pset._data["lat"][:, 0], lat, atol=1e-12)
107+
assert np.allclose([p.lon for p in pset], lon, atol=1e-12)
108+
assert np.allclose([p.lat for p in pset], lat, atol=1e-12)
109109
assert np.allclose(np.diff(pset._data.trajectory), np.ones(pset._data.trajectory.size - 1), atol=1e-12)
110110

111111

0 commit comments

Comments
 (0)