基于Scipy.stats的3D分箱点集质心邻近点选取优化问询
Great question! Let's break down how to solve both of your problems to make your code cleaner and more efficient:
1. Avoid per-dimension bin statistic calculation
Your original code calls binned_statistic_dd three separate times to compute x, y, z means—but you can calculate all 3D centroids in a single call. The trick is to pass the full 3D point array as the values parameter, then use a custom statistic function to compute the mean vector for each bin. This cuts down on redundant computations and makes your code more concise.
Here's how to adjust it:
import numpy as np import pandas as pd from scipy.stats import binned_statistic_dd def grid_sample_optimized(df): bin_count = int(np.cbrt(1000)) # 10 bins per dimension # Compute 3D bin centroids in one go centroids, edges, binnumber = binned_statistic_dd( df[['x', 'y', 'z']].values, df[['x', 'y', 'z']].values, # Handle empty bins by returning NaNs (we'll drop these later) statistic=lambda bin_points: np.mean(bin_points, axis=0) if len(bin_points) > 0 else [np.nan, np.nan, np.nan], bins=bin_count ) # Flatten the 3D centroid array into a 2D dataframe and remove empty bins centroids_flat = centroids.reshape(-1, 3) xyz_mean = pd.DataFrame(centroids_flat, columns=['x_mean', 'y_mean', 'z_mean']).dropna() # ... add nearest point logic here (we'll cover this next) return xyz_mean
This way, you only run the binning logic once, which saves time especially with large datasets.
2. More efficient nearest centroid point selection
Your current apply approach checks every point for each centroid, which is slow for large datasets (O(N*M) time, where N is the number of centroids and M is the number of points). Let's look at two faster alternatives:
Option 1: Group points by bin (fastest for most cases)
Since you already have bin assignments for every point (from binnumber), you can group points by their bin and only search within each bin for the nearest point to that bin's centroid. This reduces the search space drastically—you don't need to check points outside the bin.
Here's the full optimized function using this method:
def grid_sample_grouped(df): bin_count = int(np.cbrt(1000)) # Get bin assignments for every point _, edges, binnumber = binned_statistic_dd( df[['x', 'y', 'z']].values, values=None, # We only need bin indices here statistic='count', bins=bin_count ) # Add bin number to the dataframe df['bin_id'] = binnumber # Compute centroids for each non-empty bin bin_centroids = df.groupby('bin_id')[['x', 'y', 'z']].mean().dropna() # For each bin, find the point closest to its centroid def find_nearest_in_bin(group): centroid = bin_centroids.loc[group.name] # Calculate Euclidean distance from each point in the bin to the centroid distances = np.linalg.norm(group[['x', 'y', 'z']].values - centroid.values, axis=1) # Return the point with the smallest distance return group.iloc[distances.argmin()] # Apply the function to each bin group and clean up the result sampled_points = df.groupby('bin_id').apply(find_nearest_in_bin).reset_index(drop=True) sampled_points = sampled_points.drop(columns=['bin_id']) return sampled_points
This method runs in O(M) time (since every point is processed once) and is the most efficient choice for large datasets.
Option 2: Use KDTree for batch nearest neighbor search
If you prefer a more general-purpose nearest neighbor approach, you can use scipy.spatial.KDTree to build an index of your points, then batch-query all centroids at once. This is faster than your original apply method (O(M log N) time) and works well if you need to query multiple sets of centroids later.
Here's how to implement it:
from scipy.spatial import KDTree def grid_sample_kdtree(df): bin_count = int(np.cbrt(1000)) # Compute centroids (using the single-call method from earlier) centroids, edges, binnumber = binned_statistic_dd( df[['x', 'y', 'z']].values, df[['x', 'y', 'z']].values, statistic=lambda bin_points: np.mean(bin_points, axis=0) if len(bin_points) > 0 else [np.nan, np.nan, np.nan], bins=bin_count ) centroids_flat = centroids.reshape(-1, 3) xyz_mean = pd.DataFrame(centroids_flat, columns=['x_mean', 'y_mean', 'z_mean']).dropna() # Build a KDTree from the original points point_tree = KDTree(df[['x', 'y', 'z']].values) # Find the nearest point for each centroid (k=1 means closest single point) _, nearest_indices = point_tree.query(xyz_mean.values, k=1) # Extract the nearest points from the original dataframe sampled_points = df.iloc[nearest_indices].reset_index(drop=True) return sampled_points
Which option should you choose?
- Group by bin: Best for your specific use case, as it leverages the binning structure to minimize computation.
- KDTree: Better if you need to reuse the nearest neighbor index for other queries, or if your binning logic changes later.
内容的提问来源于stack exchange,提问作者Tom Hemmes

