Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion src/parcels/_core/fieldset.py
Original file line number Diff line number Diff line change
Expand Up @@ -131,7 +131,7 @@ def time_interval(self):
time_intervals = [t for t in time_intervals if t is not None]
if len(time_intervals) == 0: # All fields are constant fields
return None
return functools.reduce(lambda x, y: x.intersection(y), time_intervals)
return functools.reduce(lambda x, y: x.intersection(y) if x is not None else None, time_intervals)

def add_field(self, field: Field, name: str | None = None):
"""Add a :class:`parcels.field.Field` object to the FieldSet.
Expand Down
42 changes: 30 additions & 12 deletions tests/test_fieldset.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,13 +9,12 @@
import xarray as xr

import parcels.tutorial
from parcels import Field, ParticleFile, ParticleSet, XGrid, convert, open_raw_zarr
from parcels import ParticleFile, ParticleSet, convert, open_raw_zarr
from parcels._core.fieldset import FieldSet, _datetime_to_msg
from parcels._core.model import _default_vector_field_components
from parcels._datasets.structured.generic import datasets as datasets_structured
from parcels._datasets.structured.generic import datasets_sgrid
from parcels._datasets.unstructured.generic import datasets as datasets_unstructured
from parcels.interpolators import XLinear
from tests import utils

ds = datasets_structured["ds_2d_left"]
Expand Down Expand Up @@ -221,24 +220,43 @@ def test_default_vector_field_components(data_vars, expected):
assert got == expected


# TODO restructure: use adding of fieldset notation to test this
@pytest.mark.skip("Needs updating after refactoring from https://github.com/Parcels-code/Parcels/pull/2646")
def test_fieldset_time_interval():
grid1 = XGrid.from_dataset(ds, mesh="flat")
field1 = Field("field1", ds["U_A_grid"], grid1, interp_method=XLinear)
def test_multi_model_time_interval():
ds1 = datasets_structured["ds_2d_left"][["U_A_grid", "V_A_grid", "grid"]]
fieldset = FieldSet.from_sgrid_conventions(ds1, mesh="flat")

ds2 = ds.copy()
ds2 = ds1.copy().rename({"U_A_grid": "U2", "V_A_grid": "V2"})
ds2["time"] = (ds2["time"].dims, ds2["time"].data + np.timedelta64(timedelta(days=1)), ds2["time"].attrs)
grid2 = XGrid.from_dataset(ds2, mesh="flat")
field2 = Field("field2", ds2["U_A_grid"], grid2, interp_method=XLinear)
fieldset += FieldSet.from_sgrid_conventions(ds2, mesh="flat")

ds3 = ds1.copy().rename({"U_A_grid": "U3", "V_A_grid": "V3"})
ds3["time"] = (ds3["time"].dims, ds3["time"].data + np.timedelta64(timedelta(days=2)), ds3["time"].attrs)
fieldset += FieldSet.from_sgrid_conventions(ds3, mesh="flat")

fieldset = FieldSet([field1, field2])
fieldset.add_constant_field("constant_field", 1.0, mesh="flat")

assert fieldset.time_interval.left == np.datetime64("2000-01-02")
assert len(fieldset.models) == 4
assert fieldset.time_interval.left == np.datetime64("2000-01-03")
assert fieldset.time_interval.right == np.datetime64("2001-01-01")


def test_multi_model_nonoverlapping_time_interval():
ds1 = datasets_structured["ds_2d_left"][["U_A_grid", "V_A_grid", "grid"]]
fieldset = FieldSet.from_sgrid_conventions(ds1, mesh="flat")

ds2 = ds1.copy().rename({"U_A_grid": "U2", "V_A_grid": "V2"})
ds2["time"] = (ds2["time"].dims, ds2["time"].data + np.timedelta64(timedelta(days=1000)), ds2["time"].attrs)
fieldset += FieldSet.from_sgrid_conventions(ds2, mesh="flat")

ds3 = ds1.copy().rename({"U_A_grid": "U3", "V_A_grid": "V3"})
ds3["time"] = (ds3["time"].dims, ds3["time"].data + np.timedelta64(timedelta(days=2000)), ds3["time"].attrs)
fieldset += FieldSet.from_sgrid_conventions(ds3, mesh="flat")

fieldset.add_constant_field("constant_field", 1.0, mesh="flat")

assert len(fieldset.models) == 4
assert fieldset.time_interval is None


def test_fieldset_time_interval_constant_fields():
fieldset = FieldSet([])
fieldset.add_constant_field("constant_field", 1.0)
Expand Down