Pandas按Cohort分组后按多条件高效替换Val值
Pandas分组替换val值解决方案
嘿,我来帮你搞定这个Pandas的分组替换需求,完全不用写繁琐的循环!
先明确你的问题
你有这样一个DataFrame:
import pandas as pd df = pd.DataFrame({ 'date': ['2001-01-01', '2001-02-01', '2001-03-01', '2001-04-01', '2001-02-01', '2001-03-01', '2001-04-01'], 'cohort': ['2001-01-01', '2001-01-01', '2001-01-01', '2001-01-01', '2001-02-01', '2001-02-01', '2001-02-01'], 'val': [100, 101, 102, 101, 200, 201, 201] })
原始输出:
date cohort val 0 2001-01-01 2001-01-01 100 1 2001-02-01 2001-01-01 101 2 2001-03-01 2001-01-01 102 3 2001-04-01 2001-01-01 101 4 2001-02-01 2001-02-01 200 5 2001-03-01 2001-02-01 201 6 2001-04-01 2001-02-01 201
你的需求是:按cohort分组,把每组里date早于该组val最大值对应date的所有行,val都替换成这个最大值,最终得到目标结果:
date cohort val 0 2001-01-01 2001-01-01 102 1 2001-02-01 2001-01-01 102 2 2001-03-01 2001-01-01 102 3 2001-04-01 2001-01-01 101 4 2001-02-01 2001-02-01 201 5 2001-03-01 2001-02-01 201 6 2001-04-01 2001-02-01 201
不用循环的实现步骤
下面是高效的实现方式,完全利用Pandas的内置功能:
1. 先把date转成datetime类型
日期字符串直接比较容易出问题,先转成时间格式:
df['date'] = pd.to_datetime(df['date'])
2. 分组计算每组的关键信息
对每个cohort组,我们需要两个核心数据:val的最大值,以及这个最大值对应的date。用groupby+agg一次性搞定:
group_info = df.groupby('cohort').agg( max_val=('val', 'max'), max_date=('date', lambda x: x[df.loc[x.index, 'val'] == df.loc[x.index, 'val'].max()].iloc[0]) ).reset_index()
这里的lambda函数是为了精准定位到该组中val取最大值时的那个date(如果有多个相同最大值,取第一个出现的,符合你的示例要求)。
3. 把分组信息合并回原数据
用merge把每组的max_val和max_date匹配到原DataFrame的每一行:
df = df.merge(group_info, on='cohort', how='left')
4. 根据条件替换val值
用numpy.where来做条件判断替换,非常高效:
import numpy as np df['val'] = np.where(df['date'] <= df['max_date'], df['max_val'], df['val'])
逻辑很简单:如果当前行的date小于等于该组最大值对应的date,就把val换成max_val,否则保留原来的val。
5. 清理临时列
最后把我们临时加的max_val和max_date删掉,还原清爽的结构:
df = df.drop(columns=['max_val', 'max_date'])
完整代码
把上面的步骤整合起来,完整代码如下:
import pandas as pd import numpy as np df = pd.DataFrame({ 'date': ['2001-01-01', '2001-02-01', '2001-03-01', '2001-04-01', '2001-02-01', '2001-03-01', '2001-04-01'], 'cohort': ['2001-01-01', '2001-01-01', '2001-01-01', '2001-01-01', '2001-02-01', '2001-02-01', '2001-02-01'], 'val': [100, 101, 102, 101, 200, 201, 201] }) # 转换日期格式 df['date'] = pd.to_datetime(df['date']) # 计算每组的最大值和对应日期 group_info = df.groupby('cohort').agg( max_val=('val', 'max'), max_date=('date', lambda x: x[df.loc[x.index, 'val'] == df.loc[x.index, 'val'].max()].iloc[0]) ).reset_index() # 合并信息并替换val df = df.merge(group_info, on='cohort', how='left') df['val'] = np.where(df['date'] <= df['max_date'], df['max_val'], df['val']) # 清理临时列 df = df.drop(columns=['max_val', 'max_date']) print(df)
运行这段代码就能得到你想要的结果啦!
内容的提问来源于stack exchange,提问作者Gaurav Bansal
相关产品推荐
相关产品推荐

