基于分组聚合且排除当前记录的Python中位数计算问题
分组后排除当前记录计算value1的中位数
需求:在Pandas DataFrame中新增一列,值为按key字段分组后,排除当前记录的value1字段的中位数。
示例数据
import pandas as pd df = pd.DataFrame({ 'key': ['A', 'A', 'A', 'A', 'A', 'A', 'A', 'B', 'B', 'B', 'B', 'B', 'B', 'B', 'B', 'B', 'C', 'C', 'C', 'C', 'C', 'C', 'C', 'C', 'C'], 'value1': [0.1, 0.244, 0.373, 0.514, 0.663, 0.786, 0.902, 1.01, 1.151, 1.295, 1.434, 1.541, 1.679, 1.793, 1.94, 2.049, 2.164, 2.284, 2.432, 2.533, 2.68, 2.786, 2.906, 3.008, 3.136], 'value2': ['Dept1', 'Dept2', 'Dept3', 'Dept4', 'Dept5', 'Dept6', 'Dept7', 'Dept8', 'Dept9', 'Dept10', 'Dept11', 'Dept12', 'Dept13', 'Dept14', 'Dept15', 'Dept16', 'Dept17', 'Dept18', 'Dept19', 'Dept20', 'Dept21', 'Dept22', 'Dept23', 'Dept24', 'Dept25'], 'desired_output':['0.589', '0.589', '0.589', '0.518', '0.444', '0.444', '0.444', '1.61', '1.61', '1.61', '1.61', '1.557', '1.488', '1.488', '1.488', '1.488', '2.733', '2.733', '2.733', '2.733', '2.66', '2.607', '2.607', '2.607', '2.607'] })
错误代码分析
你尝试的代码触发Transform function invalid for data types错误:
df['desired_output2'] = df.groupby(['key'], as_index = False)['value1'].transform(lambda x: [x[(v['value2'] != x['value2'])].median() for v in x])
问题出在:
transform处理的x是分组后的value1单列序列,无法直接访问value2字段;- 遍历
x得到的是单个数值,不是行对象,导致索引匹配逻辑完全错误。
解决方案
方法1:直观遍历计算(适合小数据集)
通过groupby.apply处理每个分组,对每行排除自身后计算中位数:
def get_excluded_median(group): medians = [] for idx in group.index: # 排除当前行的value1 filtered_vals = group['value1'].drop(idx) medians.append(filtered_vals.median()) group['desired_output2'] = medians return group # 应用函数并重置索引 df = df.groupby('key').apply(get_excluded_median).reset_index(drop=True) # 保留3位小数,和示例输出格式对齐 df['desired_output2'] = df['desired_output2'].round(3).astype(str)
方法2:排序优化法(适合大数据集)
利用中位数依赖排序的特性,预先对分组内value1排序,直接通过位置计算排除后的中位数,避免重复过滤操作:
import numpy as np def efficient_excluded_median(group): sorted_vals = group['value1'].sort_values().values n = len(sorted_vals) # 记录每个value1在排序数组中的位置 group['sort_pos'] = group['value1'].apply(lambda x: np.argmax(sorted_vals == x)) medians = [] for pos in group['sort_pos']: # 拼接排除当前位置后的数组 excluded = np.concatenate([sorted_vals[:pos], sorted_vals[pos+1:]]) m = len(excluded) # 计算中位数 if m % 2 == 1: median = excluded[m//2] else: median = (excluded[m//2 - 1] + excluded[m//2]) / 2 medians.append(median) group['desired_output2'] = medians return group.drop('sort_pos', axis=1) df = df.groupby('key').apply(efficient_excluded_median).reset_index(drop=True) df['desired_output2'] = df['desired_output2'].round(3).astype(str)
两种方法均能生成与desired_output一致的结果,方法2在数据量较大时性能更优。
内容的提问来源于stack exchange,提问作者wolfy
相关产品推荐
相关产品推荐

