Skip to content

Commit c33b049

Browse files
Adding Support for Custom Variables
And moving from temporary TestParticle to original Particle class
1 parent fff19b3 commit c33b049

3 files changed

Lines changed: 69 additions & 65 deletions

File tree

parcels/particle.py

Lines changed: 18 additions & 41 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,7 @@
1-
from operator import attrgetter
21
from typing import Literal
32

43
import numpy as np
54

6-
from parcels.tools.statuscodes import StatusCode
7-
85
__all__ = ["InteractionParticle", "Particle", "Variable"]
96

107

@@ -110,44 +107,24 @@ class Particle:
110107
Additional Variables can be added via the :Class Variable: objects
111108
"""
112109

113-
lon = Variable("lon", dtype=np.float32)
114-
lon_nextloop = Variable("lon_nextloop", dtype=np.float32, to_write=False)
115-
lat = Variable("lat", dtype=np.float32)
116-
lat_nextloop = Variable("lat_nextloop", dtype=np.float32, to_write=False)
117-
depth = Variable("depth", dtype=np.float32)
118-
depth_nextloop = Variable("depth_nextloop", dtype=np.float32, to_write=False)
119-
time = Variable("time", dtype=np.float64)
120-
time_nextloop = Variable("time_nextloop", dtype=np.float64, to_write=False)
121-
id = Variable("id", dtype=np.int64, to_write="once")
122-
obs_written = Variable("obs_written", dtype=np.int32, initial=0, to_write=False)
123-
dt = Variable("dt", dtype=np.float64, to_write=False)
124-
state = Variable("state", dtype=np.int32, initial=StatusCode.Evaluate, to_write=False)
125-
126-
lastID = 0 # class-level variable keeping track of last Particle ID used
127-
128-
def __init__(self, lon, lat, pid, fieldset=None, ngrids=None, depth=0.0, time=0.0, cptr=None):
129-
# Enforce default values through Variable descriptor
130-
type(self).lon.initial = lon
131-
type(self).lon_nextloop.initial = lon
132-
type(self).lat.initial = lat
133-
type(self).lat_nextloop.initial = lat
134-
type(self).depth.initial = depth
135-
type(self).depth_nextloop.initial = depth
136-
type(self).time.initial = time
137-
type(self).time_nextloop.initial = time
138-
type(self).id.initial = pid
139-
type(self).lastID = max(type(self).lastID, pid)
140-
type(self).obs_written.initial = 0
141-
type(self).dt.initial = None
142-
143-
ptype = self.getPType()
144-
# Explicit initialisation of all particle variables
145-
for v in ptype.variables:
146-
if isinstance(v.initial, attrgetter):
147-
initial = v.initial(self)
148-
else:
149-
initial = v.initial
150-
setattr(self, v.name, v.dtype(initial))
110+
def __init__(self, data, index=None):
111+
self._data = data
112+
self._index = index
113+
114+
def __getattr__(self, name):
115+
if name in ["_data", "_index"]:
116+
return object.__getattribute__(self, name)
117+
_data = object.__getattribute__(self, "_data")
118+
if name in _data:
119+
return _data[name].values[self._index]
120+
else:
121+
return False
122+
123+
def __setattr__(self, name, value):
124+
if name in ["_data", "_index"]:
125+
object.__setattr__(self, name, value)
126+
else:
127+
self._data[name][self._index] = value
151128

152129
def __repr__(self):
153130
time_string = "not_yet_set" if self.time is None or np.isnan(self.time) else f"{self.time:f}"

parcels/particleset.py

Lines changed: 17 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
import sys
22
import warnings
33
from collections.abc import Iterable
4+
from operator import attrgetter
45

56
import numpy as np
67
import xarray as xr
@@ -13,7 +14,7 @@
1314
from parcels.basegrid import GridType
1415
from parcels.interaction.interactionkernel import InteractionKernel
1516
from parcels.kernel import Kernel
16-
from parcels.particle import Particle
17+
from parcels.particle import Particle, Variable
1718
from parcels.particlefile import ParticleFile
1819
from parcels.tools.converters import convert_to_flat_array
1920
from parcels.tools.loggers import logger
@@ -23,28 +24,6 @@
2324
__all__ = ["ParticleSet"]
2425

2526

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-
4827
class ParticleSet:
4928
"""Class for storing particle and executing kernel over them.
5029
@@ -179,6 +158,20 @@ def __init__(
179158
"ptype": self._pclass.getPType(), # TODO check why both pclass and ptype needed
180159
},
181160
)
161+
# add extra fields from the custom Particle class
162+
for v in self.pclass.__dict__.values():
163+
if isinstance(v, Variable):
164+
if isinstance(v.initial, attrgetter):
165+
initial = v.initial(self).values
166+
else:
167+
initial = v.initial * np.ones(len(pid_orig), dtype=v.dtype)
168+
self._data[v.name] = (["trajectory"], initial)
169+
170+
# update initial values provided on ParticleSet creation
171+
for kwvar, kwval in kwargs.items():
172+
if not hasattr(pclass, kwvar):
173+
raise RuntimeError(f"Particle class does not have Variable {kwvar}")
174+
self._data[kwvar][:] = kwval
182175

