Skip to content
Draft
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
4 changes: 4 additions & 0 deletions benchmarks/mpas_ocean.py
Original file line number Diff line number Diff line change
Expand Up @@ -325,6 +325,10 @@ def time_zonal_average(self, resolution):
lat_step = 10
self.uxds['bottomDepth'].zonal_mean(lat=(-45, 45, lat_step))

def time_zonal_average_robust(self, resolution):
lat_step = 10
self.uxds['bottomDepth'].zonal_mean(lat=(-45, 45, lat_step), use_robust_weights=True)


class ZonalAveragePeakMem:
"""Peak memory of a cold-start non-conservative zonal-mean sweep.
Expand Down
65 changes: 26 additions & 39 deletions test/grid/integrate/test_zonal_intervals.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,7 @@
import numpy as np
import numpy.testing as nt
import pandas as pd
import polars as pl
import pytest
from polars.testing import assert_frame_equal

import uxarray as ux
from uxarray.constants import ERROR_TOLERANCE, INT_FILL_VALUE
Expand Down Expand Up @@ -34,7 +32,7 @@ def test_get_zonal_face_interval():
})
expected_interval_df_sorted = expected_interval_df.sort_values(by='start').reset_index(drop=True)

actual_values_sorted = interval_df[['start', 'end']].to_numpy()
actual_values_sorted = interval_df
expected_values_sorted = expected_interval_df_sorted[['start', 'end']].to_numpy()

nt.assert_array_almost_equal(actual_values_sorted, expected_values_sorted, decimal=13)
Expand Down Expand Up @@ -63,8 +61,7 @@ def test_get_zonal_face_interval_empty_interval():
])

res = _get_zonal_face_interval(face_edges_cart, latitude_cart, face_latlon_bounds)
expected_res = pl.DataFrame({"start": [0.0], "end": [0.0]})
assert_frame_equal(res, expected_res)
nt.assert_array_equal(res, [[0.0, 0.0]])


def test_get_zonal_face_interval_encompass_pole():
Expand Down Expand Up @@ -92,14 +89,12 @@ def test_get_zonal_face_interval_encompass_pole():
[np.arcsin(0.9992005658145248), 0.5 * np.pi],
[0, 2 * np.pi]
])
expected_df = pl.DataFrame({
'start': [0.000000, 1.101091, 2.357728, 3.614365, 4.871002, 6.127640],
'end': [0.331721, 1.588358, 2.844995, 4.101632, 5.358270, 6.283185]
})
expected_starts = [0.000000, 1.101091, 2.357728, 3.614365, 4.871002, 6.127640]
expected_ends = [0.331721, 1.588358, 2.844995, 4.101632, 5.358270, 6.283185]

res = _get_zonal_face_interval(face_edges_cart, latitude_cart, face_latlon_bounds)

assert_frame_equal(res, expected_df)
nt.assert_allclose(res, np.column_stack([expected_starts, expected_ends]), rtol=1e-5, atol=1e-8)


def test_get_zonal_face_interval_FILL_VALUE():
Expand Down Expand Up @@ -127,7 +122,7 @@ def test_get_zonal_face_interval_FILL_VALUE():
})
expected_interval_df_sorted = expected_interval_df.sort_values(by='start').reset_index(drop=True)

actual_values_sorted = interval_df[['start', 'end']].to_numpy()
actual_values_sorted = interval_df
expected_values_sorted = expected_interval_df_sorted[['start', 'end']].to_numpy()

nt.assert_array_almost_equal(actual_values_sorted, expected_values_sorted, decimal=13)
Expand Down Expand Up @@ -157,7 +152,7 @@ def test_get_zonal_face_interval_GCA_constLat():
})
expected_interval_df_sorted = expected_interval_df.sort_values(by='start').reset_index(drop=True)

actual_values_sorted = interval_df[['start', 'end']].to_numpy()
actual_values_sorted = interval_df
expected_values_sorted = expected_interval_df_sorted[['start', 'end']].to_numpy()

nt.assert_array_almost_equal(actual_values_sorted, expected_values_sorted, decimal=13)
Expand Down Expand Up @@ -185,7 +180,7 @@ def test_get_zonal_face_interval_equator():
})
expected_interval_df_sorted = expected_interval_df.sort_values(by='start').reset_index(drop=True)

actual_values_sorted = interval_df[['start', 'end']].to_numpy()
actual_values_sorted = interval_df
expected_values_sorted = expected_interval_df_sorted[['start', 'end']].to_numpy()

nt.assert_array_almost_equal(actual_values_sorted, expected_values_sorted, decimal=13)
Expand Down Expand Up @@ -220,14 +215,9 @@ def test_process_overlapped_intervals_overlap_and_gap():
},
]

# Create Polars DataFrame with explicit types
df = pl.DataFrame(
{
'start': pl.Series([x['start'] for x in intervals_data], dtype=pl.Float64),
'end': pl.Series([x['end'] for x in intervals_data], dtype=pl.Float64),
'face_index': pl.Series([x['face_index'] for x in intervals_data], dtype=pl.Int64)
}
)
starts = np.array([x['start'] for x in intervals_data], dtype=np.float64)
ends = np.array([x['end'] for x in intervals_data], dtype=np.float64)
face_indices = np.array([x['face_index'] for x in intervals_data], dtype=np.int64)

# Expected results
expected_overlap_contributions = {
Expand All @@ -239,7 +229,9 @@ def test_process_overlapped_intervals_overlap_and_gap():
}

# Process intervals
overlap_contributions, total_length = _process_overlapped_intervals(df)
overlap_contributions, total_length = _process_overlapped_intervals(
starts, ends, face_indices, len(expected_overlap_contributions)
)

# Assertions
assert abs(total_length - 340.0) < 1e-10 # Using small epsilon for float comparison
Expand All @@ -250,11 +242,11 @@ def test_process_overlapped_intervals_overlap_and_gap():
f"Mismatch for face_index {face_idx}: expected {expected_value}, got {overlap_contributions[face_idx]}"

# Check that we have all expected face indices
assert set(overlap_contributions.keys()) == set(expected_overlap_contributions.keys()), \
assert len(overlap_contributions) == len(expected_overlap_contributions), \
"Mismatch in face indices"

# Check total contributions sum matches total length
assert abs(sum(overlap_contributions.values()) - total_length) < 1e-10, \
assert abs(overlap_contributions.sum() - total_length) < 1e-10, \
"Sum of contributions doesn't match total length"


Expand Down Expand Up @@ -287,14 +279,9 @@ def test_process_overlapped_intervals_antimeridian():
},
]

# Create Polars DataFrame with explicit types
df = pl.DataFrame(
{
'start': pl.Series([x['start'] for x in intervals_data], dtype=pl.Float64),
'end': pl.Series([x['end'] for x in intervals_data], dtype=pl.Float64),
'face_index': pl.Series([x['face_index'] for x in intervals_data], dtype=pl.Int64)
}
)
starts = np.array([x['start'] for x in intervals_data], dtype=np.float64)
ends = np.array([x['end'] for x in intervals_data], dtype=np.float64)
face_indices = np.array([x['face_index'] for x in intervals_data], dtype=np.int64)

# Expected results for antimeridian case
expected_overlap_contributions = {
Expand All @@ -305,7 +292,9 @@ def test_process_overlapped_intervals_antimeridian():
}

# Process intervals
overlap_contributions, total_length = _process_overlapped_intervals(df)
overlap_contributions, total_length = _process_overlapped_intervals(
starts, ends, face_indices, len(expected_overlap_contributions)
)

# Assert total length
assert abs(total_length - 350.0) < 1e-10, \
Expand All @@ -317,7 +306,7 @@ def test_process_overlapped_intervals_antimeridian():
f"Mismatch for face_index {face_idx}: expected {expected_value}, got {overlap_contributions[face_idx]}"

# Verify all expected face indices are present
assert set(overlap_contributions.keys()) == set(expected_overlap_contributions.keys()), \
assert len(overlap_contributions) == len(expected_overlap_contributions), \
"Mismatch in face indices"


Expand All @@ -341,9 +330,7 @@ def test_get_zonal_face_interval_pole():
])
constLat_cart = -0.9986295347545738

weight_df = _get_zonal_face_interval(face_edges_cart, constLat_cart, face_bounds)
df_null_counts = weight_df.null_count()

total_nulls = df_null_counts.to_numpy().sum()
intervals = _get_zonal_face_interval(face_edges_cart, constLat_cart, face_bounds)
total_nans = np.isnan(intervals).sum()

assert total_nulls == 0, f"Found {total_nulls} null values in the DataFrame"
assert total_nans == 0, f"Found {total_nans} NaN values in the intervals"
4 changes: 2 additions & 2 deletions test/grid/integrate/test_zonal_weights.py
Original file line number Diff line number Diff line change
Expand Up @@ -427,7 +427,7 @@ def test_get_zonal_faces_weight_at_constLat_latlonface():
face_0_edge_nodes, face_1_edge_nodes, face_2_edge_nodes
]), np.sin(np.deg2rad(20)), latlon_bounds, is_latlonface=True)

nt.assert_array_almost_equal(weight_df, expected_weight_df, decimal=3)
nt.assert_array_almost_equal(weight_df, expected_weight_df["weight"], decimal=3)


def test_compare_zonal_weights(gridpath):
Expand Down Expand Up @@ -468,6 +468,6 @@ def test_compare_zonal_weights(gridpath):

existing_weights = _zonal_face_weights_robust(
face_edge_nodes_xyz_candidate, z, bounds_candidate
)["weight"].to_numpy()
)

nt.assert_allclose(new_weights, existing_weights, atol=ERROR_TOLERANCE)
2 changes: 1 addition & 1 deletion uxarray/core/zonal.py
Original file line number Diff line number Diff line change
Expand Up @@ -79,7 +79,7 @@ def _compute_non_conservative_zonal_mean(uxda, latitudes, use_robust_weights=Fal
b = bounds[face_indices]

if use_robust_weights:
w = _zonal_face_weights_robust(fe, z, b)["weight"].to_numpy()
w = _zonal_face_weights_robust(fe, z, b)
else:
w = _zonal_face_weights(fe, b, nn, z)

Expand Down
Loading
Loading