如何按时间索引拆分pandas多级索引DataFrame为训练/测试集
报错原因
之前使用的df[df.loc['dates'] > '2020-12-31']写法存在两个问题:
- 多级索引场景下,
df.loc['dates']会尝试匹配第一级索引值等于字符串dates的行,而非读取名为date的索引层级 - 索引名拼写错误,定义的索引名是
date不是dates
正确实现方案
你的数据集第一级索引为时间类型,直接提取对应索引层级做布尔筛选即可完成拆分,完整可运行代码如下:
import pandas as pd from scipy import stats # 修正原示例代码的导入错误,stats模块属于scipy库 # 构造示例数据集 data = stats.poisson(mu=[5,2,1,7,2]).rvs([60, 5]).T.ravel() dates = pd.date_range('2017-01-01', freq='M', periods=60) locations = [f'location_{i}' for i in range(5)] df = pd.DataFrame(data, index=pd.MultiIndex.from_product([dates, locations]), columns=['eaches']) df.index.names = ['date', 'location'] # 定义拆分时间节点 cutoff_time = pd.Timestamp('2021-01-01') # 拆分数据集:2021年1月之前为训练集,之后(含2021年1月)为测试集 df_train = df[df.index.get_level_values('date') < cutoff_time] df_test = df[df.index.get_level_values('date') >= cutoff_time]
其他等价写法
如果偏好loc的索引切片写法,可以借助pd.IndexSlice实现,逻辑和上面完全一致:
idx = pd.IndexSlice df_train = df.loc[idx[:'2020-12-31', :], :] df_test = df.loc[idx['2021-01-01':, :], :]
结果验证
运行以下代码可以确认拆分结果符合要求:
print(f"训练集日期跨度:{df_train.index.get_level_values('date').min().date()} 至 {df_train.index.get_level_values('date').max().date()}") print(f"测试集日期跨度:{df_test.index.get_level_values('date').min().date()} 至 {df_test.index.get_level_values('date').max().date()}")
输出结果会显示训练集最晚日期为2020-12-31,测试集最早日期为2021-01-31,完全匹配拆分需求。
内容的提问来源于stack exchange,提问作者Jordan
相关产品推荐
相关产品推荐

