Pandas分组聚合:groupby 与 agg 实战
描述性统计里最常写的一句话是「按某变量分组,算各组的均值和标准差」。pandas 用 groupby 做这件事,背后是分裂—应用—合并(split-apply-combine)三步:按分组键把表切成若干块,对每块单独算,再把结果拼回一张表。理解这三步,后面 agg、transform、filter、pivot_table 的差别就好懂了——它们都是「应用」这一步的变体。
import pandas as pdfrom sklearn.datasets import load_iris
df = load_iris(as_frame=True).framedf.columns = ["sepal_len", "sepal_wid", "petal_len", "petal_wid", "species"]df["species"] = df["species"].map({0: "setosa", 1: "versicolor", 2: "virginica"})
print(df.groupby("species")["petal_len"].mean())print(type(df.groupby("species")["petal_len"].mean()).__name__)speciessetosa 1.462versicolor 4.260virginica 5.552Name: petal_len, dtype: float64Series返回的是 Series 不是 DataFrame,索引换成了分组键。这是 groupby 结果的第一特征:分组键变成索引。要接着画图或者写文件,通常得把它变回普通列,用 reset_index() 或者一开始就写 as_index=False:
print(df.groupby("species", as_index=False)["petal_len"].mean()) species petal_len0 setosa 1.4621 versicolor 4.2602 virginica 5.552习惯 ggplot2 或 dplyr 的人容易在这里卡住:R 里 summarise() 出来的分组列还在,pandas 默认把它挪到了索引上。两种写法都记住,看代码时才不会以为别人漏了一步。
分组键有缺失值时还有个静默丢失:默认把缺失组整个丢掉。
g = pd.DataFrame({"group": ["A", "A", "B", None], "value": [1, 2, 3, 4]})
print(g.groupby("group")["value"].sum())print(g.groupby("group", dropna=False)["value"].sum())groupA 3B 3Name: value, dtype: int64groupA 3B 3NaN 4Name: value, dtype: int64四个样本只汇总出两个组,value=4 那行没进任何统计,也没有任何提示。加 dropna=False 才会把它单独列成一个 NaN 组。分组键是清洗后留下的缺失(比如物种标签丢了)时,这会直接少算样本量,做之前先数一下分组键的缺失个数——方法见 Pandas数据清洗。
agg:一次算多个统计量
Section titled “agg:一次算多个统计量”agg 接受函数名列表,一次算出多列结果:
print(df.groupby("species")["petal_len"].agg(["count", "mean", "std"])) count mean stdspeciessetosa 50 1.462 0.173664versicolor 50 4.260 0.469911virginica 50 5.552 0.551895列名就是函数名。count 数的是非缺失值个数,和 size 不同——size 数行数,缺失值也算。数据有缺失时这两个数会不一样,报「每组样本量」时用错会很难看。
要对不同列用不同函数,传字典(dict):
print(df.groupby("species").agg({"petal_len": "mean", "sepal_len": "max"})) petal_len sepal_lenspeciessetosa 1.462 5.8versicolor 4.260 7.0virginica 5.552 7.9字典写法的问题是结果列名就是原列名:petal_len 这一列到底是均值还是最大值,单看表头看不出来。表要直接贴进论文时,列名得写清楚,用命名聚合。
命名聚合:结果列名由你决定
Section titled “命名聚合:结果列名由你决定”print(df.groupby("species").agg( n=("petal_len", "size"), petal_mean=("petal_len", "mean"), sepal_sd=("sepal_len", "std"),)) n petal_mean sepal_sdspeciessetosa 50 1.462 0.352490versicolor 50 4.260 0.516171virginica 50 5.552 0.635880写法是 新列名=("原列名", "函数名"),官方叫命名聚合(named aggregation)。做正式分析一律用这种写法——输出的表直接能贴进报告,不需要再改列名;字典那种写法省下的几个字符,后面要在改列名上还回去。列多的时候把多个元组并排写,注意每组都要带上圆括号。
多层索引:多列分组的结果
Section titled “多层索引:多列分组的结果”按两个变量分组时,分组键都进索引,形成多层索引(MultiIndex):
df["long_petal"] = df["petal_len"] > 4.0print(df.groupby(["species", "long_petal"])["sepal_len"].mean())print("索引名:", df.groupby(["species", "long_petal"])["sepal_len"].mean().index.names)species long_petalsetosa False 5.006000versicolor False 5.487500 True 6.147059virginica True 6.588000Name: sepal_len, dtype: float64索引名: ['species', 'long_petal']setosa 组全是短花瓣,所以只有 False 一行;versicolor 下面有 False 和 True 两层。读的时候注意 versicolor 的第二行外层是空白的,表示它和上一行同属 versicolor——这是 pandas 显示多层索引时省略了重复的外层,不是数据缺失。
取值要传元组,写成 .loc[("versicolor", True)] 这样。层级多了很啰嗦,通常直接 reset_index() 拍平:
print(df.groupby(["species", "long_petal"])["sepal_len"].mean().reset_index().head(4)) species long_petal sepal_len0 setosa False 5.0060001 versicolor False 5.4875002 versicolor True 6.1470593 virginica True 6.588000行数没变,列数多了一个——原来的索引层级变成了普通列。这是拆多层索引的标准动作,之后就能按列名筛选、排序、写文件。
transform:保持行数不变
Section titled “transform:保持行数不变”要的是「每个样本减去它所在组的均值」这种组内中心化(centering),结果必须和原表一样长,agg 做不到,用 transform:
df["petal_center"] = df.groupby("species")["petal_len"].transform("mean")df["petal_dev"] = df["petal_len"] - df["petal_center"]
print(df.groupby("species")[["petal_len", "petal_dev"]].mean())print(len(df)) petal_len petal_devspeciessetosa 1.462 2.220446e-17versicolor 4.260 2.131628e-16virginica 5.552 -5.506706e-16150transform 返回的 Series 长度和原表一致,可以直接赋值成新列。上面每组偏离均值的均值是 2.22e-17 而不是 0——浮点数加减的舍入误差,量级在 1e-16,实际就是 0。判断这类结果不要用 == 0,用 abs(x) < 1e-10 之类的阈值。
agg 返回「每组一行」,transform 返回「每行一个值」,这是两者的根本区别。有人试图用 df.groupby(...)["x"].mean() 赋回原表,结果长度对不上,赋值变成一堆缺失值——要往原表里加分组统计量,就用 transform。
filter:按组的条件筛行
Section titled “filter:按组的条件筛行”filter 筛的是组,不是行。传进去的函数作用于每个子表,返回 True 的组整组保留:
big = df.groupby("species").filter(lambda x: x["sepal_len"].mean() > 6.0)print(big["species"].value_counts())speciesvirginica 50Name: count, dtype: int64三个组的花萼平均长度是 5.006、5.936、6.588,只有 virginica 超过 6.0,所以只留下它。写 df[df["sepal_len"] > 6.0] 是筛行,写法短但结果完全不同:后者会把 setosa 里少数花萼偏长的样本也留下,破坏分组完整性。做组间比较时,样本量不一致会直接影响检验的自由度。
pivot_table 与 groupby 的关系
Section titled “pivot_table 与 groupby 的关系”透视表不是另一套东西,它做的就是「分组 + 聚合 + 把一组键变成列」。同一份数据,两种写法:
print(df.pivot_table(index="species", values="petal_len", aggfunc="mean"))print(df.groupby("species")["petal_len"].mean()) petal_lenspeciessetosa 1.462versicolor 4.260virginica 5.552speciessetosa 1.462versicolor 4.260virginica 5.552Name: petal_len, dtype: float64数值完全一致,差别在结构:pivot_table 返回 DataFrame,groupby 返回 Series。columns 参数才是透视表的用处所在——它把某个分类变量的取值铺成列:
print(df.pivot_table(index="species", columns="long_petal", values="petal_len", aggfunc="mean"))long_petal False Truespeciessetosa 1.4620 NaNversicolor 3.7125 4.517647virginica NaN 5.552000出现 NaN 的空位说明这一组没有样本(setosa 没有长花瓣,virginica 没有短花瓣),不是数据丢了。pivot_table 默认对缺失组合填 NaN,可以传 fill_value=0 换掉——但要想清楚,把「没有观测」写成 0 会让下游的均值、比值全部偏移,多数情况下留着 NaN 更诚实。
aggfunc 不传时默认是 "mean",这也是个常见的惊吓点:有人只想把长表摆成宽表,结果数值被默默平均掉了。只想改形状、不做聚合,用 pivot(见 Pandas数据清洗),但它要求每个(行, 列)组合唯一,有重复就会报错。
分组聚合之后再接合并、清洗,链条就完整了。R 语言的 dplyr分组聚合 是同一套分裂—应用—合并思想:group_by() + summarise() 对应 groupby() + agg(),mutate() 对应 transform,filter() 在两边的语义也一致——都有一批动词按组的粒度工作。