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

如何设置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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 17:45:47