如何简化Pandas DataFrame中train/valid/test序列划分的代码?
简洁实现方案
不用嵌套循环,直接用Pandas的分组和序列计数功能就能高效搞定,核心思路是按用户分组后,给每条记录标记逆序位置,再根据位置映射对应的数据集标签:
方法一(清晰易读)
import pandas as pd import numpy as np # 假设你的DataFrame是df,用户标识列名为user_id # 按用户分组,计算每条记录在组内的逆序位置(最后一条为1,倒数第二条为2,以此类推) df['_temp_rank'] = df.groupby('user_id')['cnt_seq'].cumcount(ascending=False) + 1 # 根据逆序位置批量赋值标签 df['split'] = np.select( [df['_temp_rank'] == 1, df['_temp_rank'] == 2], ['test', 'valid'], default='train' ) # 移除临时列 df.drop('_temp_rank', axis=1, inplace=True)
方法二(紧凑写法)
如果不想用临时列,可以直接在transform里完成逻辑:
df['split'] = df.groupby('user_id')['cnt_seq'].transform( lambda group: np.select( [group.rank(ascending=False, method='first') == 1, group.rank(ascending=False, method='first') == 2], ['test', 'valid'], default='train' ) )
关键逻辑说明
cumcount(ascending=False):在每个用户组内从最后一条记录开始计数(从0起始),加1后让最后一条对应1,倒数第二条对应2np.select:批量处理多条件赋值,比循环效率高很多,数据量越大优势越明显rank(ascending=False, method='first'):给组内记录做逆序排名,method='first'确保即使cnt_seq有重复值,每条记录的排名也唯一,避免标签分配出错
内容的提问来源于stack exchange,提问作者Dang
相关产品推荐
相关产品推荐

