diff --git a/nomad/stop_detection/density_algs.py b/nomad/stop_detection/density_algs.py index c82aa987..f5ff9138 100644 --- a/nomad/stop_detection/density_algs.py +++ b/nomad/stop_detection/density_algs.py @@ -206,6 +206,7 @@ def ta_dbscan( passthrough_cols=None, keep_col_names=True, traj_cols=None, + passthrough_agg=None, **kwargs ): """ @@ -227,6 +228,8 @@ def ta_dbscan( Include extra stats if True (default: False). passthrough_cols : list, optional Columns to retain per stop. + passthrough_agg : dict, optional + Aggregation functions for selected passthrough columns. traj_cols : dict, optional Mapping for column names. **kwargs @@ -279,6 +282,7 @@ def ta_dbscan( complete_output=complete_output, dur_min=dur_min, passthrough_cols=passthrough_cols, + passthrough_agg=passthrough_agg, keep_col_names=keep_col_names, traj_cols=traj_cols, **kwargs, @@ -295,6 +299,7 @@ def ta_dbscan_per_user( traj_cols=None, n_jobs=1, print_progress=False, + passthrough_agg=None, **kwargs ): """ @@ -320,6 +325,7 @@ def ta_dbscan_per_user( "dur_min": dur_min, "complete_output": complete_output, "passthrough_cols": pt_cols, + "passthrough_agg": passthrough_agg, "traj_cols": traj_cols, **kwargs, }, @@ -536,6 +542,7 @@ def dbstop( passthrough_cols=None, keep_col_names=True, traj_cols=None, + passthrough_agg=None, **kwargs ): """ @@ -557,6 +564,8 @@ def dbstop( Include extra stats if True (default: False). passthrough_cols : list, optional Columns to retain per stop. + passthrough_agg : dict, optional + Aggregation functions for selected passthrough columns. traj_cols : dict, optional Mapping for column names. **kwargs @@ -609,6 +618,7 @@ def dbstop( complete_output=complete_output, dur_min=dur_min, passthrough_cols=passthrough_cols, + passthrough_agg=passthrough_agg, keep_col_names=keep_col_names, traj_cols=traj_cols, **kwargs, @@ -626,6 +636,7 @@ def dbstop_per_user( traj_cols=None, n_jobs=1, print_progress=False, + passthrough_agg=None, **kwargs ): """ @@ -651,6 +662,7 @@ def dbstop_per_user( "dur_min": dur_min, "complete_output": complete_output, "passthrough_cols": pt_cols, + "passthrough_agg": passthrough_agg, "keep_col_names": keep_col_names, "traj_cols": traj_cols, **kwargs, @@ -967,6 +979,7 @@ def seqscan( passthrough_cols=None, keep_col_names=True, traj_cols=None, + passthrough_agg=None, **kwargs ): """ @@ -988,6 +1001,8 @@ def seqscan( Include extra stats if True (default: False). passthrough_cols : list, optional Columns to retain per stop. + passthrough_agg : dict, optional + Aggregation functions for selected passthrough columns. traj_cols : dict, optional Mapping for column names. **kwargs @@ -1039,6 +1054,7 @@ def seqscan( labels, complete_output=complete_output, passthrough_cols=passthrough_cols, + passthrough_agg=passthrough_agg, keep_col_names=keep_col_names, traj_cols=traj_cols, **kwargs, @@ -1057,6 +1073,7 @@ def seqscan_per_user( traj_cols=None, n_jobs=1, print_progress=False, + passthrough_agg=None, **kwargs ): """ @@ -1082,6 +1099,7 @@ def seqscan_per_user( "dur_min": dur_min, "complete_output": complete_output, "passthrough_cols": pt_cols, + "passthrough_agg": passthrough_agg, "keep_col_names": keep_col_names, "traj_cols": traj_cols, **kwargs, @@ -1932,6 +1950,7 @@ def st_hdbscan( complete_output=False, passthrough_cols=None, traj_cols=None, + passthrough_agg=None, **kwargs ): """ @@ -1953,6 +1972,8 @@ def st_hdbscan( If True, include extra stats. passthrough_cols : list, optional Columns to passthrough to final stop table + passthrough_agg : dict, optional + Aggregation functions for selected passthrough columns. traj_cols : dict, optional Mapping for key columns. **kwargs @@ -1989,6 +2010,7 @@ def st_hdbscan( labels, complete_output=complete_output, passthrough_cols=passthrough_cols, + passthrough_agg=passthrough_agg, keep_col_names=True, traj_cols=traj_cols, **kwargs, @@ -2005,6 +2027,7 @@ def st_hdbscan_per_user( traj_cols=None, n_jobs=1, print_progress=False, + passthrough_agg=None, **kwargs ): """ @@ -2031,6 +2054,7 @@ def st_hdbscan_per_user( "dur_min": dur_min, "complete_output": complete_output, "passthrough_cols": pt_cols, + "passthrough_agg": passthrough_agg, "traj_cols": traj_cols, **kwargs, }, diff --git a/nomad/stop_detection/sequential_algs.py b/nomad/stop_detection/sequential_algs.py index b4dc3826..6ca411ce 100644 --- a/nomad/stop_detection/sequential_algs.py +++ b/nomad/stop_detection/sequential_algs.py @@ -129,6 +129,7 @@ def detect_stops( passthrough_cols=None, keep_col_names=True, traj_cols=None, + passthrough_agg=None, **kwargs ): """ @@ -152,6 +153,8 @@ def detect_stops( If True, include additional summary statistics in output. passthrough_cols : list, optional Columns to retain (and summarize/propagate) per stop. + passthrough_agg : dict, optional + Aggregation functions for selected passthrough columns. keep_col_names : bool, default True Whether to keep original column names in output. traj_cols : dict, optional @@ -193,6 +196,7 @@ def detect_stops( labels, complete_output=complete_output, passthrough_cols=passthrough_cols, + passthrough_agg=passthrough_agg, keep_col_names=keep_col_names, traj_cols=traj_cols, **kwargs, @@ -211,6 +215,7 @@ def detect_stops_per_user( traj_cols=None, n_jobs=1, print_progress=False, + passthrough_agg=None, **kwargs ): """ @@ -232,6 +237,8 @@ def detect_stops_per_user( If True, include additional summary statistics in output. passthrough_cols : list, optional Columns to retain (and summarize/propagate) per stop. + passthrough_agg : dict, optional + Aggregation functions for selected passthrough columns. keep_col_names : bool, default True Whether to keep original column names in output. traj_cols : dict, optional @@ -271,6 +278,7 @@ def detect_stops_per_user( "method": method, "complete_output": complete_output, "passthrough_cols": pt_cols, + "passthrough_agg": passthrough_agg, "keep_col_names": keep_col_names, "traj_cols": traj_cols, **kwargs, @@ -497,6 +505,7 @@ def lachesis( postprocessing=None, eps=None, traj_cols=None, + passthrough_agg=None, **kwargs ): """ @@ -518,6 +527,8 @@ def lachesis( Passed along to the column‐detection helper. passthrough_cols : list, optional Columns to retain (and summarize/propagate) per stop. + passthrough_agg : dict, optional + Aggregation functions for selected passthrough columns. postprocessing : {None, 'dbscan'}, optional Optional stop postprocessing method. eps : float, optional @@ -559,6 +570,7 @@ def lachesis( labels, complete_output=complete_output, passthrough_cols=passthrough_cols, + passthrough_agg=passthrough_agg, keep_col_names=keep_col_names, traj_cols=traj_cols, **kwargs, @@ -577,6 +589,7 @@ def lachesis_per_user( traj_cols=None, n_jobs=1, print_progress=False, + passthrough_agg=None, **kwargs ): """ @@ -596,6 +609,8 @@ def lachesis_per_user( If True, include additional summary statistics in output. passthrough_cols : list, optional Columns to retain (and summarize/propagate) per stop. + passthrough_agg : dict, optional + Aggregation functions for selected passthrough columns. postprocessing : {None, 'dbscan'}, optional Optional stop postprocessing method applied separately to each user. eps : float, optional @@ -639,6 +654,7 @@ def lachesis_per_user( "dur_min": dur_min, "complete_output": complete_output, "passthrough_cols": pt_cols, + "passthrough_agg": passthrough_agg, "postprocessing": postprocessing, "eps": eps, "traj_cols": traj_cols, @@ -784,6 +800,7 @@ def grid_based( complete_output=False, passthrough_cols=None, traj_cols=None, + passthrough_agg=None, **kwargs ): """ @@ -801,6 +818,10 @@ def grid_based( Minimum duration in minutes for a valid stop. Default is 5. complete_output : bool, optional If True, include additional stop statistics in the output. + passthrough_cols : list, optional + Columns to retain per stop. + passthrough_agg : dict, optional + Aggregation functions for selected passthrough columns. traj_cols : dict, optional Mapping for 'timestamp', 'datetime', or 'location_id' column names. **kwargs @@ -844,6 +865,7 @@ def grid_based( traj_cols=traj_cols, keep_col_names=True, passthrough_cols=passthrough_cols, + passthrough_agg=passthrough_agg, **kwargs ), include_groups=False @@ -864,6 +886,7 @@ def grid_based_per_user( traj_cols=None, n_jobs=1, print_progress=False, + passthrough_agg=None, **kwargs ): """ @@ -892,6 +915,7 @@ def grid_based_per_user( "dur_min": dur_min, "complete_output": complete_output, "passthrough_cols": pt_cols, + "passthrough_agg": passthrough_agg, "traj_cols": traj_cols, **kwargs, }, diff --git a/nomad/stop_detection/utils.py b/nomad/stop_detection/utils.py index 97e1c35c..f0d6612d 100644 --- a/nomad/stop_detection/utils.py +++ b/nomad/stop_detection/utils.py @@ -402,6 +402,7 @@ def summarize_stop_grid( keep_col_names=True, passthrough_cols=None, traj_cols=None, + passthrough_agg=None, **kwargs ): """ @@ -418,6 +419,9 @@ def summarize_stop_grid( if True, they use the user's time‐column name. passthrough_cols : list[str], optional Additional columns (e.g. 'user_id') to carry through. + passthrough_agg : dict, optional + Pandas-compatible aggregation function for each passthrough column that + should not use its first value. traj_cols : dict, optional Column‐name overrides. @@ -432,6 +436,8 @@ def summarize_stop_grid( """ if passthrough_cols is None: passthrough_cols = [] + if passthrough_agg is None: + passthrough_agg = {} # 1) pick time key t_key, use_datetime = loader._fallback_time_cols_dt(grouped_data.columns, traj_cols, kwargs) @@ -486,7 +492,11 @@ def summarize_stop_grid( for c in to_pass: if c in grouped_data.columns: - out[c] = grouped_data[c].iloc[0] + out[c] = ( + grouped_data[c].agg(passthrough_agg[c]) + if c in passthrough_agg + else grouped_data[c].iloc[0] + ) return pd.Series(out, dtype='object') diff --git a/nomad/tests/test_stop_detection.py b/nomad/tests/test_stop_detection.py index 8b7b359b..dc759aa3 100644 --- a/nomad/tests/test_stop_detection.py +++ b/nomad/tests/test_stop_detection.py @@ -5,6 +5,8 @@ from pathlib import Path import nomad.io.base as loader from nomad import filters +import nomad.stop_detection.density_algs as density_algs +import nomad.stop_detection.sequential_algs as sequential_algs from nomad.stop_detection.density_algs import ( dbstop, dbstop_labels, @@ -30,6 +32,7 @@ detect_stops_per_user, grid_based, grid_based_labels, + grid_based_per_user, lachesis, lachesis_labels, lachesis_labels_per_user, @@ -397,6 +400,21 @@ def simple_traj(stop_test_params): df["tz_offset"] = 0 return df +@pytest.fixture +def passthrough_traj(): + return pd.DataFrame({ + "user_id": ["a", "a", "a", "b", "b", "b"], + "timestamp": [0, 600, 1200, 0, 600, 1200], + "x": [0.0, 0.0, 0.0, 10.0, 10.0, 10.0], + "y": [0.0, 0.0, 0.0, 10.0, 10.0, 10.0], + "h3_cell": ["cell-a"] * 3 + ["cell-b"] * 3, + "location_id": pd.Series( + ["first-a", "mode-a", "mode-a", "first-b", "mode-b", "mode-b"], + dtype="string", + ), + "place_code": pd.Series([1, 2, 2, 3, 4, 4], dtype="Int64"), + }) + @pytest.fixture(scope="module") def agent_traj_ground_truth(): test_dir = Path(__file__).resolve().parent @@ -1671,3 +1689,97 @@ def test_st_hdbscan_delta_roam(hdbscan_traj): delta_roam=50, traj_cols=traj_cols ) assert isinstance(stops, pd.DataFrame) + + +@pytest.mark.parametrize( + ("module", "stop_name", "label_name", "algorithm_kwargs"), + [ + (sequential_algs, "detect_stops", "detect_stops_labels", {"delta_roam": 100, "dt_max": 60}), + (sequential_algs, "lachesis", "lachesis_labels", {"delta_roam": 100, "dt_max": 60}), + (density_algs, "ta_dbscan", "ta_dbscan_labels", {"dist_thresh": 100, "min_pts": 2, "time_thresh": 60}), + (density_algs, "dbstop", "dbstop_labels", {"dist_thresh": 100, "min_pts": 2, "time_thresh": 60}), + (density_algs, "seqscan", "seqscan_labels", {"dist_thresh": 100, "min_pts": 2, "time_thresh": 60}), + (density_algs, "st_hdbscan", "hdbscan_labels", {"time_thresh": 60}), + ], +) +def test_direct_stop_apis_forward_passthrough_agg_only_to_summarization( + passthrough_traj, + monkeypatch, + module, + stop_name, + label_name, + algorithm_kwargs, +): + label_kwargs = [] + + def labels(data, **kwargs): + label_kwargs.append(kwargs) + return pd.Series(0, index=data.index, name="cluster") + + monkeypatch.setattr(module, label_name, labels) + stops = getattr(module, stop_name)( + passthrough_traj.iloc[:3], + dur_min=0, + passthrough_cols=["location_id"], + passthrough_agg={"location_id": lambda values: values.mode().iloc[0]}, + **algorithm_kwargs, + ) + + assert stops["location_id"].tolist() == ["mode-a"] + assert "passthrough_agg" not in label_kwargs[0] + + +def test_per_user_stop_api_forwards_nullable_passthrough_agg(passthrough_traj): + stops = detect_stops_per_user( + passthrough_traj, + delta_roam=100, + dt_max=60, + dur_min=0, + passthrough_cols=["location_id", "place_code"], + passthrough_agg={ + "location_id": lambda values: values.mode().iloc[0], + "place_code": lambda values: values.mode().iloc[0], + }, + n_jobs=1, + ) + + assert stops["location_id"].tolist() == ["mode-a", "mode-b"] + assert stops["place_code"].tolist() == [2, 4] + assert str(stops["place_code"].dtype) == "Int64" + + +def test_grid_based_applies_passthrough_agg_and_preserves_default(passthrough_traj): + kwargs = { + "time_thresh": 60, + "min_cluster_size": 2, + "dur_min": 0, + "passthrough_cols": ["location_id"], + "traj_cols": {"location_id": "h3_cell"}, + } + + default_stops = grid_based(passthrough_traj.iloc[:3], **kwargs) + modal_stops = grid_based( + passthrough_traj.iloc[:3], + passthrough_agg={"location_id": lambda values: values.mode().iloc[0]}, + **kwargs, + ) + + assert default_stops["location_id"].tolist() == ["first-a"] + assert modal_stops["location_id"].tolist() == ["mode-a"] + + +def test_grid_based_per_user_preserves_passthrough_schema_for_all_noise(passthrough_traj): + stops = grid_based_per_user( + passthrough_traj, + time_thresh=60, + min_cluster_size=4, + dur_min=0, + passthrough_cols=["location_id"], + passthrough_agg={"location_id": lambda values: values.mode().iloc[0]}, + traj_cols={"location_id": "h3_cell"}, + n_jobs=1, + ) + + assert stops.empty + assert "location_id" in stops.columns + assert "user_id" in stops.columns diff --git a/nomad/tests/test_visit_location_detection.py b/nomad/tests/test_visit_location_detection.py index d5b6bd1e..c195f2e5 100644 --- a/nomad/tests/test_visit_location_detection.py +++ b/nomad/tests/test_visit_location_detection.py @@ -1,8 +1,11 @@ import geopandas as gpd +import h3 import pandas as pd import pytest +from shapely.geometry import LineString, Point -from nomad.visit_attribution.visit_attribution import detect_locations +import nomad.visit_attribution.visit_attribution as visit_attribution +from nomad.visit_attribution.visit_attribution import detect_locations, poi_map @pytest.fixture @@ -218,3 +221,148 @@ def test_detect_locations_does_not_mutate_input(cartesian_points): detect_locations(cartesian_points, epsilon=2) pd.testing.assert_frame_equal(cartesian_points, original) + + +def test_poi_map_attributes_unique_h3_cells_and_preserves_alignment(monkeypatch): + poi_cell = h3.latlng_to_cell(39.95, -75.16, 10) + adjacent_cell = next(iter(h3.grid_ring(poi_cell, 1))) + unmatched_cell = next(iter(h3.grid_ring(poi_cell, 2))) + latitude, longitude = h3.cell_to_latlng(poi_cell) + stops = pd.DataFrame( + {"h3_cell": [poi_cell, poi_cell, adjacent_cell, unmatched_cell, pd.NA]}, + index=[3, 3, 7, 9, 11], + ) + pois = gpd.GeoDataFrame( + {"building_id": ["library"]}, + geometry=[Point(longitude, latitude).buffer(0.000001)], + crs="EPSG:4326", + ) + batch_sizes = [] + grid_disk_distances = visit_attribution.h3ronpy.grid_disk_distances + + def record_batch_size(cells, max_distance): + batch_sizes.append(len(cells)) + return grid_disk_distances(cells, max_distance) + + monkeypatch.setattr( + visit_attribution.h3ronpy, + "grid_disk_distances", + record_batch_size, + ) + + locations = poi_map( + stops, + pois, + max_distance=1, + location_id="building_id", + ) + + assert batch_sizes == [3] + assert locations.index.tolist() == [3, 3, 7, 9, 11] + assert locations.name == "building_id" + assert locations.iloc[:3].tolist() == ["library", "library", "library"] + assert locations.iloc[3:].isna().all() + + +def test_poi_map_h3_supports_projected_pois_and_column_overrides(): + h3_cell = h3.latlng_to_cell(39.95, -75.16, 10) + latitude, longitude = h3.cell_to_latlng(h3_cell) + stops = pd.DataFrame({"containment_area": [h3_cell]}) + pois = gpd.GeoDataFrame( + {"place": [42]}, + geometry=[Point(longitude, latitude).buffer(0.000001)], + crs="EPSG:4326", + ).to_crs("EPSG:3857") + + with pytest.warns(UserWarning, match="Reprojecting for H3 attribution"): + locations = poi_map( + stops, + pois, + location_id="place", + traj_cols={"h3_cell": "containment_area"}, + ) + + assert locations.tolist() == [42] + assert locations.name == "place" + + +def test_poi_map_h3_breaks_equidistant_ties_by_poi_order(): + stop_cell = h3.latlng_to_cell(39.95, -75.16, 10) + poi_cells = list(h3.grid_ring(stop_cell, 1))[:2] + centers = [h3.cell_to_latlng(cell) for cell in poi_cells] + pois = gpd.GeoDataFrame( + {"location_id": ["first", "second"]}, + geometry=[ + Point(longitude, latitude).buffer(0.000001) + for latitude, longitude in centers + ], + crs="EPSG:4326", + ) + + locations = poi_map( + pd.DataFrame({"h3_cell": [stop_cell]}), + pois, + max_distance=1, + location_id="location_id", + ) + + assert locations.tolist() == ["first"] + + +def test_poi_map_h3_returns_aligned_empty_result(): + pois = gpd.GeoDataFrame( + {"location_id": ["unused"]}, + geometry=[Point(-75.16, 39.95).buffer(0.000001)], + crs="EPSG:4326", + ) + stops = pd.DataFrame({"cell": pd.Series(dtype="string")}) + + locations = poi_map( + stops, + pois, + location_id="location_id", + traj_cols={"h3_cell": "cell"}, + ) + + assert locations.empty + assert locations.index.equals(stops.index) + assert locations.name == "location_id" + + +def test_poi_map_h3_maps_multi_cell_poi_and_falls_back_to_index(): + first_cell = h3.latlng_to_cell(39.95, -75.16, 10) + second_cell = next(iter(h3.grid_ring(first_cell, 1))) + centers = [h3.cell_to_latlng(cell) for cell in [first_cell, second_cell]] + poi = gpd.GeoDataFrame( + geometry=[ + LineString([ + (longitude, latitude) for latitude, longitude in centers + ]).buffer(0.000001) + ], + index=pd.Index(["building-a"]), + crs="EPSG:4326", + ) + + with pytest.warns(UserWarning, match="using poi_table.index"): + locations = poi_map( + pd.DataFrame({"h3_cell": [first_cell, second_cell]}), + poi, + ) + + assert locations.tolist() == ["building-a", "building-a"] + assert locations.name == "location_id" + + +def test_poi_map_h3_requires_one_resolution_per_call(): + cells = [ + h3.latlng_to_cell(39.95, -75.16, 9), + h3.latlng_to_cell(39.95, -75.16, 10), + ] + poi = gpd.GeoDataFrame( + geometry=[Point(-75.16, 39.95).buffer(0.000001)], + crs="EPSG:4326", + ) + + with pytest.warns(UserWarning, match="using poi_table.index"): + with pytest.raises(ValueError, match="same resolution"): + poi_map(pd.DataFrame({"h3_cell": cells}), poi) diff --git a/nomad/visit_attribution/visit_attribution.py b/nomad/visit_attribution/visit_attribution.py index 13eefc33..cce257f2 100644 --- a/nomad/visit_attribution/visit_attribution.py +++ b/nomad/visit_attribution/visit_attribution.py @@ -1,7 +1,11 @@ import geopandas as gpd +import h3ronpy +import h3ronpy.vector as h3vector import warnings import pandas as pd import pyproj +import pyarrow as pa +import pyarrow.compute as pc import numpy as np from sklearn.cluster import DBSCAN import nomad.io.base as loader @@ -137,19 +141,24 @@ def point_in_polygon(data, poi_table, method='centroid', data_crs=None, max_dist # change to point_in_polygon, move to filters.py def poi_map(data, poi_table, max_distance=0, data_crs=None, location_id=None, traj_cols=None, **kwargs): """ - Assign each point in `data` to a polygon in `poi_table`, using containment when - `max_distance==0` or the nearest neighbor within `max_distance` otherwise. + Assign each point or H3 containment area in `data` to a polygon in `poi_table`. + + Points use geometric containment when `max_distance==0` or the nearest neighbor + within `max_distance` otherwise. H3 cells use intersecting POI cells or the POI + cell with the smallest H3 grid distance. Parameters ---------- data : pd.DataFrame or gpd.GeoDataFrame - Input points, either as a DataFrame with coordinate columns or a GeoDataFrame. + Input points, either as a DataFrame with coordinate columns or a GeoDataFrame, + or a table containing H3 cells. poi_table : gpd.GeoDataFrame Polygons to match against, indexed or with `location_id` column. traj_cols : list of str, optional Names of the coordinate columns in `data` when it is a DataFrame. max_distance : float, default 0 - Maximum search radius for nearest‐neighbor matching; zero invokes a point‐in‐polygon test. + Maximum search radius for nearest-neighbor matching. For H3 input, this is + the maximum grid distance in cells. data_crs : str or pyproj.CRS, optional CRS for `data` if it is a DataFrame; ignored for GeoDataFrames. location_id : str, optional @@ -160,15 +169,83 @@ def poi_map(data, poi_table, max_distance=0, data_crs=None, location_id=None, tr Returns ------- pd.Series - Indexed like `data`, with each entry set to the matching polygon’s ID (from - `location_id` or `poi_table.index`). Points not contained or beyond `max_distance` - yield NaN. When multiple polygons overlap a point, only the first match is kept. + Indexed like `data`, with each entry set to the matching polygon's ID (from + `location_id` or `poi_table.index`). Points or cells not contained or beyond + `max_distance` yield NaN. Ties retain the first POI in `poi_table`. """ # column name handling traj_cols = loader._parse_traj_cols(data.columns, traj_cols, kwargs, defaults={}) - + if poi_table.crs is None: raise ValueError(f"poi_table must have crs attribute for spatial join.") + + if location_id is None and "location_id" in traj_cols: + location_id = traj_cols["location_id"] + + out_col = location_id if location_id is not None else "location_id" + # choose where IDs come from: poi_table column (if it exists) else poi_table.index + use_col = (location_id is not None) and (location_id in poi_table.columns) + + if location_id is None: + warnings.warn("location_id not provided; using poi_table.index for spatial join.") + elif not use_col: + warnings.warn(f"{location_id} column not found in poi_table; using poi_table.index for spatial join.") + + h3_col = traj_cols.get("h3_cell", "h3_cell") + if h3_col in data.columns: + locations = pd.Series(index=data.index, name=out_col, dtype="object") + unique_cells = pd.Series(data[h3_col].dropna().unique(), dtype="string") + if unique_cells.empty or poi_table.empty: + return locations + + parsed_cells = pa.array(h3ronpy.cells_parse(pa.array(unique_cells, type=pa.string()))) + resolutions = pc.unique(pa.array(h3ronpy.cells_resolution(parsed_cells))).to_pylist() + if len(resolutions) != 1: + raise ValueError("All h3_cell values must have the same resolution.") + + h3_crs = pyproj.CRS("EPSG:4326") + if not h3_crs.equals(pyproj.CRS(poi_table.crs)): + poi_table = poi_table.to_crs(h3_crs) + warnings.warn("CRS for `poi_table` is not EPSG:4326. Reprojecting for H3 attribution...") + + poi_cell_lists = pa.array(h3vector.wkb_to_cells( + pa.array(poi_table.geometry.to_wkb()), + resolutions[0], + containment_mode=h3ronpy.ContainmentMode.Covers, + )) + poi_cells = pc.list_flatten(poi_cell_lists) + if len(poi_cells) == 0: + return locations + + poi_positions = pc.list_parent_indices(poi_cell_lists).to_numpy() + poi_ids = ( + poi_table[location_id] if use_col else pd.Series(poi_table.index) + ).reset_index(drop=True) + poi_coverage = pd.DataFrame({ + "_candidate_cell": poi_cells.to_numpy(), + "_poi_position": poi_positions, + out_col: poi_ids.iloc[poi_positions].to_numpy(), + }) + + disk_distances = pa.table(h3ronpy.grid_disk_distances(parsed_cells, max_distance)) + candidate_lists = disk_distances["cell"].combine_chunks() + candidates = pd.DataFrame({ + "_input_position": pc.list_parent_indices(candidate_lists).to_numpy(), + "_candidate_cell": pc.list_flatten(candidate_lists).to_numpy(), + "_distance": pc.list_flatten( + disk_distances["k"].combine_chunks() + ).to_numpy(), + }) + matches = candidates.merge(poi_coverage, on="_candidate_cell") + if matches.empty: + return locations + + matches = matches.sort_values( + ["_input_position", "_distance", "_poi_position"], kind="stable" + ).drop_duplicates("_input_position") + lookup = pd.Series(index=unique_cells, dtype="object") + lookup.iloc[matches["_input_position"].to_numpy()] = matches[out_col].to_numpy() + return data[h3_col].astype("string").map(lookup).rename(out_col) # Determine which geometry to use if isinstance(data, gpd.GeoDataFrame): @@ -215,15 +292,6 @@ def poi_map(data, poi_table, max_distance=0, data_crs=None, location_id=None, tr poi_table = poi_table.to_crs(data_crs) warnings.warn("CRS for `poi_table` does not match crs for `data`. Reprojecting...") - out_col = location_id if location_id is not None else "location_id" - # choose where IDs come from: poi_table column (if it exists) else poi_table.index - use_col = (location_id is not None) and (location_id in poi_table.columns) - - if location_id is None: - warnings.warn("location_id not provided; using poi_table.index for spatial join.") - elif not use_col: - warnings.warn(f"{location_id} column not found in poi_table; using poi_table.index for spatial join.") - if max_distance>0: if data_crs.is_geographic: warnings.warn(f"Provided CRS {data_crs.name} is a geographic coordinate system. " diff --git a/setup.py b/setup.py index 7da959a0..47c225c2 100644 --- a/setup.py +++ b/setup.py @@ -34,6 +34,7 @@ 'pyarrow', 's3fs', 'h3', + 'h3ronpy', 'pydeck' ],