如何在Pandas中分组过滤计算滚动均值并全量填充?
问题描述
测试数据
import pandas as pd import numpy as np df = pd.DataFrame({ 'num_legs': [4,4,5,6,7,4,2,3,4, 2,4,4,5,6,7,4,2,3,3,5,5,6], 'num_wings': [2,7,21,0,21,13,23,43, 2,7,21,13,23,43,23,23,23,11,26,32,75,13], 'new_col': np.arange(22) })
需求说明
- 按
num_legs列分组,对new_col列计算窗口大小为3、最小周期为1的滚动均值; - 计算滚动均值时,仅选取分组内
num_wings > 10的new_col值; - 使用
transform或等价方法将计算结果填充到原DataFrame的所有行中; - 补充要求:对于
num_wings > 10的行,需确保能获取到分组内对应的滚动均值(可通过前/后向填充处理特殊情况)。
尝试的错误代码
df.groupby('num_legs')['new_col'].transform(lambda x: df.loc[df['num_wings']>10, 'new_col'].rolling(3))
解决方案
你的代码核心问题是:未在分组内部筛选符合条件的数据,而是全局筛选;且未完成滚动均值的计算逻辑。以下是两种可行实现方式:
方式1:使用groupby.apply(逻辑更直观)
import pandas as pd import numpy as np # 初始化数据 df = pd.DataFrame({ 'num_legs': [4,4,5,6,7,4,2,3,4, 2,4,4,5,6,7,4,2,3,3,5,5,6], 'num_wings': [2,7,21,0,21,13,23,43, 2,7,21,13,23,43,23,23,23,11,26,32,75,13], 'new_col': np.arange(22) }) def group_rolling_calc(group): # 1. 筛选分组内符合条件的行,提取new_col值 filtered_vals = group[group['num_wings'] > 10]['new_col'] # 2. 计算滚动均值:窗口3,最小周期1(保证不足3个数据时仍能计算) rolling_result = filtered_vals.rolling(window=3, min_periods=1).mean() # 3. 将结果映射回原分组所有行,不符合条件的行填充NaN group['rolling_mean'] = group.index.map(rolling_result) # 4. 处理补充要求:确保符合条件的行必有值(前向填充兜底,无前置值则用后向填充) mask = group['num_wings'] > 10 group.loc[mask, 'rolling_mean'] = group.loc[mask, 'rolling_mean'].fillna(method='ffill').fillna(method='bfill') return group # 按分组应用逻辑,合并结果到原DataFrame df = df.groupby('num_legs').apply(group_rolling_calc)
方式2:使用groupby.transform(贴合需求中的方法要求)
def transform_rolling_calc(x): # 获取当前分组的完整数据(通过索引关联num_wings列) group = df.loc[x.index] # 筛选符合条件的new_col值并计算滚动均值 filtered = group[group['num_wings'] > 10]['new_col'] rolling_result = filtered.rolling(window=3, min_periods=1).mean() # 映射回当前分组的所有行,返回同长度序列 return pd.Series([rolling_result.get(idx, np.nan) for idx in x.index], index=x.index) # 生成滚动均值列 df['rolling_mean'] = df.groupby('num_legs')['new_col'].transform(transform_rolling_calc) # 处理补充要求:确保符合条件的行无缺失值 mask = df['num_wings'] > 10 df.loc[mask, 'rolling_mean'] = df.loc[mask, 'rolling_mean'].fillna(method='ffill').fillna(method='bfill')
内容的提问来源于stack exchange,提问作者tjt
相关产品推荐
相关产品推荐

