Pandas分组内计算滚动求和遇索引错误,求正确实现方法
分组内滚动统计量的正确实现方式
错误原因
你遇到的TypeError是因为groupby.rolling()返回的结果带有多级索引(外层是分组键class,内层是原DataFrame的索引),直接赋值给原DataFrame列时,索引无法对齐。
解决方案
以下两种方法均可解决索引对齐问题,且能得到你需要的结果:
方法1:使用transform自动对齐(推荐)
transform会自动将分组计算的结果按原DataFrame的索引重新排列,无需手动处理索引:
import pandas as pd df = pd.DataFrame.from_dict({'class': ['a', 'b', 'b', 'c', 'c', 'c', 'b', 'a', 'b'], 'val': [1, 2, 3, 4, 5, 6, 7, 8, 9]}) # 窗口为2的滚动求和,min_periods=1确保第一个元素保留原值 df['sum2_per_class'] = df.groupby('class')['val'].transform( lambda x: x.rolling(2, min_periods=1).sum() )
方法2:手动重置索引对齐
通过reset_index去掉分组键的索引层,保留原DataFrame的索引后再赋值:
df['sum2_per_class'] = df.groupby('class')['val']\ .rolling(2, min_periods=1).sum()\ .reset_index(level=0, drop=True)
关键参数说明
min_periods=1:默认窗口为2时,第一个元素因缺少前一个值会返回NaN,设置该参数后,允许窗口中至少存在1个元素就计算统计量,从而保留第一个元素的原值。
最终结果
执行上述代码后,df的sum2_per_class列会和你给出的目标列完全一致:
class val sum2_per_class 0 a 1 1.0 1 b 2 2.0 2 b 3 5.0 3 c 4 4.0 4 c 5 9.0 5 c 6 11.0 6 b 7 10.0 7 a 8 9.0 8 b 9 16.0
内容的提问来源于stack exchange,提问作者Baron Yugovich
相关产品推荐
相关产品推荐

