Pandas groupby aggregate无法识别id列问题排查
我正在处理一个大型数据库,已使用Pandas的apply方法根据客户消费的产品类型对客户进行分类。示例代码如下:
import pandas as pd import numpy as np from datetime import datetime num_variables = 1000 rng = np.random.default_rng() data = pd.DataFrame({ "id" : np.random.randint(1,999999999,num_variables), "date" : [np.random.choice(pd.date_range(datetime(2021,1,1),datetime(2022,12,31))) for i in range(num_variables)], "product" : [np.random.choice(['giftcards', 'afiliates']) for i in range(num_variables)], "brand" : [np.random.choice(['brand_1', 'brand_2', 'brand_4', 'brand_6']) for i in range(num_variables)], "gmv" : rng.random(num_variables) * 100, "revenue" : rng.random(num_variables) * 100,}) data = data.astype({'product':'category', 'brand':'category'}) base = data.groupby(['id', 'product']).aggregate({'product' : 'count'}) base = base.unstack()
接下来我需要按"type"列对客户分组,统计每组的客户数量,先执行分类函数及应用:
def setup(row): if row[('product', 'afiliates')] >= 1 and row[('product', 'giftcards')] == 0: return 'afiliates' if row[('product', 'afiliates')] == 0 and row[('product', 'giftcards')] >= 1: return 'gift' if row[('product', 'afiliates')] >= 1 and row[('product', 'giftcards')] >= 1: return 'both' base['type'] = base.apply(setup, axis=1) base.reset_index(inplace=True)
目前一切正常,执行results = base[['type','id']].groupby(['type'], dropna=False).agg('count')可得到正常结果,但改用results = base[['type','id']].groupby(['type']).aggregate({'id': 'count'})时,触发KeyError,提示"Column(s) ['id'] do not exist",请问我忽略了什么?
原因分析
执行base = base.unstack()后,DataFrame的列变为多级索引(MultiIndex)。后续通过reset_index()将id转为列、新增type列时,Pandas会自动将这些普通列名整合进MultiIndex体系,最终id和type的列名实际是元组形式(比如('id', '')、('type', '')),而非单纯的字符串'id'/'type'。
当你用aggregate({'id': 'count'})时,传入的字符串'id'无法匹配实际的元组列名,因此触发KeyError;而agg('count')是对所有列统一应用统计逻辑,不需要指定列名,所以能正常运行。
解决方法
方法1:使用正确的元组列名指定
直接用实际的元组列名替换字符串'id':
results = base[['type','id']].groupby(['type']).aggregate({('id', ''): 'count'})
方法2:扁平化多级列名
在reset_index()后,将所有列名转为普通字符串,避免元组列名的困扰:
# 扁平化多级列名,例如将('product', 'afiliates')转为'product_afiliates' base.columns = ['_'.join(col).strip('_') for col in base.columns.values] # 之后正常执行分组统计 results = base[['type','id']].groupby(['type']).aggregate({'id': 'count'})
方法3:更高效的统计方式
无需指定列名,直接使用size或value_counts统计客户数,代码更简洁高效:
# 方式1:用size统计每组客户数 results = base.groupby('type', dropna=False)['id'].size().reset_index(name='count') # 方式2:用value_counts直接统计 results = base['type'].value_counts(dropna=False).reset_index(name='count').rename(columns={'index':'type'})
内容的提问来源于stack exchange,提问作者FábioRB

