Skip to content

Commit f4c3f8d

Browse files
Adding unt test for backends
1 parent 8b8fb2c commit f4c3f8d

1 file changed

Lines changed: 62 additions & 11 deletions

File tree

tests/test_fieldset.py

Lines changed: 62 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
import pandas as pd
88
import pytest
99

10+
import parcels.tutorial
1011
from parcels import Field, ParticleFile, ParticleSet, XGrid, convert
1112
from parcels._core.fieldset import FieldSet, _datetime_to_msg
1213
from parcels._core.model import _default_vector_field_components
@@ -398,21 +399,71 @@ def test_fieldset_describe(fieldset_two_models: FieldSet):
398399
fieldset = fieldset_two_models
399400
io = StringIO()
400401
expected = """\
401-
| Name | Type | Grid number | Interp method / value |
402-
|:---------------|:------------|:--------------|:------------------------|
403-
| my_list | Context | - | [1, 2, 'hello'] |
404-
| my_value | Context | - | 2.0 |
405-
| U | Field | 0 | XLinear(...) |
406-
| V | Field | 0 | XLinear(...) |
407-
| UV | VectorField | 0 | XLinear_Velocity(...) |
408-
| U_wind | Field | 1 | XLinear(...) |
409-
| V_wind | Field | 1 | XLinear(...) |
410-
| UV_wind | VectorField | 1 | XLinear_Velocity(...) |
411-
| constant_field | Field | 2 | XConstantField(...) |
402+
| Name | Type | Grid number | Interp method / value | Backend |
403+
|:---------------|:------------|:--------------|:------------------------|:----------|
404+
| my_list | Context | - | [1, 2, 'hello'] | - |
405+
| my_value | Context | - | 2.0 | - |
406+
| U | Field | 0 | XLinear(...) | NumPy |
407+
| V | Field | 0 | XLinear(...) | NumPy |
408+
| UV | VectorField | 0 | XLinear_Velocity(...) | - |
409+
| U_wind | Field | 1 | XLinear(...) | NumPy |
410+
| V_wind | Field | 1 | XLinear(...) | NumPy |
411+
| UV_wind | VectorField | 1 | XLinear_Velocity(...) | - |
412+
| constant_field | Field | 2 | XConstantField(...) | NumPy |
412413
413414
mesh: flat
414415
time interval: (np.datetime64('2000-01-01T00:00:00.000000000'), np.datetime64('2001-01-01T00:00:00.000000000'))
415416
"""
416417
fieldset.describe(io)
417418
actual = io.getvalue()
418419
assert actual == expected
420+
421+
422+
def test_fieldset_describe_backends():
423+
ds_u = parcels.tutorial.open_dataset("NemoNorthSeaORCA025-N006_data/U")
424+
ds_v = parcels.tutorial.open_dataset("NemoNorthSeaORCA025-N006_data/V")
425+
ds_w = parcels.tutorial.open_dataset("NemoNorthSeaORCA025-N006_data/W")
426+
ds_coords = parcels.tutorial.open_dataset("NemoNorthSeaORCA025-N006_data/mesh_mask")[["glamf", "gphif"]]
427+
428+
ds_fset = convert.nemo_to_sgrid(
429+
fields={"U": ds_u["uo"], "V": ds_v["vo"], "W": ds_w["wo"]},
430+
coords=ds_coords,
431+
)
432+
fieldset = FieldSet.from_sgrid_conventions(ds_fset)
433+
434+
io = StringIO()
435+
expected = """\
436+
| Name | Type | Grid number | Interp method / value | Backend |
437+
|:-------|:------------|--------------:|:------------------------|:----------|
438+
| U | Field | 0 | XLinear(...) | Dask |
439+
| V | Field | 0 | XLinear(...) | Dask |
440+
| W | Field | 0 | XLinear(...) | Dask |
441+
| UV | VectorField | 0 | CGrid_Velocity(...) | - |
442+
| UVW | VectorField | 0 | CGrid_Velocity(...) | - |
443+
444+
mesh: spherical
445+
time interval: (np.datetime64('2000-01-02T12:00:00.000000000'), np.datetime64('2000-01-27T12:00:00.000000000'))
446+
"""
447+
fieldset.describe(io)
448+
actual = io.getvalue()
449+
assert actual == expected
450+
451+
# Also run with WindowedArray backend
452+
fieldset = fieldset.to_windowed_arrays()
453+
454+
io = StringIO()
455+
expected = """\
456+
| Name | Type | Grid number | Interp method / value | Backend |
457+
|:-------|:------------|--------------:|:------------------------|:--------------|
458+
| U | Field | 0 | XLinear(...) | WindowedArray |
459+
| V | Field | 0 | XLinear(...) | WindowedArray |
460+
| W | Field | 0 | XLinear(...) | WindowedArray |
461+
| UV | VectorField | 0 | CGrid_Velocity(...) | - |
462+
| UVW | VectorField | 0 | CGrid_Velocity(...) | - |
463+
464+
mesh: spherical
465+
time interval: (np.datetime64('2000-01-02T12:00:00.000000000'), np.datetime64('2000-01-27T12:00:00.000000000'))
466+
"""
467+
fieldset.describe(io)
468+
actual = io.getvalue()
469+
assert actual == expected

0 commit comments

Comments
 (0)