From 2fd5368344be8ccdba1413bfef1719b9e81e4945 Mon Sep 17 00:00:00 2001 From: cmdupuis3 Date: Wed, 30 Sep 2026 13:12:59 -0500 Subject: [PATCH 1/2] optimization sketch --- benchmarks/mpas_ocean.py | 4 + uxarray/grid/integrate.py | 300 ++++++++++++++++++++++++-------------- 2 files changed, 198 insertions(+), 106 deletions(-) diff --git a/benchmarks/mpas_ocean.py b/benchmarks/mpas_ocean.py index e01e2fccc..a051f960c 100644 --- a/benchmarks/mpas_ocean.py +++ b/benchmarks/mpas_ocean.py @@ -312,6 +312,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.""" diff --git a/uxarray/grid/integrate.py b/uxarray/grid/integrate.py index b15d9aeef..7548fedee 100644 --- a/uxarray/grid/integrate.py +++ b/uxarray/grid/integrate.py @@ -1,3 +1,5 @@ +import math + import numpy as np import polars as pl from numba import njit, prange, types @@ -11,45 +13,10 @@ get_number_of_intersections, ) -DUMMY_EDGE_VALUE = [INT_FILL_VALUE, INT_FILL_VALUE, INT_FILL_VALUE] - point_type = types.UniTuple(types.float64, 3) edge_type = types.UniTuple(types.int64, 2) -def _is_edge_gca(is_GCA_list, is_latlonface, edges_z): - """Determine if each edge is a Great Circle Arc (GCA) or a constant - latitude line in a vectorized manner. - - Parameters: - ---------- - is_GCA_list : np.ndarray or None - An array indicating whether each edge is a GCA (True) or a constant latitude line (False). - Shape: (n_edges). If None, edge types are determined based on `is_latlonface` and the z-coordinates. - is_latlonface : bool - Flag indicating if all edges should be considered as lat-lon faces, which implies all edges - are either constant latitude or longitude lines. - edges_z : np.ndarray - Array containing the z-coordinates for each vertex of the edges. This is used to determine - whether edges are on the equator or if they are aligned in latitude when `is_GCA_list` is None. - Shape should be (n_edges, 2). - - Returns: - ------- - np.ndarray - A boolean array where each element indicates whether the corresponding edge is considered a GCA. - True for GCA, False for constant latitude line. - """ - if is_GCA_list is not None: - return is_GCA_list - if is_latlonface: - return ~np.isclose(edges_z[:, 0], edges_z[:, 1], atol=ERROR_TOLERANCE) - return ~( - np.isclose(edges_z[:, 0], 0, atol=ERROR_TOLERANCE) - & np.isclose(edges_z[:, 1], 0, atol=ERROR_TOLERANCE) - ) - - def _zonal_face_weights_robust( faces_edges_cart_candidate: np.ndarray, latitude_cart: float, @@ -352,12 +319,194 @@ def _process_overlapped_intervals(intervals_df: pl.DataFrame): return overlap_contributions, total_length +@njit(cache=True, inline="always") +def _isclose_scalar(a, b, atol): + """Scalar form of ``np.isclose``, including its default relative tolerance.""" + return abs(a - b) <= atol + 1.0e-5 * abs(b) + + +@njit(cache=True, inline="always") +def _edge_is_valid(face_edges_cart, e): + """Determine whether edge ``e`` is a real edge rather than a dummy one.""" + for i in range(2): + for j in range(3): + if face_edges_cart[e, i, j] == INT_FILL_VALUE: + return False + return True + + +@njit(cache=True, inline="always") +def _edge_is_gca(z0, z1, is_GCA_list, is_latlonface, edge_index): + """Determine if an edge is a Great Circle Arc (GCA) or a constant latitude + line, from the z-coordinates of its two vertices. + + An explicit `is_GCA_list` entry wins; lat-lon faces treat edges of equal + latitude as constant latitude lines; otherwise only edges lying on the + equator are constant latitude lines. + """ + if is_GCA_list is not None: + return is_GCA_list[edge_index] + if is_latlonface: + return not _isclose_scalar(z0, z1, ERROR_TOLERANCE) + return not ( + _isclose_scalar(z0, 0.0, ERROR_TOLERANCE) + and _isclose_scalar(z1, 0.0, ERROR_TOLERANCE) + ) + + +@njit(cache=True, inline="always") +def _lon_rad_from_xyz(x, y, z): + """Longitude of a Cartesian point in radians, in the range [0, 2*pi]. + + Scalar counterpart of `_xyz_to_lonlat_rad`, keeping only the longitude. + """ + denom = (x**2 + y**2 + z**2) ** 0.5 + x_norm = x / denom + y_norm = y / denom + z_norm = z / denom + + # Longitude is undefined at the poles, matching `_xyz_to_lonlat_rad` + if abs(z_norm) > 1.0 - ERROR_TOLERANCE: + return 0.0 + + lon = math.atan2(y_norm, x_norm) + if lon < 0.0: + lon += 2.0 * np.pi + return lon + + +@njit(cache=True) +def _unique_rows(points, n_points): + """In-place equivalent of ``np.unique(points[:n_points], axis=0)``. + + Sorts the first `n_points` rows lexicographically and compacts duplicates to + the front, returning how many unique rows there are. An insertion sort is + used because a face only ever contributes a handful of points. + """ + for i in range(1, n_points): + x = points[i, 0] + y = points[i, 1] + z = points[i, 2] + j = i - 1 + while j >= 0: + prev_x = points[j, 0] + prev_y = points[j, 1] + prev_z = points[j, 2] + if prev_x != x: + if prev_x < x: + break + elif prev_y != y: + if prev_y < y: + break + elif not prev_z > z: + break + points[j + 1, 0] = prev_x + points[j + 1, 1] = prev_y + points[j + 1, 2] = prev_z + j -= 1 + points[j + 1, 0] = x + points[j + 1, 1] = y + points[j + 1, 2] = z + + # Duplicates are adjacent now that the rows are sorted + n_unique = 0 + for i in range(n_points): + if n_unique > 0 and ( + points[n_unique - 1, 0] == points[i, 0] + and points[n_unique - 1, 1] == points[i, 1] + and points[n_unique - 1, 2] == points[i, 2] + ): + continue + points[n_unique, 0] = points[i, 0] + points[n_unique, 1] = points[i, 1] + points[n_unique, 2] = points[i, 2] + n_unique += 1 + return n_unique + + +@njit(cache=True) +def _get_faces_constLat_intersection_info_numba( + face_edges_cart, latitude_cart, is_GCA_list, is_latlonface +): + """Numba kernel behind `_get_faces_constLat_intersection_info`. + + Returns the unique intersection points along with the minimum and maximum + longitude across them. When the face is only touched by the latitude there + is a single point and the longitudes are meaningless; the Python wrapper + replaces them with None. + """ + n_edges = face_edges_cart.shape[0] + + # Each edge contributes at most two intersection points + points = np.empty((2 * n_edges + 2, 3), dtype=np.float64) + n_points = 0 + n_valid = 0 + + for e in range(n_edges): + if not _edge_is_valid(face_edges_cart, e): + continue + + z0 = face_edges_cart[e, 0, 2] + z1 = face_edges_cart[e, 1, 2] + is_gca = _edge_is_gca(z0, z1, is_GCA_list, is_latlonface, n_valid) + n_valid += 1 + + if not is_gca: + # A constant latitude edge lying on the latitude is itself the whole + # intersection, so it replaces anything the other edges contribute + if _isclose_scalar(z0, latitude_cart, ERROR_TOLERANCE) and _isclose_scalar( + z1, latitude_cart, ERROR_TOLERANCE + ): + for i in range(2): + for j in range(3): + points[i, j] = face_edges_cart[e, i, j] + n_points = 2 + break + continue + + intersections = gca_const_lat_intersection(face_edges_cart[e], latitude_cart) + n_intersections = get_number_of_intersections(intersections) + for r in range(n_intersections): + points[n_points, 0] = intersections[r, 0] + points[n_points, 1] = intersections[r, 1] + points[n_points, 2] = intersections[r, 2] + n_points += 1 + + n_unique = _unique_rows(points, n_points) + unique_intersections = points[:n_unique] + + if n_unique == 0: + raise ValueError( + "No intersections are found for the face, please make sure the " + "build_latlon_box generates the correct results" + ) + # More than two intersections per edge means the face is concave + if n_unique > 2 and n_unique > 2 * n_valid: + raise ValueError( + "UXarray doesn't support concave face with intersections points as currently, please modify your grids accordingly" + ) + + pt_lon_min = np.inf + pt_lon_max = -np.inf + for i in range(n_unique): + lon = _lon_rad_from_xyz( + unique_intersections[i, 0], + unique_intersections[i, 1], + unique_intersections[i, 2], + ) + if lon < pt_lon_min: + pt_lon_min = lon + if lon > pt_lon_max: + pt_lon_max = lon + + return unique_intersections, pt_lon_min, pt_lon_max + + def _get_faces_constLat_intersection_info( face_edges_cart, latitude_cart, is_GCA_list, is_latlonface ): - """Processes each edge of a face polygon in a vectorized manner to - determine overlaps and calculate the intersections for a given latitude and - the faces. + """Processes each edge of a face polygon to determine overlaps and + calculate the intersections for a given latitude and the faces. Parameters: ---------- @@ -381,78 +530,17 @@ def _get_faces_constLat_intersection_info( - pt_lon_min (float): The min longnitude of the interseted intercal in radian if any; otherwise, None.. - pt_lon_max (float): The max longnitude of the interseted intercal in radian, if any; otherwise, None. """ - valid_edges_mask = ~(np.any(face_edges_cart == DUMMY_EDGE_VALUE, axis=(1, 2))) - - # Apply mask to filter out dummy edges - valid_edges = face_edges_cart[valid_edges_mask] - - # Extract Z coordinates for edge determination - edges_z = valid_edges[:, :, 2] - - # Determine if each edge is GCA or constant latitude - is_GCA = _is_edge_gca(is_GCA_list, is_latlonface, edges_z) - - # Check overlap with latitude - overlaps_with_latitude = np.all( - np.isclose(edges_z, latitude_cart, atol=ERROR_TOLERANCE), axis=1 - ) - overlap_flag = np.any(overlaps_with_latitude & ~is_GCA) - - # Identify overlap edges if needed - intersections_pts_list_cart = [] - if overlap_flag: - overlap_index = np.where(overlaps_with_latitude & ~is_GCA)[0][0] - intersections_pts_list_cart.extend(valid_edges[overlap_index]) - else: - # Calculate intersections (assuming a batch-capable intersection function) - for idx, edge in enumerate(valid_edges): - if is_GCA[idx]: - intersections = gca_const_lat_intersection(edge, latitude_cart) - n_intersections = get_number_of_intersections(intersections) - if n_intersections == 0: - continue - elif n_intersections == 2: - intersections_pts_list_cart.extend(intersections) - else: - intersections_pts_list_cart.append(intersections[0]) - - # Find the unique intersection points - unique_intersections = np.unique(intersections_pts_list_cart, axis=0) - - if len(unique_intersections) == 2: - unique_intersection_lonlat = np.array( - [_xyz_to_lonlat_rad(pt[0], pt[1], pt[2]) for pt in unique_intersections] + unique_intersections, pt_lon_min, pt_lon_max = ( + _get_faces_constLat_intersection_info_numba( + face_edges_cart, latitude_cart, is_GCA_list, is_latlonface ) + ) - sorted_lonlat = np.sort(unique_intersection_lonlat, axis=0) - pt_lon_min, pt_lon_max = sorted_lonlat[:, 0] - return unique_intersections, pt_lon_min, pt_lon_max - elif len(unique_intersections) == 1: + if len(unique_intersections) == 1: + # The face is only touched by the latitude, so there is no interval return unique_intersections, None, None - elif len(unique_intersections) != 0 and len(unique_intersections) != 1: - # If the unique intersections numbers is larger than n_edges * 2, then it means the face is concave - if len(unique_intersections) > len(valid_edges) * 2: - raise ValueError( - "UXarray doesn't support concave face with intersections points as currently, please modify your grids accordingly" - ) - else: - # Now return all the intersections points and the pt_lon_min, pt_lon_max - unique_intersection_lonlat = np.array( - [_xyz_to_lonlat_rad(pt[0], pt[1], pt[2]) for pt in unique_intersections] - ) - - sorted_lonlat = np.sort(unique_intersection_lonlat, axis=0) - # Extract the minimum and maximum longitudes - pt_lon_min, pt_lon_max = ( - np.min(sorted_lonlat[:, 0]), - np.max(sorted_lonlat[:, 0]), - ) - return unique_intersections, pt_lon_min, pt_lon_max - elif len(unique_intersections) == 0: - raise ValueError( - "No intersections are found for the face, please make sure the build_latlon_box generates the correct results" - ) + return unique_intersections, pt_lon_min, pt_lon_max @njit(cache=True) From b867747138079408d28968da4c433ab3d405f271 Mon Sep 17 00:00:00 2001 From: cmdupuis3 Date: Wed, 30 Sep 2026 14:07:06 -0500 Subject: [PATCH 2/2] Drop polars from the robust zonal weights and move the face loop into Numba The robust path built several tiny polars DataFrames per candidate face, and that fixed per-call overhead was nearly its whole runtime. Each latitude is now two Numba calls: one computing every candidate face's longitude intervals, and one running the overlap sweep line. Results are bit-identical, and robust zonal means are 52-86x faster, now within 1.3-2.7x of the default path instead of 100-200x slower. Both new per-latitude kernels release the GIL. Each call does milliseconds of work, so from 4 threads the robust weights take 1.09x the single-thread wall time instead of 4.03x. The private helpers now use arrays instead of DataFrames: _zonal_face_weights_robust returns the weights, _get_zonal_face_interval returns (start, end) rows, and _process_overlapped_intervals takes starts/ends/face indices and returns per-face contributions. A face whose crossing longitudes cannot be paired into intervals now raises ValueError rather than polars' ShapeError. Co-Authored-By: Claude Opus 5.5 --- test/grid/integrate/test_zonal_intervals.py | 65 ++- test/grid/integrate/test_zonal_weights.py | 4 +- uxarray/core/zonal.py | 2 +- uxarray/grid/integrate.py | 481 ++++++++++++-------- 4 files changed, 329 insertions(+), 223 deletions(-) diff --git a/test/grid/integrate/test_zonal_intervals.py b/test/grid/integrate/test_zonal_intervals.py index 132a3be44..fcacc4039 100644 --- a/test/grid/integrate/test_zonal_intervals.py +++ b/test/grid/integrate/test_zonal_intervals.py @@ -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 @@ -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) @@ -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(): @@ -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(): @@ -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) @@ -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) @@ -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) @@ -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 = { @@ -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 @@ -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" @@ -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 = { @@ -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, \ @@ -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" @@ -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" diff --git a/test/grid/integrate/test_zonal_weights.py b/test/grid/integrate/test_zonal_weights.py index 858af81b1..efec126ef 100644 --- a/test/grid/integrate/test_zonal_weights.py +++ b/test/grid/integrate/test_zonal_weights.py @@ -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): @@ -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) diff --git a/uxarray/core/zonal.py b/uxarray/core/zonal.py index 2beb3abf3..b3bb75aea 100644 --- a/uxarray/core/zonal.py +++ b/uxarray/core/zonal.py @@ -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) diff --git a/uxarray/grid/integrate.py b/uxarray/grid/integrate.py index 6595fae38..2af876b46 100644 --- a/uxarray/grid/integrate.py +++ b/uxarray/grid/integrate.py @@ -1,13 +1,11 @@ import math import numpy as np -import polars as pl from numba import njit, prange, types from numba.typed import List from uxarray.constants import ERROR_TOLERANCE, INT_FILL_VALUE from uxarray.grid.arcs import compute_arc_length -from uxarray.grid.coordinates import _xyz_to_lonlat_rad from uxarray.grid.intersections import ( gca_const_lat_intersection, get_number_of_intersections, @@ -23,10 +21,10 @@ def _zonal_face_weights_robust( face_latlon_bound_candidate: np.ndarray, is_latlonface: bool = False, is_face_GCA_list: np.ndarray | None = None, -) -> pl.DataFrame: +) -> np.ndarray: """ Utilize the sweep line algorithm to calculate the weight of each face at - a constant latitude, returning a Polars DataFrame. + a constant latitude. Parameters ---------- @@ -45,95 +43,44 @@ def _zonal_face_weights_robust( Returns ------- - weights_df : pl.DataFrame - DataFrame with columns ["face_index", "weight"], containing the per-face weights - (as a fraction of the total length of intersection). + weights : np.ndarray + Shape (n_faces,), the weight of each candidate face as a fraction of the + total length of intersection. """ + n_faces = len(faces_edges_cart_candidate) # Special case: latitude_cart close to +1 or -1 (near poles) if np.isclose(latitude_cart, 1, atol=ERROR_TOLERANCE) or np.isclose( latitude_cart, -1, atol=ERROR_TOLERANCE ): # Evenly distribute weight among candidate faces - n_faces = len(faces_edges_cart_candidate) - weights = {face_index: 1.0 / n_faces for face_index in range(n_faces)} - # Convert dict to Polars DataFrame - return pl.DataFrame( - list(weights.items()), schema=["face_index", "weight"], orient="row" - ) - - intervals_list = [] - - # Iterate over faces - for face_index, face_edges in enumerate(faces_edges_cart_candidate): - # Remove edges that contain INT_FILL_VALUE - face_edges = face_edges[np.all(face_edges != INT_FILL_VALUE, axis=(1, 2))] - - # Which edges are GCA vs constant-lat? - if is_face_GCA_list is not None: - is_GCA_list = is_face_GCA_list[face_index] - else: - is_GCA_list = None - - # Retrieve intervals for the current face - face_interval_df = _get_zonal_face_interval( - face_edges, + return np.ones(n_faces) / n_faces + + starts, ends, face_indices, bad_face = _zonal_face_intervals_numba( + faces_edges_cart_candidate, + latitude_cart, + face_latlon_bound_candidate, + is_face_GCA_list, + is_latlonface, + ) + if bad_face >= 0: + face_edges = faces_edges_cart_candidate[bad_face] + _raise_zonal_face_interval_error( + face_edges[np.all(face_edges != INT_FILL_VALUE, axis=(1, 2))], latitude_cart, - face_latlon_bound_candidate[face_index], - is_latlonface=is_latlonface, - is_GCA_list=is_GCA_list, - ) - - # Check if there are any null values in face_interval_df - has_null = face_interval_df.select(pl.col("*").is_null().any()).row(0)[0] - if has_null: - # Skip this face (only "touched" by the latitude) - continue - - # Check if all start == 0 and all end == 0 - all_start_zero = face_interval_df.select((pl.col("start") == 0).all()).row(0)[0] - all_end_zero = face_interval_df.select((pl.col("end") == 0).all()).row(0)[0] - if all_start_zero and all_end_zero: - # Skip face being merely touched - continue - - # Add each interval row to intervals_list - for row in face_interval_df.iter_rows(named=True): - intervals_list.append( - { - "start": row["start"], - "end": row["end"], - "face_index": face_index, - } - ) - - # Build a Polars DataFrame from intervals - intervals_df = pl.DataFrame(intervals_list) - - # Process intervals to get overlap contributions - try: - overlap_contributions, total_length = _process_overlapped_intervals( - intervals_df + face_latlon_bound_candidate[bad_face], + is_latlonface, + None if is_face_GCA_list is None else is_face_GCA_list[bad_face], ) - # Build final weights dict - weights = {} - n_faces = len(faces_edges_cart_candidate) - for face_index in range(n_faces): - # fraction of total for this face - weights[face_index] = ( - overlap_contributions.get(face_index, 0.0) / total_length - ) - - # Return as Polars DataFrame - weights_df = pl.DataFrame( - list(weights.items()), schema=["face_index", "weight"], orient="row" - ) - return weights_df + overlap_contributions, total_length = _process_overlapped_intervals( + starts, ends, face_indices, n_faces + ) + if total_length == 0.0: + # Every candidate face is only touched by the latitude + raise ZeroDivisionError("float division by zero") - except ValueError: - # If an exception occurs, you can print debug info here if needed - raise + return overlap_contributions / total_length def _get_zonal_face_interval( @@ -142,12 +89,12 @@ def _get_zonal_face_interval( face_latlon_bound: np.ndarray, is_latlonface: bool = False, is_GCA_list: np.ndarray | None = None, -) -> pl.DataFrame: +) -> np.ndarray: """ Processes a face polygon represented by edges in Cartesian coordinates to find intervals where the face intersects with a given latitude. This function handles directed and undirected Great Circle Arcs (GCAs) and edges - at constant latitude, returning a Polars DataFrame with columns ["start", "end"]. + at constant latitude, returning the intervals as (start, end) longitude pairs. Requires the face edges to be sorted in counter-clockwise order, and the span of the face in longitude should be less than pi. Also, all arcs/edges length should be within pi. @@ -174,57 +121,44 @@ def _get_zonal_face_interval( Returns ------- - Intervals_df : pl.DataFrame - a Polars DataFrame with columns ["start", "end"] + intervals : np.ndarray + Shape (n_intervals, 2), the (start, end) longitudes in radians of each + interval, sorted by start. A face only touched by the latitude has the + single interval (0, 0). """ - - face_lon_bound_left, face_lon_bound_right = face_latlon_bound[1] - - try: - unique_intersections, pt_lon_min, pt_lon_max = ( - _get_faces_constLat_intersection_info( - face_edges_cart, latitude_cart, is_GCA_list, is_latlonface - ) - ) - - # If there's exactly one intersection, the face is only "touched" - if len(unique_intersections) == 1: - return pl.DataFrame({"start": [0.0], "end": [0.0]}) - - # Convert intersection points to (lon, lat) in radians - longitudes = np.array( - [_xyz_to_lonlat_rad(*pt.tolist())[0] for pt in unique_intersections] + intervals = np.empty((face_edges_cart.shape[0] + 1, 2)) + n_intervals = _face_zonal_intervals_numba( + face_edges_cart, + latitude_cart, + face_latlon_bound[1], + is_GCA_list, + is_latlonface, + intervals, + ) + if n_intervals < 0: + _raise_zonal_face_interval_error( + face_edges_cart, + latitude_cart, + face_latlon_bound, + is_latlonface, + is_GCA_list, ) - # Handle special wrap-around cases (crossing anti-meridian, etc.) - if face_lon_bound_left >= face_lon_bound_right or ( - face_lon_bound_left == 0 and face_lon_bound_right == 2 * np.pi - ): - if not ( - (pt_lon_max >= np.pi and pt_lon_min >= np.pi) - or (0 <= pt_lon_max <= np.pi and 0 <= pt_lon_min <= np.pi) - ): - if pt_lon_max != 2 * np.pi and pt_lon_min != 0: - # Add wrap-around points - longitudes = np.append(longitudes, [0.0, 2 * np.pi]) - elif pt_lon_max >= np.pi and pt_lon_min == 0: - # If min is 0, but we really need 2*pi - longitudes[longitudes == 0] = 2.0 * np.pi - - # Sort unique longitudes - longitudes = np.unique(longitudes) - longitudes.sort() + return intervals[:n_intervals] - # Pair sorted longitudes into intervals - starts = longitudes[::2] - ends = longitudes[1::2] - # Create Polars DataFrame - intervals_df = pl.DataFrame({"start": starts, "end": ends}) - intervals_df_sorted = intervals_df.sort("start") - - return intervals_df_sorted +def _raise_zonal_face_interval_error( + face_edges_cart, latitude_cart, face_latlon_bound, is_latlonface, is_GCA_list +): + """Raise the error for a face `_face_zonal_intervals_numba` could not process. + The kernel cannot build messages that format the face array, so this re-runs + the face's intersection step in Python, which raises with the full message. + """ + try: + _get_faces_constLat_intersection_info( + face_edges_cart, latitude_cart, is_GCA_list, is_latlonface + ) except ValueError as e: default_print_options = np.get_printoptions() # TODO: what is build_latlon_box? @@ -249,8 +183,17 @@ def _get_zonal_face_interval( np.set_printoptions(**default_print_options) raise + # The intersections are valid, so it is their longitudes that cannot be paired + raise ValueError( + "Found an odd number of intersection longitudes for this face, so they " + "cannot be paired into intervals." + f"\nFace edges cartesian coordinates: {face_edges_cart}" + ) + -def _process_overlapped_intervals(intervals_df: pl.DataFrame): +def _process_overlapped_intervals( + starts: np.ndarray, ends: np.ndarray, face_indices: np.ndarray, n_faces: int +): """Process the overlapped intervals using the sweep line algorithm. This function processes multiple intervals per face using a sweep line algorithm, @@ -260,62 +203,34 @@ def _process_overlapped_intervals(intervals_df: pl.DataFrame): Parameters ---------- - intervals_df : pl.DataFrame - A Polars DataFrame containing the intervals and corresponding face indices. - Required columns: - - start : numeric - Starting position of each interval - - end : numeric - Ending position of each interval - - face_index : int or str - Identifier for the face associated with each interval + starts : np.ndarray + Starting position of each interval. + ends : np.ndarray + Ending position of each interval. + face_indices : np.ndarray + Index of the face each interval belongs to, in ``[0, n_faces)``. + n_faces : int + The number of faces. Returns ------- - tuple[dict, float] + tuple[np.ndarray, float] A tuple containing: - - dict: Maps face indices to their contributions to the total length, + - np.ndarray: Shape (n_faces,), each face's contribution to the total length, where overlapping segments are weighted equally among active faces - float: The total length of all intervals considering their overlaps """ - - events = [] - # Iterate Polars rows as dictionaries - for row in intervals_df.iter_rows(named=True): - events.append((row["start"], "start", row["face_index"])) - events.append((row["end"], "end", row["face_index"])) - - # Sort the events by (position, event_type) - # so that 'start' comes before 'end' if position ties - events.sort(key=lambda x: (x[0], x[1])) - - active_faces = set() - last_position = None - total_length = 0.0 - overlap_contributions = {} - - for position, event_type, face_idx in events: - if last_position is not None and active_faces: - segment_length = position - last_position - # Each face gets an equal share of this segment - segment_weight = segment_length / len(active_faces) - for active_face in active_faces: - overlap_contributions[active_face] = ( - overlap_contributions.get(active_face, 0.0) + segment_weight - ) - total_length += segment_length - - if event_type == "start": - active_faces.add(face_idx) - elif event_type == "end": - if face_idx in active_faces: - active_faces.remove(face_idx) - else: - raise ValueError( - f"Cannot end interval for currently-inactive face_idx {face_idx}, at position {position}." - ) - - last_position = position + overlap_contributions, total_length, bad_row = _process_overlapped_intervals_numba( + np.asarray(starts, dtype=np.float64), + np.asarray(ends, dtype=np.float64), + np.asarray(face_indices, dtype=np.int64), + n_faces, + ) + if bad_row >= 0: + raise ValueError( + f"Cannot end interval for currently-inactive face_idx {face_indices[bad_row]}, " + f"at position {ends[bad_row]}." + ) return overlap_contributions, total_length @@ -546,6 +461,211 @@ def _get_faces_constLat_intersection_info( return unique_intersections, pt_lon_min, pt_lon_max +@njit(cache=True) +def _face_zonal_intervals_numba( + face_edges_cart, + latitude_cart, + face_lon_bounds, + is_GCA_list, + is_latlonface, + intervals, +): + """Numba kernel behind `_get_zonal_face_interval`. + + Writes the face's intervals as (start, end) rows of `intervals`, which needs + one more row than the face has edges, and returns how many there are. Returns + -1 when the face cannot be processed, either because its intersections are + invalid or because their longitudes cannot be paired into intervals; the + Python wrappers re-run such a face to raise the error. + """ + unique_intersections, pt_lon_min, pt_lon_max, n_valid_edges = ( + _get_faces_constLat_intersection_info_numba( + face_edges_cart, latitude_cart, is_GCA_list, is_latlonface + ) + ) + n_unique = unique_intersections.shape[0] + + # The cases `_get_faces_constLat_intersection_info` raises on + if n_unique == 0 or n_unique > 2 * n_valid_edges: + return -1 + + # If there's exactly one intersection, the face is only "touched" + if n_unique == 1: + intervals[0, 0] = 0.0 + intervals[0, 1] = 0.0 + return 1 + + # Room for the two wrap-around points added below + longitudes = np.empty(n_unique + 2) + for i in range(n_unique): + longitudes[i] = _lon_rad_from_xyz( + unique_intersections[i, 0], + unique_intersections[i, 1], + unique_intersections[i, 2], + ) + n_longitudes = n_unique + + # Handle special wrap-around cases (crossing anti-meridian, etc.) + face_lon_bound_left = face_lon_bounds[0] + face_lon_bound_right = face_lon_bounds[1] + if face_lon_bound_left >= face_lon_bound_right or ( + face_lon_bound_left == 0 and face_lon_bound_right == 2 * np.pi + ): + if not ( + (pt_lon_max >= np.pi and pt_lon_min >= np.pi) + or (0 <= pt_lon_max <= np.pi and 0 <= pt_lon_min <= np.pi) + ): + if pt_lon_max != 2 * np.pi and pt_lon_min != 0: + # Add wrap-around points + longitudes[n_longitudes] = 0.0 + longitudes[n_longitudes + 1] = 2 * np.pi + n_longitudes += 2 + elif pt_lon_max >= np.pi and pt_lon_min == 0: + # If min is 0, but we really need 2*pi + for i in range(n_longitudes): + if longitudes[i] == 0: + longitudes[i] = 2.0 * np.pi + + # Pair the sorted unique longitudes into intervals + longitudes = np.unique(longitudes[:n_longitudes]) + if longitudes.shape[0] % 2 != 0: + return -1 + + n_intervals = longitudes.shape[0] // 2 + for i in range(n_intervals): + intervals[i, 0] = longitudes[2 * i] + intervals[i, 1] = longitudes[2 * i + 1] + return n_intervals + + +@njit(cache=True, nogil=True) +def _zonal_face_intervals_numba( + faces_edges_cart, + latitude_cart, + face_latlon_bounds, + is_face_GCA_list, + is_latlonface, +): + """The intervals of every candidate face along a line of constant latitude. + + Returns the intervals in face order as `starts`, `ends` and the `face_indices` + they belong to, leaving out faces only touched by the latitude, along with + the index of the first face that cannot be processed (-1 if there is none). + """ + n_faces = faces_edges_cart.shape[0] + max_intervals = faces_edges_cart.shape[1] + 1 + + face_intervals = np.empty((max_intervals, 2)) + starts = np.empty(n_faces * max_intervals) + ends = np.empty(n_faces * max_intervals) + face_indices = np.empty(n_faces * max_intervals, dtype=np.int64) + n_intervals = 0 + + for face_index in range(n_faces): + if is_face_GCA_list is None: + is_GCA_list = None + else: + is_GCA_list = is_face_GCA_list[face_index] + + n_face_intervals = _face_zonal_intervals_numba( + faces_edges_cart[face_index], + latitude_cart, + face_latlon_bounds[face_index, 1], + is_GCA_list, + is_latlonface, + face_intervals, + ) + if n_face_intervals < 0: + return starts[:0], ends[:0], face_indices[:0], face_index + + # Skip faces being merely touched, where every interval is (0, 0) + touched = True + for i in range(n_face_intervals): + if face_intervals[i, 0] != 0 or face_intervals[i, 1] != 0: + touched = False + if touched: + continue + + for i in range(n_face_intervals): + starts[n_intervals] = face_intervals[i, 0] + ends[n_intervals] = face_intervals[i, 1] + face_indices[n_intervals] = face_index + n_intervals += 1 + + return ( + starts[:n_intervals], + ends[:n_intervals], + face_indices[:n_intervals], + -1, + ) + + +@njit(cache=True, nogil=True) +def _process_overlapped_intervals_numba(starts, ends, face_indices, n_faces): + """Numba kernel behind `_process_overlapped_intervals`. + + Returns each face's contribution and the total length, plus the row of the + first interval whose end is reached while its face is inactive (-1 if there + is none), for the Python wrapper to raise on. + """ + n_rows = starts.shape[0] + + # Each interval contributes a start and an end event, in row order + positions = np.empty(2 * n_rows) + is_start = np.empty(2 * n_rows, dtype=np.int64) + rows = np.empty(2 * n_rows, dtype=np.int64) + for i in range(n_rows): + positions[2 * i] = starts[i] + is_start[2 * i] = 1 + rows[2 * i] = i + positions[2 * i + 1] = ends[i] + is_start[2 * i + 1] = 0 + rows[2 * i + 1] = i + + # Sort the events by position, with ends before starts at the same position + # and otherwise keeping their order: two stable sorts, minor key first + order = np.argsort(is_start, kind="mergesort") + order = order[np.argsort(positions[order], kind="mergesort")] + + overlap_contributions = np.zeros(n_faces) + is_active = np.zeros(n_faces, dtype=np.bool_) + active_faces = np.empty(n_faces, dtype=np.int64) + n_active = 0 + total_length = 0.0 + last_position = 0.0 + + for k in range(2 * n_rows): + event = order[k] + position = positions[event] + if k > 0 and n_active > 0: + segment_length = position - last_position + # Each face gets an equal share of this segment + segment_weight = segment_length / n_active + for a in range(n_active): + overlap_contributions[active_faces[a]] += segment_weight + total_length += segment_length + + face_index = face_indices[rows[event]] + if is_start[event]: + if not is_active[face_index]: + is_active[face_index] = True + active_faces[n_active] = face_index + n_active += 1 + else: + if not is_active[face_index]: + return overlap_contributions, total_length, rows[event] + is_active[face_index] = False + for a in range(n_active): + if active_faces[a] == face_index: + active_faces[a] = active_faces[n_active - 1] + break + n_active -= 1 + + last_position = position + + return overlap_contributions, total_length, -1 + + @njit(cache=True) def _add_edge(edges_list, i0, i1): """Insert an edge into a list of edges in sorted order, ensuring no duplicates. @@ -759,8 +879,7 @@ def _zonal_face_weights( if check_equator: # If near equator, use original approach if np.isclose(z, 0.0, atol=ERROR_TOLERANCE): - overlap_result = _zonal_face_weights_robust(face_edges_xyz, z, face_bounds) - return overlap_result["weight"].to_numpy() + return _zonal_face_weights_robust(face_edges_xyz, z, face_bounds) # Otherwise, use the Numba approach return _zonal_face_weights_util_numba(face_edges_xyz, n_edges_per_face, z)