Python中如何拆分历史数据点≤12的商品销售时间序列数据集用于6个月预测
提取历史数据点数量≤12的商品数据集方法
现有一份商品销量时间序列数据集,包含Item ID、Date(时间范围2015-01-01至2022-12-01)、Quantities of items sold三列,需要拆分出历史数据点数量≤12的Item ID对应的完整数据集。
你的尝试代码问题分析
你写的代码:
grouped_data = df.groupby('item_id').apply(lambda x: x[x['date'].count() <= 12])
存在逻辑冗余且效率偏低的问题:x['date'].count()返回的是当前分组的总行数,用这个值去逐行判断,结果要么保留整个分组(如果行数≤12),要么完全剔除该分组,写法绕了弯路,不如直接针对分组行数做筛选。
两种高效实现方法
方法1:先筛选合法Item ID,再提取数据
这种方式逻辑清晰,适合需要单独保留合法ID列表的场景:
# 统计每个Item ID的数据点数量,筛选出数量≤12的ID valid_item_ids = df.groupby('Item ID').size()[lambda s: s <= 12].index # 提取对应ID的所有数据 filtered_df = df[df['Item ID'].isin(valid_item_ids)]
groupby('Item ID').size()会直接统计每个分组的行数,比count()更高效;后续通过isin()过滤原数据集,快速得到目标结果。
方法2:使用groupby的filter方法(更简洁)
filter方法会直接保留满足条件的整个分组,代码更紧凑:
filtered_df = df.groupby('Item ID').filter(lambda x: len(x) <= 12)
len(x)直接获取当前分组的行数,判断是否≤12,符合条件的分组会被完整保留到结果中。
内容的提问来源于stack exchange,提问作者Mufasa
相关产品推荐
相关产品推荐

