Skip to content
62 changes: 62 additions & 0 deletions test/io/test_esmf.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,68 @@ def test_read_esmf_dataset(gridpath, datasetpath):
for dim in dims:
assert dim in uxds.dims

@pytest.mark.parametrize("mask_and_scale", [True, False])
@pytest.mark.parametrize(
"num_element_conn", [[4, 3, 3], [4, 4, 4]], ids=["counted", "miscounted"]
)
def test_read_esmf_padding_independent_of_cf_decoding(
num_element_conn, mask_and_scale, tmp_path
):
"""Padding is recognized whether or not xarray decoded the fill value.

ESMF pads a short face in `elementConn` with that variable's `_FillValue`.
With CF decoding on, xarray replaces it with NaN and promotes the array to
float; with decoding off, the raw -1 comes through. Casting first and checking
for INT_FILL_VALUE afterwards recognizes neither: the raw -1 becomes the index
-2, and the NaN cast is platform-dependent -- arm64 gives 0, so the padding
decodes to -1, a negative index that silently wraps to the last node.

So the padding has to be located before the cast, from `numElementConn`.
"""
node_lon = np.array([0.0, 120.0, 120.0, 0.0, 240.0])
node_lat = np.array([0.0, 0.0, 10.0, 10.0, 5.0])

# 1-based and -1 padded, as ESMF specifies: one quad and two triangles
in_ds = xr.Dataset(
{
"nodeCoords": xr.DataArray(
np.column_stack([node_lon, node_lat]),
dims=("nodeCount", "coordDim"),
attrs={"units": "degrees"},
),
"elementConn": xr.DataArray(
np.array([[1, 2, 3, 4], [2, 5, 3, -1], [1, 4, 5, -1]], dtype=np.int32),
dims=("elementCount", "maxNodePElement"),
attrs={"_FillValue": np.int32(-1)},
),
"numElementConn": xr.DataArray(
np.array(num_element_conn, dtype=np.byte), dims="elementCount"
),
}
)

path = tmp_path / "esmf_ragged.nc"
in_ds.to_netcdf(path)

with xr.open_dataset(path, mask_and_scale=mask_and_scale) as raw:
uxgrid = ux.open_grid(raw)

np.testing.assert_array_equal(
uxgrid.face_node_connectivity.values,
np.array([
[0, 1, 2, 3],
[1, 4, 2, INT_FILL_VALUE],
[0, 3, 4, INT_FILL_VALUE],
]),
)
# The connectivity knows each face's real size whatever the file claimed
real_sizes = (uxgrid.face_node_connectivity.values != INT_FILL_VALUE).sum(axis=1)
np.testing.assert_array_equal(real_sizes, [4, 3, 3])

# "n_nodes_per_face" is still the declared count, so a miscount survives here
np.testing.assert_array_equal(uxgrid.n_nodes_per_face.values, num_element_conn)
assert uxgrid.n_node == 5

def test_esmf_round_trip_consistency(gridpath):
"""Test round-trip serialization of grid objects through ESMF xarray format.

Expand Down
27 changes: 22 additions & 5 deletions uxarray/io/_esmf.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,12 +93,29 @@ def _read_esmf(in_ds):
# assume start index is 1 if one is not provided
start_index = 1

face_node_connectivity = in_ds["elementConn"].astype(INT_DTYPE)
face_node_connectivity = xr.where(
face_node_connectivity != INT_FILL_VALUE,
face_node_connectivity - start_index,
face_node_connectivity,
element_conn = in_ds["elementConn"]
face_dim, node_dim = element_conn.dims

# "numElementConn" gives the face size, so locate the padding positionally.
# Matching the sentinel means guessing: CF decoding turns it into NaN, and the
# cast below preserves neither NaN nor the raw value as INT_FILL_VALUE.
positions = xr.DataArray(
np.arange(element_conn.sizes[node_dim], dtype=INT_DTYPE), dims=node_dim
)
fill_mask = (positions >= n_nodes_per_face).transpose(face_dim, node_dim)

# NaN is never a usable index, whatever "numElementConn" claims
fill_mask = fill_mask | element_conn.isnull()
Comment thread
cmdupuis3 marked this conversation as resolved.

# ...and neither is the declared sentinel, whatever "numElementConn" claims
sentinel = element_conn.attrs.get(
"_FillValue", element_conn.encoding.get("_FillValue")
)
if sentinel is not None:
fill_mask = fill_mask | (element_conn == sentinel)

face_node_connectivity = element_conn.fillna(0).astype(INT_DTYPE) - start_index
face_node_connectivity = xr.where(fill_mask, INT_FILL_VALUE, face_node_connectivity)

out_ds["face_node_connectivity"] = xr.DataArray(
data=face_node_connectivity,
Expand Down
Loading