跳到正文

Pandas分组聚合:groupby 与 agg 实战

描述性统计里最常写的一句话是「按某变量分组,算各组的均值和标准差」。pandas 用 groupby 做这件事,背后是分裂—应用—合并(split-apply-combine)三步:按分组键把表切成若干块,对每块单独算,再把结果拼回一张表。理解这三步,后面 agg、transform、filter、pivot_table 的差别就好懂了——它们都是「应用」这一步的变体。

import pandas as pd
from sklearn.datasets import load_iris
df = load_iris(as_frame=True).frame
df.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__)
species
setosa 1.462
versicolor 4.260
virginica 5.552
Name: petal_len, dtype: float64
Series

返回的是 Series 不是 DataFrame,索引换成了分组键。这是 groupby 结果的第一特征:分组键变成索引。要接着画图或者写文件,通常得把它变回普通列,用 reset_index() 或者一开始就写 as_index=False

print(df.groupby("species", as_index=False)["petal_len"].mean())
species petal_len
0 setosa 1.462
1 versicolor 4.260
2 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())
group
A 3
B 3
Name: value, dtype: int64
group
A 3
B 3
NaN 4
Name: value, dtype: int64

四个样本只汇总出两个组,value=4 那行没进任何统计,也没有任何提示。加 dropna=False 才会把它单独列成一个 NaN 组。分组键是清洗后留下的缺失(比如物种标签丢了)时,这会直接少算样本量,做之前先数一下分组键的缺失个数——方法见 Pandas数据清洗

agg 接受函数名列表,一次算出多列结果:

print(df.groupby("species")["petal_len"].agg(["count", "mean", "std"]))
count mean std
species
setosa 50 1.462 0.173664
versicolor 50 4.260 0.469911
virginica 50 5.552 0.551895

列名就是函数名。count 数的是非缺失值个数,和 size 不同——size 数行数,缺失值也算。数据有缺失时这两个数会不一样,报「每组样本量」时用错会很难看。

要对不同列用不同函数,传字典(dict):

print(df.groupby("species").agg({"petal_len": "mean", "sepal_len": "max"}))
petal_len sepal_len
species
setosa 1.462 5.8
versicolor 4.260 7.0
virginica 5.552 7.9

字典写法的问题是结果列名就是原列名:petal_len 这一列到底是均值还是最大值,单看表头看不出来。表要直接贴进论文时,列名得写清楚,用命名聚合。

print(df.groupby("species").agg(
n=("petal_len", "size"),
petal_mean=("petal_len", "mean"),
sepal_sd=("sepal_len", "std"),
))
n petal_mean sepal_sd
species
setosa 50 1.462 0.352490
versicolor 50 4.260 0.516171
virginica 50 5.552 0.635880

写法是 新列名=("原列名", "函数名"),官方叫命名聚合(named aggregation)。做正式分析一律用这种写法——输出的表直接能贴进报告,不需要再改列名;字典那种写法省下的几个字符,后面要在改列名上还回去。列多的时候把多个元组并排写,注意每组都要带上圆括号。

按两个变量分组时,分组键都进索引,形成多层索引(MultiIndex):

df["long_petal"] = df["petal_len"] > 4.0
print(df.groupby(["species", "long_petal"])["sepal_len"].mean())
print("索引名:", df.groupby(["species", "long_petal"])["sepal_len"].mean().index.names)
species long_petal
setosa False 5.006000
versicolor False 5.487500
True 6.147059
virginica True 6.588000
Name: 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_len
0 setosa False 5.006000
1 versicolor False 5.487500
2 versicolor True 6.147059
3 virginica True 6.588000

行数没变,列数多了一个——原来的索引层级变成了普通列。这是拆多层索引的标准动作,之后就能按列名筛选、排序、写文件。

要的是「每个样本减去它所在组的均值」这种组内中心化(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_dev
species
setosa 1.462 2.220446e-17
versicolor 4.260 2.131628e-16
virginica 5.552 -5.506706e-16
150

transform 返回的 Series 长度和原表一致,可以直接赋值成新列。上面每组偏离均值的均值是 2.22e-17 而不是 0——浮点数加减的舍入误差,量级在 1e-16,实际就是 0。判断这类结果不要用 == 0,用 abs(x) < 1e-10 之类的阈值。

agg 返回「每组一行」,transform 返回「每行一个值」,这是两者的根本区别。有人试图用 df.groupby(...)["x"].mean() 赋回原表,结果长度对不上,赋值变成一堆缺失值——要往原表里加分组统计量,就用 transform

filter 筛的是,不是行。传进去的函数作用于每个子表,返回 True 的组整组保留:

big = df.groupby("species").filter(lambda x: x["sepal_len"].mean() > 6.0)
print(big["species"].value_counts())
species
virginica 50
Name: count, dtype: int64

三个组的花萼平均长度是 5.006、5.936、6.588,只有 virginica 超过 6.0,所以只留下它。写 df[df["sepal_len"] > 6.0] 是筛行,写法短但结果完全不同:后者会把 setosa 里少数花萼偏长的样本也留下,破坏分组完整性。做组间比较时,样本量不一致会直接影响检验的自由度。

透视表不是另一套东西,它做的就是「分组 + 聚合 + 把一组键变成列」。同一份数据,两种写法:

print(df.pivot_table(index="species", values="petal_len", aggfunc="mean"))
print(df.groupby("species")["petal_len"].mean())
petal_len
species
setosa 1.462
versicolor 4.260
virginica 5.552
species
setosa 1.462
versicolor 4.260
virginica 5.552
Name: 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 True
species
setosa 1.4620 NaN
versicolor 3.7125 4.517647
virginica 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() 对应 transformfilter() 在两边的语义也一致——都有一批动词按组的粒度工作。