From 7b5cb97de4bc82ac335287481b83326569119196 Mon Sep 17 00:00:00 2001 From: j-atkins <106238905+j-atkins@users.noreply.github.com> Date: Mon, 14 Sep 2026 14:59:18 +0200 Subject: [PATCH 01/14] experiment optimising adcp depth fetching --- src/virtualship/instruments/adcp.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/src/virtualship/instruments/adcp.py b/src/virtualship/instruments/adcp.py index 0b9630d3..3fab2395 100644 --- a/src/virtualship/instruments/adcp.py +++ b/src/virtualship/instruments/adcp.py @@ -52,12 +52,18 @@ def __init__(self, expedition, from_data): """Initialize ADCPInstrument.""" variables = expedition.instruments_config.adcp_config.active_variables() + # TODO: this may be unnecessary for performance optimisation now that _via_tmp_ds() is not being used for underway instruments (see base.py)? + fetch_spec = FetchSpec( + depth_min=0, # ensures copernicusmarine fetches properly + depth_max=expedition.instruments_config.adcp_config.max_depth_meter, + ) + super().__init__( expedition, variables, add_bathymetry=False, verbose_progress=False, - fetch_spec=FetchSpec(), + fetch_spec=fetch_spec, from_data=from_data, ) From a25ec0cbaf1a8f1058e07719b91aaf22bc6692a8 Mon Sep 17 00:00:00 2001 From: j-atkins <106238905+j-atkins@users.noreply.github.com> Date: Tue, 15 Sep 2026 09:20:15 +0200 Subject: [PATCH 02/14] remove via_tmp_ds step and move to ChunkCachedArray backed fs --- src/virtualship/instruments/base.py | 21 +++++++++------------ 1 file changed, 9 insertions(+), 12 deletions(-) diff --git a/src/virtualship/instruments/base.py b/src/virtualship/instruments/base.py index f13a8338..f33dac9b 100644 --- a/src/virtualship/instruments/base.py +++ b/src/virtualship/instruments/base.py @@ -194,12 +194,10 @@ def _generate_fieldset(self) -> parcels.FieldSet: TODO: the need for this step may be removed as Parcels x copernicusmarine integration improves, tracked in https://github.com/Parcels-code/Parcels/issues/2756 and xref'd in VirtualShip #357 (https://github.com/Parcels-code/virtualship/issues/357) """ combined_fieldset = None - keys = list(self.variables.keys()) - time_buffer = self.fetch_spec.time_buffer + is_underway = self.instrument_type.is_underway - for key in keys: - var = self.variables[key] + for key, var in self.variables.items(): physical = var in COPERNICUSMARINE_PHYS_VARIABLES if self.from_data is not None: # load from local data @@ -228,16 +226,15 @@ def _generate_fieldset(self) -> parcels.FieldSet: fields = {key: ds[field_var_name]} ds_fset = parcels.convert.copernicusmarine_to_sgrid(fields=fields) - # streaming data performance is improved by writing to a temporary file, unnecessary for local data - if self.from_data is None: - ds_fset = self._via_tmp_ds(ds_fset) + # operations only necessary for non-underway instruments + if not is_underway: + fs = parcels.FieldSet.from_sgrid_conventions(ds_fset) - fs = parcels.FieldSet.from_sgrid_conventions(ds_fset) + # to ChunkCachedArrays for better Dask/memory management + fs = fs.to_chunk_cached_arrays() - # non-underway instruments to windowed arrays, just in case any ds is Dask backed - # underway instruments should not to converted to windowed arrays, as they use one direct fieldset.eval() call which could cause a big memory usage if the fieldset is windowed - if not self.instrument_type.is_underway: - fs = fs.to_windowed_arrays() + else: + fs = parcels.FieldSet.from_sgrid_conventions(ds_fset) combined_fieldset = combined_fieldset + fs if combined_fieldset else fs From fad859e59a7cf89fbf4e6ae34930c842f00f3efe Mon Sep 17 00:00:00 2001 From: j-atkins <106238905+j-atkins@users.noreply.github.com> Date: Tue, 15 Sep 2026 09:21:25 +0200 Subject: [PATCH 03/14] remove via_tmp_ds method --- src/virtualship/instruments/base.py | 28 ---------------------------- 1 file changed, 28 deletions(-) diff --git a/src/virtualship/instruments/base.py b/src/virtualship/instruments/base.py index f33dac9b..8789111f 100644 --- a/src/virtualship/instruments/base.py +++ b/src/virtualship/instruments/base.py @@ -328,34 +328,6 @@ def _get_local_ds(self, files: list[Path]) -> xr.Dataset: return ds - def _via_tmp_ds(self, ds: xr.Dataset) -> xr.Dataset: - """Create and re-load a temporary local dataset without loading everything into RAM, using local Zarr store for improved performance and concurrent chunk writing.""" - tmp_dir = tempfile.TemporaryDirectory() - self._tmp_dirs.append(tmp_dir) - tmp_store = Path(tmp_dir.name) / f"tmp_{id(ds)}.zarr" - - # strip pre-existing per-variable encoding, which may interfere with zarr defaults - ds_to_write = ds.copy() - for variable in ds_to_write.variables.values(): - variable.encoding = {} - - # TODO: potential trade off between speed and memory usage here... could remove to reduce memory footprint, but may slow down writing (?) - ds_to_write = ds_to_write.chunk( - {dim: size for dim, size in ds_to_write.sizes.items()} - ) - - ds_to_write.to_zarr( - tmp_store, - mode="w", - consolidated=False, - ) - - loaded_ds = xr.open_zarr( - tmp_store, chunks=None, consolidated=False - ) # chunks=None to avoid Dask backed - - return loaded_ds - @staticmethod def _sample_initial( pset: parcels.ParticleSet, From 15ed4ad6501a1fdf13234c82dadd0b99ca52000c Mon Sep 17 00:00:00 2001 From: j-atkins <106238905+j-atkins@users.noreply.github.com> Date: Tue, 15 Sep 2026 09:24:54 +0200 Subject: [PATCH 04/14] temporary fix in argo kernels to handle empty selections --- src/virtualship/instruments/argo_float.py | 14 ++++++++++++-- 1 file changed, 12 insertions(+), 2 deletions(-) diff --git a/src/virtualship/instruments/argo_float.py b/src/virtualship/instruments/argo_float.py index dd686e81..11ab5230 100644 --- a/src/virtualship/instruments/argo_float.py +++ b/src/virtualship/instruments/argo_float.py @@ -165,7 +165,12 @@ def _argo_sample_temperature(particles, fieldset): # Phase 3: ascending — sample temperature phase_mask = particles.cycle_phase == 3 depth_mask = particles.z < particles.min_depth # still ascending - sampling_particles = particles[np.logical_and(phase_mask, depth_mask)] + mask = np.logical_and(phase_mask, depth_mask) + if not np.any(mask): + # TODO: tmp fix avoiding IndexError in Parcels' ChunkCachedArray vectorized indexing when sampling with an empty ParticleSet (Parcels issue: #2906) + # TODO: can be removed when fixed upstream in Parcels + return + sampling_particles = particles[mask] sampling_particles.temperature = fieldset.T[sampling_particles] @@ -173,7 +178,12 @@ def _argo_sample_salinity(particles, fieldset): # Phase 3: ascending — sample salinity phase_mask = particles.cycle_phase == 3 depth_mask = particles.z < particles.min_depth # still ascending - sampling_particles = particles[np.logical_and(phase_mask, depth_mask)] + mask = np.logical_and(phase_mask, depth_mask) + if not np.any(mask): + # TODO: tmp fix avoiding IndexError in Parcels' ChunkCachedArray vectorized indexing when sampling with an empty ParticleSet (Parcels issue: #2906) + # TODO: can be removed when fixed upstream in Parcels + return + sampling_particles = particles[mask] sampling_particles.salinity = fieldset.S[sampling_particles] From 44686606f1bc4e5d06828692234ad344e708f5cb Mon Sep 17 00:00:00 2001 From: j-atkins <106238905+j-atkins@users.noreply.github.com> Date: Tue, 15 Sep 2026 09:33:24 +0200 Subject: [PATCH 05/14] remove unnecessary context management for instruments --- src/virtualship/cli/_run.py | 21 ++++++++++++--------- src/virtualship/instruments/base.py | 22 ---------------------- 2 files changed, 12 insertions(+), 31 deletions(-) diff --git a/src/virtualship/cli/_run.py b/src/virtualship/cli/_run.py index 61e629ce..98fed408 100644 --- a/src/virtualship/cli/_run.py +++ b/src/virtualship/cli/_run.py @@ -191,17 +191,20 @@ def _run( attr = MeasurementsToSimulate.get_attr_for_instrumenttype(itype) measurements = getattr(schedule_results.measurements_to_simulate, attr) - # initialise instrument, execute simulation within context manager - with instrument_class( + # initialise instrument + instrument = instrument_class( expedition=expedition, from_data=Path(from_data) if from_data is not None else None, - ) as instrument: - instrument.execute( - measurements=measurements, - out_path=expedition_dir.joinpath( - RESULTS, f"{itype.name.lower()}.parquet" - ), - ) + ) + + # execute simulation + instrument.execute( + measurements=measurements, + out_path=expedition_dir.joinpath( + RESULTS, f"{itype.name.lower()}.parquet" + ), + ) + except Exception as e: # clean up if unexpected error occurs if os.path.exists(problems_dir): diff --git a/src/virtualship/instruments/base.py b/src/virtualship/instruments/base.py index 8789111f..89c78fd3 100644 --- a/src/virtualship/instruments/base.py +++ b/src/virtualship/instruments/base.py @@ -3,7 +3,6 @@ import abc import collections import inspect -import tempfile from dataclasses import dataclass from datetime import timedelta from itertools import pairwise @@ -83,7 +82,6 @@ def __init__( self.add_bathymetry = add_bathymetry self.verbose_progress = verbose_progress self.fetch_spec = fetch_spec or FetchSpec() - self._tmp_dirs: list[tempfile.TemporaryDirectory] = [] # only waypoints relevant to this instrument; avoid needlessly ballooning fieldset to full expedition schedule relevant_waypoints = _get_instrument_relevant_waypoints( @@ -104,26 +102,6 @@ def __init__( self.min_lat, self.max_lat = min(wp_lats), max(wp_lats) self.min_lon, self.max_lon = min(wp_lons), max(wp_lons) - def close(self): - """Explicitly cleanup all tmp dirs.""" - tmp_dirs = getattr(self, "_tmp_dirs", None) - if not tmp_dirs: - return - for tmp_dir in tmp_dirs: - try: - tmp_dir.cleanup() - except Exception: - pass # i.e. best effort clean up - self._tmp_dirs = [] - - def __enter__(self): - """Enter the context manager.""" - return self - - def __exit__(self, exc_type, exc_val, exc_tb): - """Exit context manager, ensuring resource cleanup.""" - self.close() - def load_input_data(self) -> parcels.FieldSet: """Load and return the input data as a FieldSet for the instrument.""" try: From 40ede59d92cd8da79d9b51022f234763afbcf1bd Mon Sep 17 00:00:00 2001 From: j-atkins <106238905+j-atkins@users.noreply.github.com> Date: Tue, 15 Sep 2026 09:39:44 +0200 Subject: [PATCH 06/14] adapt tests to using chunk cached arrays and no context manager --- tests/instruments/test_base.py | 63 ++-------------------------------- 1 file changed, 3 insertions(+), 60 deletions(-) diff --git a/tests/instruments/test_base.py b/tests/instruments/test_base.py index b6955491..3213eea4 100644 --- a/tests/instruments/test_base.py +++ b/tests/instruments/test_base.py @@ -107,7 +107,7 @@ def __init__(self, **fields): setattr(self, name, value) self.fields = {} - def to_windowed_arrays(self): + def to_chunk_cached_arrays(self): """Mimic FieldSet.to_windowed_arrays.""" return self @@ -139,7 +139,6 @@ def test_load_input_data(): return_value="dummy_product_id", ), patch("copernicusmarine.open_dataset"), - patch.object(dummy, "_via_tmp_ds", side_effect=lambda ds: ds), patch("parcels.convert.copernicusmarine_to_sgrid"), patch( "parcels.FieldSet.from_sgrid_conventions", return_value=fake_fieldset @@ -228,61 +227,6 @@ def test_fetch_spec_applied_to_instrument(): assert dummy.fetch_spec.depth_max is None -def test_via_tmp_ds_roundtrip(): - """_via_tmp_ds writes to a tmp file and re-opens it.""" - mock_waypoint = MagicMock() - mock_waypoint.location.latitude = 1.0 - mock_waypoint.location.longitude = 2.0 - - with DummyInstrument( - expedition=MagicMock(schedule=MagicMock(waypoints=[mock_waypoint])), - variables={"A": "a"}, - add_bathymetry=False, - verbose_progress=False, - from_data=None, - ) as dummy: - ds = xr.Dataset( - {"temperature": (["x", "y"], [[1.0, 2.0], [3.0, 4.0]])}, - coords={"x": [0, 1], "y": [10, 20]}, - ) - result = dummy._via_tmp_ds(ds) - - assert isinstance(result, xr.Dataset) - assert "temperature" in result - assert ( - result is not ds - ) # result is new object loaded from tmp file, not the original - - result.close() - ds.close() - - -def test_instrument_context_manager(): - """Test that context manager cleans up temporary directories upon exit.""" - mock_waypoint = MagicMock() - mock_waypoint.location.latitude = 1.0 - mock_waypoint.location.longitude = 2.0 - - with DummyInstrument( - expedition=MagicMock(schedule=MagicMock(waypoints=[mock_waypoint])), - variables={"A": "a"}, - add_bathymetry=False, - verbose_progress=False, - from_data=None, - ) as dummy: - ds = xr.Dataset( - {"temperature": (["x", "y"], [[1.0, 2.0], [3.0, 4.0]])}, - coords={"x": [0, 1], "y": [10, 20]}, - ) - result = dummy._via_tmp_ds(ds) - assert len(dummy._tmp_dirs) == 1 - result.close() - ds.close() - - # outside 'with' block, tmp dirs should be cleared - assert len(dummy._tmp_dirs) == 0 - - def test_generate_fieldset_combines_fields(): mock_waypoint = MagicMock() mock_waypoint.location.latitude = 1.0 @@ -299,12 +243,11 @@ def test_generate_fieldset_combines_fields(): fs_A = MagicMock() fs_B = MagicMock() - fs_A.to_windowed_arrays.return_value = fs_A - fs_B.to_windowed_arrays.return_value = fs_B + fs_A.to_chunk_cached_arrays.return_value = fs_A + fs_B.to_chunk_cached_arrays.return_value = fs_B with ( patch.object(dummy, "_get_copernicus_ds"), - patch.object(dummy, "_via_tmp_ds"), patch("parcels.convert.copernicusmarine_to_sgrid"), patch("parcels.FieldSet.from_sgrid_conventions", side_effect=[fs_A, fs_B]), ): From 19e5c346338619232e7af6b6979cea4a5ea2fb6d Mon Sep 17 00:00:00 2001 From: j-atkins <106238905+j-atkins@users.noreply.github.com> Date: Tue, 15 Sep 2026 11:22:00 +0200 Subject: [PATCH 07/14] add stricter depth fetch specs to help manage memory footprint --- src/virtualship/instruments/adcp.py | 2 -- src/virtualship/instruments/argo_float.py | 2 ++ src/virtualship/instruments/ctd.py | 6 +++++- src/virtualship/instruments/xbt.py | 6 +++++- 4 files changed, 12 insertions(+), 4 deletions(-) diff --git a/src/virtualship/instruments/adcp.py b/src/virtualship/instruments/adcp.py index 3fab2395..f7fc1800 100644 --- a/src/virtualship/instruments/adcp.py +++ b/src/virtualship/instruments/adcp.py @@ -51,8 +51,6 @@ class ADCPInstrument(UnderwayInstrument): def __init__(self, expedition, from_data): """Initialize ADCPInstrument.""" variables = expedition.instruments_config.adcp_config.active_variables() - - # TODO: this may be unnecessary for performance optimisation now that _via_tmp_ds() is not being used for underway instruments (see base.py)? fetch_spec = FetchSpec( depth_min=0, # ensures copernicusmarine fetches properly depth_max=expedition.instruments_config.adcp_config.max_depth_meter, diff --git a/src/virtualship/instruments/argo_float.py b/src/virtualship/instruments/argo_float.py index 11ab5230..419652e1 100644 --- a/src/virtualship/instruments/argo_float.py +++ b/src/virtualship/instruments/argo_float.py @@ -257,6 +257,8 @@ def __init__(self, expedition, from_data): latlon_buffer=9.0, # [degrees] time_buffer=expedition.instruments_config.argo_float_config.lifetime.total_seconds() / (24 * 3600), # [days] + depth_min=expedition.instruments_config.argo_float_config.min_depth_meter, + depth_max=expedition.instruments_config.argo_float_config.max_depth_meter, ) super().__init__( diff --git a/src/virtualship/instruments/ctd.py b/src/virtualship/instruments/ctd.py index 50d641fd..66c40a86 100644 --- a/src/virtualship/instruments/ctd.py +++ b/src/virtualship/instruments/ctd.py @@ -134,13 +134,17 @@ class CTDInstrument(Instrument): def __init__(self, expedition, from_data): """Initialize CTDInstrument.""" variables = expedition.instruments_config.ctd_config.active_variables() + fetch_spec = FetchSpec( + depth_min=expedition.instruments_config.ctd_config.min_depth_meter, + depth_max=expedition.instruments_config.ctd_config.max_depth_meter, + ) super().__init__( expedition, variables, add_bathymetry=True, verbose_progress=False, - fetch_spec=FetchSpec(), + fetch_spec=fetch_spec, from_data=from_data, ) diff --git a/src/virtualship/instruments/xbt.py b/src/virtualship/instruments/xbt.py index fb810a65..ec78d011 100644 --- a/src/virtualship/instruments/xbt.py +++ b/src/virtualship/instruments/xbt.py @@ -89,13 +89,17 @@ class XBTInstrument(Instrument): def __init__(self, expedition, from_data): """Initialize XBTInstrument.""" variables = expedition.instruments_config.xbt_config.active_variables() + fetch_spec = FetchSpec( + depth_min=expedition.instruments_config.ctd_config.min_depth_meter, + depth_max=expedition.instruments_config.ctd_config.max_depth_meter, + ) super().__init__( expedition, variables, add_bathymetry=True, verbose_progress=False, - fetch_spec=FetchSpec(), + fetch_spec=fetch_spec, from_data=from_data, ) From c0b2701cfb73dc5cf125a1c328424b7046cf92e8 Mon Sep 17 00:00:00 2001 From: j-atkins <106238905+j-atkins@users.noreply.github.com> Date: Tue, 15 Sep 2026 11:27:36 +0200 Subject: [PATCH 08/14] limit chunk cache size for memory management --- src/virtualship/instruments/base.py | 3 ++- src/virtualship/instruments/xbt.py | 4 ++-- src/virtualship/utils.py | 5 ++++- 3 files changed, 8 insertions(+), 4 deletions(-) diff --git a/src/virtualship/instruments/base.py b/src/virtualship/instruments/base.py index 89c78fd3..357d26a4 100644 --- a/src/virtualship/instruments/base.py +++ b/src/virtualship/instruments/base.py @@ -22,6 +22,7 @@ from virtualship.utils import ( COPERNICUSMARINE_PHYS_VARIABLES, INSTRUMENT_CLASS_MAP, + MAX_CACHE_BYTES, _find_files_in_timerange, _find_nc_file_with_variable, _get_bathy_data, @@ -209,7 +210,7 @@ def _generate_fieldset(self) -> parcels.FieldSet: fs = parcels.FieldSet.from_sgrid_conventions(ds_fset) # to ChunkCachedArrays for better Dask/memory management - fs = fs.to_chunk_cached_arrays() + fs = fs.to_chunk_cached_arrays(max_cache_bytes=MAX_CACHE_BYTES) else: fs = parcels.FieldSet.from_sgrid_conventions(ds_fset) diff --git a/src/virtualship/instruments/xbt.py b/src/virtualship/instruments/xbt.py index ec78d011..34dad3be 100644 --- a/src/virtualship/instruments/xbt.py +++ b/src/virtualship/instruments/xbt.py @@ -90,8 +90,8 @@ def __init__(self, expedition, from_data): """Initialize XBTInstrument.""" variables = expedition.instruments_config.xbt_config.active_variables() fetch_spec = FetchSpec( - depth_min=expedition.instruments_config.ctd_config.min_depth_meter, - depth_max=expedition.instruments_config.ctd_config.max_depth_meter, + depth_min=expedition.instruments_config.xbt_config.min_depth_meter, + depth_max=expedition.instruments_config.xbt_config.max_depth_meter, ) super().__init__( diff --git a/src/virtualship/utils.py b/src/virtualship/utils.py index a3c7ba6e..694bc837 100644 --- a/src/virtualship/utils.py +++ b/src/virtualship/utils.py @@ -45,7 +45,7 @@ # projection used to sail between waypoints PROJECTION = pyproj.Geod(ellps="WGS84") -# caching for problems module +# problems module CACHE = "cache" EXPEDITION_IDENTIFIER = "id_latest.txt" PROBLEMS_ENCOUNTERED = "problems_encountered_" + "{expedition_id}" @@ -55,6 +55,9 @@ EXPEDITION_ORIGINAL = "expedition_original.yaml" EXPEDITION_LATEST = "expedition_latest.yaml" +# Parcels cacheing +MAX_CACHE_BYTES = 300_000_000 # [Bytes per variable] + # ===================================================== # SECTION: Copernicus Marine Service constants From d08a7a2f2831538c966f08afca30cbe415c7b41f Mon Sep 17 00:00:00 2001 From: j-atkins <106238905+j-atkins@users.noreply.github.com> Date: Mon, 21 Sep 2026 09:41:20 +0200 Subject: [PATCH 09/14] fix test set up errors after merge --- tests/instruments/test_base.py | 10 +++------- 1 file changed, 3 insertions(+), 7 deletions(-) diff --git a/tests/instruments/test_base.py b/tests/instruments/test_base.py index 389dd790..adff0661 100644 --- a/tests/instruments/test_base.py +++ b/tests/instruments/test_base.py @@ -167,8 +167,8 @@ def __init__(self, **fields): setattr(self, name, value) self.fields = {} - def to_chunk_cached_arrays(self): - """Mimic FieldSet.to_windowed_arrays.""" + def to_chunk_cached_arrays(self, **kwargs): + """Mimic FieldSet.to_chunk_cached_arrays.""" return self @@ -268,11 +268,7 @@ def test_fetch_spec_applied_to_instrument(mock_expedition): assert dummy.fetch_spec.depth_max is None -def test_generate_fieldset_combines_fields(): - mock_waypoint = MagicMock() - mock_waypoint.location.latitude = 1.0 - mock_waypoint.location.longitude = 2.0 - +def test_generate_fieldset_combines_fields(mock_expedition): dummy = DummyInstrument( expedition=mock_expedition, variables={"A": "a", "B": "b"}, From 0e355e30b4afa3749c6536090bea7de0673fcf04 Mon Sep 17 00:00:00 2001 From: j-atkins <106238905+j-atkins@users.noreply.github.com> Date: Mon, 28 Sep 2026 14:32:13 +0200 Subject: [PATCH 10/14] read settings/config from its own subsets of particles/floats in different phases --- src/virtualship/instruments/argo_float.py | 22 +++++++++++----------- 1 file changed, 11 insertions(+), 11 deletions(-) diff --git a/src/virtualship/instruments/argo_float.py b/src/virtualship/instruments/argo_float.py index dd686e81..50b3f471 100644 --- a/src/virtualship/instruments/argo_float.py +++ b/src/virtualship/instruments/argo_float.py @@ -68,13 +68,13 @@ def _argo_float_vertical_movement(particles, fieldset): ptcls4 = particles[particles.cycle_phase == 4] # Phase 0: Sinking with vertical_speed until depth is driftdepth - ptcls0.dz += particles.vertical_speed * ptcls0.dt + ptcls0.dz += ptcls0.vertical_speed * ptcls0.dt loc_bathy = fieldset.bathymetry.eval(ptcls0.t, ptcls0.z, ptcls0.y, ptcls0.x) - driftdepth_mask = ptcls0.z + ptcls0.dz <= particles.drift_depth # noqa:has reached drift depth + driftdepth_mask = ptcls0.z + ptcls0.dz <= ptcls0.drift_depth # noqa:has reached drift depth bathysafe_mask = ptcls0.z + ptcls0.dz >= loc_bathy # noqa:has not reached bathymetry next_phase = np.logical_and(driftdepth_mask, bathysafe_mask) ptcls0.cycle_phase[next_phase] = 1 - ptcls0.dz[next_phase] = particles.drift_depth - ptcls0.z[next_phase] # noqa:avoid overshoot + ptcls0.dz[next_phase] = ptcls0.drift_depth[next_phase] - ptcls0.z[next_phase] # noqa:avoid overshoot # Phase 0.5: Check for grounding at bathymetry and raise if necessary _handle_grounding( @@ -88,18 +88,18 @@ def _argo_float_vertical_movement(particles, fieldset): # Phase 1: Drifting at depth for drifttime seconds ptcls1.drift_age += ptcls1.dt - next_phase = ptcls1.drift_age >= particles.drift_days * 86400 # [seconds] + next_phase = ptcls1.drift_age >= ptcls1.drift_days * 86400 # [seconds] ptcls1.cycle_phase[next_phase] = 2 ptcls1.drift_age[next_phase] = 0 # reset drift_age for next cycle # Phase 2: Sinking further to maxdepth - ptcls2.dz += particles.vertical_speed * ptcls2.dt + ptcls2.dz += ptcls2.vertical_speed * ptcls2.dt loc_bathy = fieldset.bathymetry.eval(ptcls2.t, ptcls2.z, ptcls2.y, ptcls2.x) - maxdepth_mask = ptcls2.z + ptcls2.dz <= particles.max_depth # noqa:has reached max depth + maxdepth_mask = ptcls2.z + ptcls2.dz <= ptcls2.max_depth # noqa:has reached max depth bathysafe_mask = ptcls2.z + ptcls2.dz >= loc_bathy # noqa:has not reached bathymetry next_phase = np.logical_and(maxdepth_mask, bathysafe_mask) ptcls2.cycle_phase[next_phase] = 3 - ptcls2.dz[next_phase] = particles.max_depth - ptcls2.z[next_phase] # noqa:avoid overshoot + ptcls2.dz[next_phase] = ptcls2.max_depth[next_phase] - ptcls2.z[next_phase] # noqa:avoid overshoot # Phase 2.5: Check for grounding at bathymetry and raise if necessary _handle_grounding( @@ -112,13 +112,13 @@ def _argo_float_vertical_movement(particles, fieldset): ) # Phase 3: Rising with vertical_speed until at surface - ptcls3.dz -= particles.vertical_speed * ptcls3.dt - next_phase = ptcls3.z + ptcls3.dz >= particles.min_depth + ptcls3.dz -= ptcls3.vertical_speed * ptcls3.dt + next_phase = ptcls3.z + ptcls3.dz >= ptcls3.min_depth ptcls3.cycle_phase[next_phase] = 4 - ptcls3.dz[next_phase] = particles.min_depth - ptcls3.z[next_phase] # noqa:avoid overshoot + ptcls3.dz[next_phase] = ptcls3.min_depth[next_phase] - ptcls3.z[next_phase] # noqa:avoid overshoot # Phase 4: Transmitting at surface until cycletime is reached - next_phase = ptcls4.cycle_age >= particles.cycle_days * 86400 + next_phase = ptcls4.cycle_age >= ptcls4.cycle_days * 86400 ptcls4.cycle_phase[next_phase] = 0 ptcls4.cycle_age[next_phase] = 0 # reset cycle_age for next cycle ptcls4.temperature = np.nan # no temperature measurement when at surface From c29742de928bcb5d84a3ee56581723692b926d4b Mon Sep 17 00:00:00 2001 From: j-atkins <106238905+j-atkins@users.noreply.github.com> Date: Mon, 28 Sep 2026 14:33:48 +0200 Subject: [PATCH 11/14] add test for multiple drifters in different phases --- tests/instruments/test_argo_float.py | 44 ++++++++++++++++++++++++++++ 1 file changed, 44 insertions(+) diff --git a/tests/instruments/test_argo_float.py b/tests/instruments/test_argo_float.py index 1f519816..f99735db 100644 --- a/tests/instruments/test_argo_float.py +++ b/tests/instruments/test_argo_float.py @@ -183,6 +183,50 @@ def test_simulate_argo_floats(tmpdir) -> None: assert var in results, f"Results don't contain {var}" +def test_simulate_argo_floats_different_phases(tmpdir) -> None: + """Handles multiple argo floats at different time steps, in differet phases.""" + lifetime_days = 1 + fieldset = create_fieldset(lifetime_days=lifetime_days) + + sensors = [ + SensorConfig(sensor_type=SensorType.TEMPERATURE), + SensorConfig(sensor_type=SensorType.SALINITY), + ] + expedition = create_dummy_expedition( + sensors, lifetime=timedelta(days=lifetime_days) + ) + + argo_instrument = ArgoFloatInstrument(expedition, None) + + # staggered deployments + deploy_offsets = [timedelta(hours=0), timedelta(hours=2), timedelta(hours=4)] + argo_floats = [ + create_argo_float( + Waypoint( + location=Location(latitude=2 + i, longitude=1 + i), + time=BASE_TIME + offset, + instrument=[InstrumentType.ARGO_FLOAT], + ) + ) + for i, offset in enumerate(deploy_offsets) + ] + + out_path = tmpdir.join("out_phases.parquet") + argo_instrument.load_input_data = lambda: fieldset + argo_instrument.simulate(argo_floats, out_path) + + results = parcels.read_particlefile(out_path) + assert np.unique(results["particle_id"].to_numpy()).size == len(argo_floats) + + # every float should have sunk to and stayed at drift depth without overshooting + for pid in np.unique(results["particle_id"].to_numpy()): + z = results.filter(pl.col("particle_id") == pid)["z"].to_numpy() + z = z[np.isfinite(z)] + assert np.isclose(z.min(), DRIFT_DEPTH), ( + f"Float {pid} should reach drift depth without overshooting" + ) + + def test_argo_float_disabled_sensor(tmpdir) -> None: """Variables for disabled sensors must not appear in the zarr output.""" fieldset = create_fieldset(include_salinity=False) From 7e01eb62d654b56762682e12f68079e774ae82b4 Mon Sep 17 00:00:00 2001 From: j-atkins <106238905+j-atkins@users.noreply.github.com> Date: Mon, 28 Sep 2026 14:42:20 +0200 Subject: [PATCH 12/14] bug in XBT as well in subset size mismatch --- src/virtualship/instruments/xbt.py | 2 +- tests/instruments/test_xbt.py | 52 ++++++++++++++++++++++++++++++ 2 files changed, 53 insertions(+), 1 deletion(-) diff --git a/src/virtualship/instruments/xbt.py b/src/virtualship/instruments/xbt.py index fb810a65..80e2c484 100644 --- a/src/virtualship/instruments/xbt.py +++ b/src/virtualship/instruments/xbt.py @@ -70,7 +70,7 @@ def _xbt_cast(particles, fieldset): # set particle depth to max depth if it's too deep too_deep = particles.z + particles.dz < particles.max_depth - particles.dz[too_deep] = particles.max_depth - particles.z[too_deep] + particles.dz[too_deep] = particles.max_depth[too_deep] - particles.z[too_deep] # ===================================================== diff --git a/tests/instruments/test_xbt.py b/tests/instruments/test_xbt.py index e3661df7..1541b427 100644 --- a/tests/instruments/test_xbt.py +++ b/tests/instruments/test_xbt.py @@ -221,6 +221,58 @@ def test_simulate_xbts(tmpdir, xbt_expedition) -> None: ) +def test_simulate_xbts_multiple_waypoints(tmpdir, xbt_expedition) -> None: + """Should handle multiple XBTs dropped at different waypoints, all reach the bottom at their own locations.""" + BATHYMETRY = -1000.0 + + # staggered drops + waypoints = [ + (Location(latitude=0.0, longitude=0.0), datetime.timedelta(seconds=0)), + (Location(latitude=0.5, longitude=0.5), datetime.timedelta(seconds=60)), + (Location(latitude=1.0, longitude=1.0), datetime.timedelta(seconds=120)), + ] + xbts = [ + XBT( + spacetime=Spacetime(location=location, time=BASE_TIME + offset), + min_depth=0, + max_depth=float("-inf"), + fall_speed=FALL_SPEED, + deceleration_coefficient=DECELERATION_COEFFICIENT, + ) + for location, offset in waypoints + ] + + fieldset = create_fieldset( + { + "V": np.zeros((2, 2, 2, 2)), + "U": np.zeros((2, 2, 2, 2)), + "T": np.full((2, 2, 2, 2), 5.0), + }, + bathymetry_val=BATHYMETRY, + ) + + xbt_instrument = XBTInstrument(xbt_expedition, None) + out_path = tmpdir.join("out_multiple.parquet") + xbt_instrument.load_input_data = lambda: fieldset + xbt_instrument.simulate(xbts, out_path) + + results = parcels.read_particlefile(out_path) + pids = np.unique(results["particle_id"].to_numpy()) + assert pids.size == len(xbts) + + for xbt, pid in zip(xbts, pids, strict=True): + xbt_df = results.filter(pl.col("particle_id") == pid) + assert np.allclose(xbt_df["y"].to_numpy(), xbt.spacetime.location.lat), ( + f"XBT {pid} should stay at its waypoint latitude" + ) + assert np.allclose(xbt_df["x"].to_numpy(), xbt.spacetime.location.lon), ( + f"XBT {pid} should stay at its waypoint longitude" + ) + assert np.isclose(xbt_df["z"].min(), BATHYMETRY, atol=1.0), ( + f"XBT {pid} should reach the bottom" + ) + + def test_xbt_sensor_config_active_variables() -> None: """active_variables() only returns variables for enabled sensors.""" config_with_temp = XBTConfig( From 54a7980e77db4a408fcafe89e5975e563a1df42b Mon Sep 17 00:00:00 2001 From: j-atkins <106238905+j-atkins@users.noreply.github.com> Date: Mon, 28 Sep 2026 14:42:40 +0200 Subject: [PATCH 13/14] add same tests for remaining at-waypoint deployed instruments --- tests/instruments/test_ctd.py | 53 +++++++++++++++ tests/instruments/test_drifter.py | 106 ++++++++++++++++++++++++++++++ 2 files changed, 159 insertions(+) diff --git a/tests/instruments/test_ctd.py b/tests/instruments/test_ctd.py index cc8a4ad1..b05525d5 100644 --- a/tests/instruments/test_ctd.py +++ b/tests/instruments/test_ctd.py @@ -272,6 +272,59 @@ def test_simulate_ctds(tmpdir) -> None: ) # rtol to handle interpolation differences at the extreme ends of the depth range +def test_simulate_ctds_multiple_waypoints(tmpdir) -> None: + """Should handle multiple CTDs cast at different waypoints, overlapping in time (some lowering while others raise).""" + BATHYMETRY = -1000.0 + + # staggered casts + waypoints = [ + (Location(latitude=0.0, longitude=0.0), datetime.timedelta(minutes=0)), + (Location(latitude=0.5, longitude=0.5), datetime.timedelta(minutes=10)), + (Location(latitude=1.0, longitude=1.0), datetime.timedelta(minutes=20)), + ] + ctds = [ + CTD( + spacetime=Spacetime(location=location, time=BASE_TIME + offset), + min_depth=0, + max_depth=float("-inf"), + ) + for location, offset in waypoints + ] + + fieldset = create_fieldset( + {"T": np.full((2, 2, 2, 2), 5.0)}, + time_range=[ + np.datetime64(BASE_TIME), + np.datetime64(BASE_TIME + datetime.timedelta(hours=2)), + ], + bathymetry_val=BATHYMETRY, + ) + + expedition = create_dummy_expedition( + [SensorConfig(sensor_type=SensorType.TEMPERATURE)] + ) + ctd_instrument = CTDInstrument(expedition, None) + out_path = tmpdir.join("out_multiple.parquet") + ctd_instrument.load_input_data = lambda: fieldset + ctd_instrument.simulate(ctds, out_path) + + results = parcels.read_particlefile(out_path) + pids = np.unique(results["particle_id"].to_numpy()) + assert pids.size == len(ctds) + + for ctd, pid in zip(ctds, pids, strict=True): + ctd_df = results.filter(pl.col("particle_id") == pid) + assert np.allclose(ctd_df["y"].to_numpy(), ctd.spacetime.location.lat), ( + f"CTD {pid} should stay at its waypoint latitude" + ) + assert np.allclose(ctd_df["x"].to_numpy(), ctd.spacetime.location.lon), ( + f"CTD {pid} should stay at its waypoint longitude" + ) + assert np.isclose(ctd_df["z"].min(), BATHYMETRY, atol=1.0), ( + f"CTD {pid} should reach the bottom" + ) + + def test_ctd_sensor_config_active_variables() -> None: """active_variables() only returns variables for enabled sensors.""" config_both = CTDConfig( diff --git a/tests/instruments/test_drifter.py b/tests/instruments/test_drifter.py index 704c9216..211a7f55 100644 --- a/tests/instruments/test_drifter.py +++ b/tests/instruments/test_drifter.py @@ -161,6 +161,112 @@ def test_simulate_drifters(tmpdir) -> None: ) +def test_simulate_drifters_multiple_waypoints(tmpdir) -> None: + """Should handle multiple drifters deployed at different waypoints and times.""" + fieldset = create_fieldset( + { + "V": np.full((2, 2, 2), 1.0), + "U": np.full((2, 2, 2), 1.0), + "T": np.full((2, 2, 2), 1.0), + } + ) + + # all drifters share the single lifetime from the expedition's drifter config + expedition = create_dummy_expedition(lifetime=datetime.timedelta(hours=12)) + lifetime = expedition.instruments_config.drifter_config.lifetime + + # staggered deployments (at multiples of the 5h output interval), so the first drifter + # reaches the end of its lifetime whilst others are still drifting + waypoints = [ + (Location(latitude=1.0, longitude=1.0), datetime.timedelta(hours=0)), + (Location(latitude=2.0, longitude=2.0), datetime.timedelta(hours=5)), + (Location(latitude=3.0, longitude=3.0), datetime.timedelta(hours=10)), + ] + drifters = [ + Drifter( + spacetime=Spacetime(location=location, time=BASE_TIME + offset), + depth=DEPLOY_DEPTH, + lifetime=lifetime, + ) + for location, offset in waypoints + ] + + drifter_instrument = DrifterInstrument(expedition, None) + out_path = tmpdir.join("out_multiple.parquet") + drifter_instrument.load_input_data = lambda: fieldset + drifter_instrument.simulate(drifters, out_path) + + results = parcels.read_particlefile(out_path) + pids = np.unique(results["particle_id"].to_numpy()) + assert pids.size == len(drifters) + + last_times = [] + for drifter, pid in zip(drifters, pids, strict=True): + drifter_df = results.filter(pl.col("particle_id") == pid).drop_nulls("t") + first = drifter_df.sort("t")[0] + last_times.append(drifter_df["t"].max()) + + assert first["t"].item() == drifter.spacetime.time, ( + f"Drifter {pid} should start at its deployment time" + ) + assert np.isclose(first["y"].item(), drifter.spacetime.location.lat, atol=0.1) + assert np.isclose(first["x"].item(), drifter.spacetime.location.lon, atol=0.1) + + assert last_times[-1] <= drifter.spacetime.time + lifetime, ( + f"Drifter {pid} should stop at the end of its lifetime" + ) + + assert last_times == sorted(last_times) and len(set(last_times)) == len( + last_times + ), "With a shared lifetime, later-deployed drifters should stop later" + + +def test_simulate_drifters_at_same_waypoint(tmpdir) -> None: + """Multiple drifters deployed at the same waypoint (same location, time and depth) are simulated as separate drifters.""" + fieldset = create_fieldset( + { + "V": np.full((2, 2, 2), 1.0), + "U": np.full((2, 2, 2), 1.0), + "T": np.full((2, 2, 2), 1.0), + } + ) + + expedition = create_dummy_expedition() + + n_drifters = 3 + drifters = [ + Drifter( + spacetime=Spacetime( + location=Location(latitude=1.0, longitude=1.0), + time=BASE_TIME, + ), + depth=DEPLOY_DEPTH, + lifetime=expedition.instruments_config.drifter_config.lifetime, + ) + for _ in range(n_drifters) + ] + + drifter_instrument = DrifterInstrument(expedition, None) + out_path = tmpdir.join("out_same_waypoint.parquet") + drifter_instrument.load_input_data = lambda: fieldset + drifter_instrument.simulate(drifters, out_path) + + results = parcels.read_particlefile(out_path) + pids = np.unique(results["particle_id"].to_numpy()) + assert pids.size == n_drifters + + # small random noise is added to release locations, so drifters at the same waypoint follow different trajectories + release_points = set() + for pid in pids: + first = results.filter(pl.col("particle_id") == pid).sort("t")[0] + assert np.isclose(first["y"].item(), 1.0, atol=0.1) + assert np.isclose(first["x"].item(), 1.0, atol=0.1) + release_points.add((first["y"].item(), first["x"].item())) + assert len(release_points) == n_drifters, ( + "Drifters at the same waypoint should have distinct release locations" + ) + + def test_drifter_depths(tmpdir) -> None: CONST_TEMPERATURE = 1.0 # constant temperature in fieldset DEPTH_FACTOR = 3.0 # factor to multiply surface values by at depth for test From 16b08ac5afd449a9e07c9482073973707f36914f Mon Sep 17 00:00:00 2001 From: j-atkins <106238905+j-atkins@users.noreply.github.com> Date: Wed, 30 Sep 2026 16:23:55 +0200 Subject: [PATCH 14/14] revert mistake from previous merge conflict --- src/virtualship/instruments/argo_float.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/virtualship/instruments/argo_float.py b/src/virtualship/instruments/argo_float.py index ea3f4608..dc9d9453 100644 --- a/src/virtualship/instruments/argo_float.py +++ b/src/virtualship/instruments/argo_float.py @@ -165,11 +165,11 @@ def _argo_sample_temperature(particles, fieldset): phase_mask = particles.cycle_phase == 3 depth_mask = particles.z < particles.min_depth # still ascending mask = np.logical_and(phase_mask, depth_mask) + particles.temperature[~mask] = np.nan # no measurement outside the ascent if not np.any(mask): # TODO: tmp fix avoiding IndexError in Parcels' ChunkCachedArray vectorized indexing when sampling with an empty ParticleSet (Parcels issue: #2906) # TODO: can be removed when fixed upstream in Parcels return - particles.temperature[~mask] = np.nan # no measurement outside the ascent sampling_particles = particles[mask] sampling_particles.temperature = fieldset.T[sampling_particles] @@ -179,11 +179,11 @@ def _argo_sample_salinity(particles, fieldset): phase_mask = particles.cycle_phase == 3 depth_mask = particles.z < particles.min_depth # still ascending mask = np.logical_and(phase_mask, depth_mask) + particles.salinity[~mask] = np.nan # no measurement outside the ascent if not np.any(mask): # TODO: tmp fix avoiding IndexError in Parcels' ChunkCachedArray vectorized indexing when sampling with an empty ParticleSet (Parcels issue: #2906) # TODO: can be removed when fixed upstream in Parcels return - particles.salinity[~mask] = np.nan # no measurement outside the ascent sampling_particles = particles[mask] sampling_particles.salinity = fieldset.S[sampling_particles]