183176
self._kernel = None
184177

@@ -216,7 +209,7 @@ def __getattr__(self, name):
216209

217210
def __getitem__(self, index):
218211
"""Get a single particle by index."""
219-
return TestParticle(self._data, index=index)
212+
return Particle(self._data, index=index)
220213

221214
@staticmethod
222215
def lonlatdepth_dtype_from_field_interp_method(field):

tests/v4/test_particleset.py

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
from contextlib import nullcontext as does_not_raise
22
from datetime import datetime, timedelta
3+
from operator import attrgetter
34

45
import numpy as np
56
import pytest
@@ -66,6 +67,39 @@ def test_create_empty_pset(fieldset):
6667
assert pset.size == 0
6768

6869

70+
def test_pset_custominit_on_pset(fieldset):
71+
MyParticle = Particle.add_variable("sample_var")
72+
73+
pset = ParticleSet(fieldset, lon=0, lat=0, pclass=MyParticle, sample_var=5)
74+
75+
pset.execute(DoNothing, dt=np.timedelta64(1, "s"), runtime=np.timedelta64(21, "s"))
76+
assert np.allclose([p.sample_var for p in pset], 5.0)
77+
78+
79+
def test_pset_custominit_on_pset_attrgetter(fieldset):
80+
MyParticle = Particle.add_variable("sample_var", initial=attrgetter("lon"))
81+
82+
pset = ParticleSet(fieldset, lon=3, lat=0, pclass=MyParticle)
83+
84+
pset.execute(DoNothing, dt=np.timedelta64(1, "s"), runtime=np.timedelta64(21, "s"))
85+
assert np.allclose([p.sample_var for p in pset], 3.0)
86+
87+
88+
@pytest.mark.parametrize("pset_override", [True, False])
89+
def test_pset_custominit_on_pclass(fieldset, pset_override):
90+
MyParticle = Particle.add_variable("sample_var", initial=4)
91+
92+
if pset_override:
93+
pset = ParticleSet(fieldset, lon=0, lat=0, pclass=MyParticle, sample_var=5)
94+
else:
95+
pset = ParticleSet(fieldset, lon=0, lat=0, pclass=MyParticle)
96+
97+
pset.execute(DoNothing, dt=np.timedelta64(1, "s"), runtime=np.timedelta64(21, "s"))
98+
99+
check_val = 5.0 if pset_override else 4.0
100+
assert np.allclose([p.sample_var for p in pset], check_val)
101+
102+
69103
@pytest.mark.parametrize(
70104
"time, expectation",
71105
[

0 commit comments

Comments
 (0)