You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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.

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.15 08:18:18