如何设置Ball Tree以返回指定类型的有效近邻?
问题描述
我有一个包含多地理位置的DataFrame,想用Ball Tree实现最近邻查找,但输出存在两个问题:
- 目标位置自身会被列为近邻(例如位置A的近邻列表里包含A)
- 同一ID不同时间的记录也会被列为该ID的近邻(例如A在不同时间的条目会出现在A的近邻里)
之前在小数据集上是返回更多近邻后再过滤,但大数据集下计算成本太高,想问有没有办法直接定义Ball Tree返回的邻居类型?
测试代码:
from sklearn.neighbors import BallTree import numpy as np import pandas as pd test_data = pd.DataFrame({'latitude':[51.51, 51.52,61.53,61.54,71.55, 71.56, 51.51, 51.52,61.53,61.54,71.55, 71.56, 51.51, 51.52,61.53,61.54,71.55, 71.56], 'longitude':[-0.13,-0.13,-0.13,-0.14,-0.13,-0.13, -0.13,-0.13,-0.13,-0.14,-0.13,-0.13, -0.13,-0.13,-0.13,-0.14,-0.13,-0.13], 'id':['A','B','C','D','E','F', 'A','B','C','D','E','F', 'A','B','C','D','E','F'], 'target':[35,410,1,100,114,78, 14,254,101,278,3578,435, 254,254,37,47,38,101], 'time':['2019-03-10 11:00:00','2019-03-10 11:00:00','2019-03-10 11:00:00','2019-03-10 11:00:00','2019-03-10 11:00:00','2019-03-10 11:00:00', '2019-03-10 11:10:00','2019-03-10 11:10:00','2019-03-10 11:10:00','2019-03-10 11:10:00','2019-03-10 11:10:00','2019-03-10 11:10:00', '2019-03-10 11:20:00','2019-03-10 11:20:00','2019-03-10 11:20:00','2019-03-10 11:20:00','2019-03-10 11:20:00','2019-03-10 11:20:00', ]}) # --- 步骤1:数据准备 test_data=test_data.reset_index() # 将纬度和经度转换为弧度 for column in test_data[['latitude','longitude']]: rad = np.deg2rad(test_data[column].values) test_data[f'{column}'] = rad # 创建时间列的副本,其中一个将设为索引 test_data['time2']=test_data['time'] # 将时间转换为datetime类型 test_data['time']=pd.to_datetime(test_data['time']) test_data = test_data.set_index('time').astype('str') # --- 步骤2:训练Ball Tree locations_a = test_data locations_b = test_data col_name = 'ss_id' latitude = "latitude" longitude = "longitude" # 创建Ball Tree ball = BallTree(locations_a[[latitude, longitude]].values, metric='haversine') # 设置返回的近邻数量 k = 6 # 计算距离 distances, indices = ball.query(locations_b[[latitude, longitude]].values, k = k) # --- 步骤3:将结果合并到DataFrame dists = pd.DataFrame(distances).stack() rel = pd.DataFrame(indices).stack() # 创建结果DataFrame neighbor_info_df = pd.merge(dists.rename('distance'), rel.rename('neighbor_idx'), right_index=True, left_index=True) # 重置并重命名索引 neighbor_info_df = neighbor_info_df.reset_index(level=1).rename({'level_1': 'neighbor_number'}, axis=1) neighbor_info_df = neighbor_info_df.reset_index().rename({'index': 'id_index_no'}, axis=1) neighbor_info_df.head(10)
解决方案
Ball Tree本身没有直接配置过滤规则的参数,但可以通过以下两种高效方式解决问题,避免返回过多无效近邻后再过滤:
1. 按ID分组构建Ball Tree,仅跨组查询
如果需求是同一ID的所有记录都不能互为近邻,可以先按ID分组,为每个ID的记录构建仅包含其他ID位置的候选数据集,这样Ball Tree查询时只会返回其他ID的近邻,从根源避免无效结果。
示例代码:
from sklearn.neighbors import BallTree import numpy as np import pandas as pd test_data = pd.DataFrame({'latitude':[51.51, 51.52,61.53,61.54,71.55, 71.56, 51.51, 51.52,61.53,61.54,71.55, 71.56, 51.51, 51.52,61.53,61.54,71.55, 71.56], 'longitude':[-0.13,-0.13,-0.13,-0.14,-0.13,-0.13, -0.13,-0.13,-0.13,-0.14,-0.13,-0.13, -0.13,-0.13,-0.13,-0.14,-0.13,-0.13], 'id':['A','B','C','D','E','F', 'A','B','C','D','E','F', 'A','B','C','D','E','F'], 'target':[35,410,1,100,114,78, 14,254,101,278,3578,435, 254,254,37,47,38,101], 'time':['2019-03-10 11:00:00','2019-03-10 11:00:00','2019-03-10 11:00:00','2019-03-10 11:00:00','2019-03-10 11:00:00','2019-03-10 11:00:00', '2019-03-10 11:10:00','2019-03-10 11:10:00','2019-03-10 11:10:00','2019-03-10 11:10:00','2019-03-10 11:10:00','2019-03-10 11:10:00', '2019-03-10 11:20:00','2019-03-10 11:20:00','2019-03-10 11:20:00','2019-03-10 11:20:00','2019-03-10 11:20:00','2019-03-10 11:20:00', ]}) # 数据预处理:转换经纬度为弧度 test_data['latitude'] = np.deg2rad(test_data['latitude']) test_data['longitude'] = np.deg2rad(test_data['longitude']) test_data['time'] = pd.to_datetime(test_data['time']) # 初始化结果列表 results = [] # 按ID分组,为每个ID的记录查询其他ID的近邻 unique_ids = test_data['id'].unique() for target_id in unique_ids: # 目标数据:当前ID的所有记录 target_df = test_data[test_data['id'] == target_id] # 候选数据:除当前ID外的所有记录 candidate_df = test_data[test_data['id'] != target_id] # 构建Ball Tree(仅用候选数据) ball_tree = BallTree(candidate_df[['latitude', 'longitude']].values, metric='haversine') # 查询k个近邻(这里k=3,按需调整) distances, indices = ball_tree.query(target_df[['latitude', 'longitude']].values, k=3) # 整理结果 for idx, row in target_df.iterrows(): for i in range(3): neighbor_row = candidate_df.iloc[indices[idx][i]] results.append({ 'source_idx': idx, 'source_id': target_id, 'source_time': row['time'], 'neighbor_idx': neighbor_row.name, 'neighbor_id': neighbor_row['id'], 'neighbor_time': neighbor_row['time'], 'distance': distances[idx][i] * 6371 # 转换为公里(地球半径) }) # 转换为DataFrame neighbor_results = pd.DataFrame(results) print(neighbor_results.head())
2. 查询后基于元信息快速过滤
如果需要更精细的过滤规则,可以在Ball Tree查询后,利用返回的索引关联原始数据的ID信息,通过向量运算快速过滤无效结果。这种方式无需设置过大的k值,仅查询实际需要的近邻数量即可。
示例代码(基于你的原始代码修改):
# --- 步骤3修改:合并结果后直接过滤 # 关联原始数据的source_id和neighbor_id neighbor_info_df = neighbor_info_df.merge( test_data[['id']].rename(columns={'id': 'source_id'}), left_on='id_index_no', right_index=True ) neighbor_info_df = neighbor_info_df.merge( test_data[['id']].rename(columns={'id': 'neighbor_id'}), left_on='neighbor_idx', right_index=True ) # 过滤规则:排除同一ID的记录,同时排除自身(距离为0的情况) neighbor_info_df = neighbor_info_df[ (neighbor_info_df['source_id'] != neighbor_info_df['neighbor_id']) & (neighbor_info_df['distance'] > 1e-8) ] print(neighbor_info_df.head())
为什么这种方式更高效?
- 无需设置过大的k值,减少Ball Tree的计算量
- 过滤操作基于pandas的向量运算,比返回大量无效数据后再筛选快得多,尤其适合大数据集
内容的提问来源于stack exchange,提问作者Rebecca James
相关产品推荐
相关产品推荐

