如何高效统计DataFrame中各user_id对应check_date的符合条件行数?
更高效的区间计数方案
你的需求很清晰——要给每一行计算对应check_date下,同user_id的区间(start_date到end_date)包含该日期的行数。原方案虽然能得到正确结果,但循环处理每个唯一check_date并重复过滤全表的方式,在数据量较大时效率会很低。这里提供一个基于二分查找的优化方案,能大幅减少重复计算,提升运行效率:
核心思路
对于每个user_id,我们可以把它的所有start_date和end_date分别排序。对于任意一个check_date,满足start_date <= check_date < end_date的区间数量,等价于**start_date <= check_date的区间数减去end_date <= check_date的区间数**。利用二分查找可以快速统计这两个数值,避免了每次都遍历全表过滤。
实现代码
import pandas as pd import bisect # 示例数据 data = [ [1, 1, '2018-11-05', '2018-12-06', '2018-11-22', 2], [2, 1, '2018-11-10', '2018-11-25', '2018-11-24', 2], [3, 1, '2018-12-05', '2018-12-31', '2018-12-20', 1], [4, 1, '2018-12-25', '2019-01-30', '2018-12-30', 2] ] df = pd.DataFrame(data, columns=['id', 'user_id', 'start_date', 'end_date', 'check_date', 'actual']) # 转换日期格式为datetime df['start_date'] = pd.to_datetime(df['start_date']) df['end_date'] = pd.to_datetime(df['end_date']) df['check_date'] = pd.to_datetime(df['check_date']) # 1. 按user_id分组,收集并排序每个用户的start和end日期 user_intervals = df.groupby('user_id').agg( starts=('start_date', lambda x: sorted(x)), ends=('end_date', lambda x: sorted(x)) ).reset_index() # 2. 定义计数函数:用二分查找计算包含check_date的区间数 def count_matching_intervals(starts, ends, check_date): # 统计start_date <= check_date的数量 cnt_start = bisect.bisect_right(starts, check_date) # 统计end_date <= check_date的数量 cnt_end = bisect.bisect_right(ends, check_date) # 符合条件的区间数 = 前者减后者 return cnt_start - cnt_end # 3. 创建所有(user_id, check_date)的组合,计算每个组合的实际行数 unique_users = df['user_id'].unique() unique_check_dates = df['check_date'].unique() # 生成笛卡尔积 user_check_pairs = pd.MultiIndex.from_product( [unique_users, unique_check_dates], names=['user_id', 'check_date'] ).to_frame(index=False) # 合并用户区间数据 user_check_pairs = user_check_pairs.merge(user_intervals, on='user_id', how='left') # 计算每个组合的actual_rows user_check_pairs['actual_rows'] = user_check_pairs.apply( lambda row: count_matching_intervals(row['starts'], row['ends'], row['check_date']), axis=1 ) # 4. 合并结果回原DataFrame df = df.merge( user_check_pairs[['user_id', 'check_date', 'actual_rows']], on=['user_id', 'check_date'], how='left' ) print(df[['id', 'user_id', 'check_date', 'actual', 'actual_rows']])
为什么更高效?
- 原方案的时间复杂度是O(K*N)(K是唯一
check_date数量,N是总行数),每次循环都要过滤全表; - 优化方案的时间复杂度是O(N log N + U*K log N)(U是唯一
user_id数量),排序是O(N log N),每个(user_id, check_date)组合的计算是O(log N),整体计算量远小于原方案,尤其当数据量较大时优势明显。
验证结果
运行上述代码后,输出的actual_rows会和示例中的actual列完全一致,说明结果正确。
内容的提问来源于stack exchange,提问作者Tamplier
相关产品推荐
相关产品推荐

