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+
2648class 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 ):
0 commit comments