如何在Pandas中按组计算行间欧氏距离并生成关联索引列表列
如何在Pandas中按组计算行间欧氏距离并生成关联索引列表列
嘿,看来你之前已经搞定过类似的Pandas行处理问题了,那咱们这次就直奔主题,解决你当前的需求~
首先我再确认下核心需求,避免理解偏差:咱们需要按trial、RECORDING_SESSION_LABEL、IP_INDEX这三个字段分组,组内从第二行开始,每一行都要和它上面所有的行,用CURRENT_FIX_X和CURRENT_FIX_Y计算欧氏距离。如果距离小于58.93,就把被比较行的CURRENT_FIX_INDEX收集起来,转成字符串后放到当前行的新列refix_list里对吧?
那我给你两种实现方式,一种是直观的逐行遍历版,适合理解调试;另一种是更高效的向量化版,适合处理大数据集。
一、直观逐行遍历版(适合小数据集或调试)
这个版本逻辑完全跟着需求步骤走,新手也能一眼看明白:
import pandas as pd import numpy as np def process_single_group(group): # 先给当前组初始化新列,默认是空字符串 group['refix_list'] = '' # 从组内第二行开始遍历(索引从1开始,因为第0行没有上面的行) for row_idx in range(1, len(group)): # 取出当前行的坐标值 curr_x, curr_y = group.iloc[row_idx][['CURRENT_FIX_X', 'CURRENT_FIX_Y']] # 取出当前行之前的所有行 previous_rows = group.iloc[:row_idx] # 计算当前行和所有前序行的欧氏距离 distances = np.sqrt( (previous_rows['CURRENT_FIX_X'] - curr_x)**2 + (previous_rows['CURRENT_FIX_Y'] - curr_y)**2 ) # 筛选距离小于58.93的行,提取它们的CURRENT_FIX_INDEX matched_indices = previous_rows[distances < 58.93]['CURRENT_FIX_INDEX'].tolist() # 把列表转成字符串,赋值给当前行的refix_list列 group.iloc[row_idx, group.columns.get_loc('refix_list')] = ','.join(map(str, matched_indices)) return group # 把处理函数应用到每个分组 df = df.groupby(['trial', 'RECORDING_SESSION_LABEL', 'IP_INDEX']).apply(process_single_group)
这个版本的关键点:
- 每个分组单独处理,绝对不会跨组比较,完全符合你的分组要求
- 第一行的
refix_list是空字符串,因为它没有前序行可以比较,完全合理 - 如果某一行和所有前序行的距离都不达标,对应的
refix_list也是空字符串
二、高效向量化版(适合大数据集)
如果你的数据集行数很多,逐行遍历会有点慢,那可以用scipy的cdist函数来批量计算距离,速度会快很多:
import pandas as pd from scipy.spatial.distance import cdist def process_group_fast(group): # 重置组内索引,避免原索引混乱,方便后续定位 group = group.reset_index(drop=True) group['refix_list'] = '' # 把坐标列转成numpy数组,方便批量计算 coords_matrix = group[['CURRENT_FIX_X', 'CURRENT_FIX_Y']].values # 同样从第二行开始遍历 for row_idx in range(1, len(group)): # 批量计算当前行和所有前序行的欧氏距离 all_distances = cdist([coords_matrix[row_idx]], coords_matrix[:row_idx], metric='euclidean')[0] # 筛选符合条件的索引 matched_indices = group.loc[all_distances < 58.93, 'CURRENT_FIX_INDEX'].tolist() # 转字符串赋值 group.loc[row_idx, 'refix_list'] = ','.join(map(str, matched_indices)) return group # 应用到分组 df = df.groupby(['trial', 'RECORDING_SESSION_LABEL', 'IP_INDEX']).apply(process_group_fast)
这个版本的优势:
cdist是底层优化过的批量计算函数,比手动循环计算距离快很多,数据量越大优势越明显- 逻辑和逐行版完全一致,只是把距离计算的部分换成了高效的批量实现
一些额外的小提示
- 如果你需要调整索引列表的分隔符,比如用分号或者空格,只需要修改
','.join里的分隔符就行 - 如果
CURRENT_FIX_INDEX是数值类型,map(str, matched_indices)会把它转成字符串,避免出现类型错误 - 要是你想给空列表的情况设置默认值(比如填
'无'),可以加个判断:','.join(matched_indices) if matched_indices else '无'
备注:内容来源于stack exchange,提问作者Eslifkin
相关产品推荐
相关产品推荐

