基于PySpark RDD:如何计算并对比用户使用频率平均值?
需求说明
某公司计划为k位用户提供两个月的应用免费使用权(以优惠券形式),目标是识别可能流失的用户,并从中筛选出k位值得留存的高价值用户。这类需留存用户的定义是:曾因高频听歌为公司带来较高价值,但近期使用频率出现显著下降。
要求仅使用PySpark RDD(禁止使用DataFrame),完成以下计算:
- 每位用户最近6个月的使用频率平均值
- 每位用户最近3个月的使用频率平均值
- 若最近3个月的平均值低于6个月平均值的50%,标记该用户为流失风险用户
- 最终从流失风险用户中选出k位历史价值最高的用户(按6个月总使用量排序)
原实现代码
from datetime import datetime from pyspark import SparkContext def k_users_to_retain(k, timestamp): # 步骤1:筛选注册时间早于指定时间的用户并收集其ID filtered_users = users_rdd.filter(lambda user: user.registered is not None and user.registered < timestamp) eligible_user_ids = filtered_users.map(lambda user: user.userid).collect() # 步骤2:筛选符合条件用户的听歌记录,计算每位用户每月的使用频率 filtered_tracks = tracks_rdd.filter(lambda track: track.userid in eligible_user_ids) \ .map(lambda track: ((track.userid, track.year, track.month), 1)) \ .reduceByKey(lambda a, b: a + b) # 步骤3:按用户分组,按年月倒序排序后取最近6条记录 grouped_by_user = filtered_tracks.map(lambda x: (x[0][0], (x[0][1], x[0][2], x[1]))) \ .groupByKey() \ .mapValues(lambda records: sorted(records, key=lambda x: (x[0], x[1]), reverse=True)[:6]) # 扁平化结果得到最终输出格式 top_6_per_user = grouped_by_user.flatMap(lambda x: [((x[0], (record[2]))) for record in x[1]]) # 收集结果 result = top_6_per_user.collect() return result # 定义输入参数 k = 5 timestamp = datetime(2009, 4, 8) # 获取按日期排序的前K用户听歌记录,仅保留每位用户最近6条数据 result = k_users_to_retain(k, timestamp) # 展示结果 for record in result: print(record)
原代码部分输出
('user_000001', 62) ('user_000001', 822) ('user_000001', 700) ('user_000001', 671) ('user_000001', 680) ('user_000001', 760) ('user_000002', 486) ('user_000002', 645) ('user_000002', 673) ('user_000002', 791) ('user_000002', 608) ('user_000002', 953) ('user_000003', 351) ('user_000003', 50) ('user_000003', 140) ('user_000003', 401) ('user_000003', 88) ('user_000003', 183) ('user_000004', 22) ('user_000004', 504) ('user_000004', 35) ('user_000004', 39) ('user_000004', 539) ('user_000004', 693)
原代码存在的问题
- 性能隐患:将符合条件的用户ID收集到Driver端(
collect()),再用track.userid in eligible_user_ids过滤记录,用户量较大时会导致Driver内存溢出,且无法利用Spark分布式计算能力。 - 核心逻辑缺失:仅获取了用户最近6个月的使用数据,未完成平均值计算、流失风险标记、高价值用户筛选的核心需求。
- 排序逻辑不严谨:仅按年月倒序取前6条,但未处理用户某月份无记录的情况,可能导致统计的6个月并非连续的最近时段。
修正后的实现代码
from datetime import datetime from pyspark import SparkContext def k_users_to_retain(k, timestamp): # 步骤1:筛选注册时间早于指定时间的用户,保留(userid, registered) filtered_users = users_rdd.filter(lambda user: user.registered is not None and user.registered < timestamp) \ .map(lambda user: (user.userid, user.registered)) # 步骤2:关联用户与听歌记录,将年月转换为可排序的整数(如200903) track_with_date = tracks_rdd.map(lambda track: (track.userid, (track.year * 100 + track.month, 1))) \ .join(filtered_users) \ .map(lambda x: (x[0], x[1][0][0], x[1][0][1])) # (userid, year_month, count) # 步骤3:计算每位用户每月的使用频率 monthly_usage = track_with_date.map(lambda x: ((x[0], x[1]), x[2])) \ .reduceByKey(lambda a, b: a + b) \ .map(lambda x: (x[0][0], (x[0][1], x[1]))) # (userid, (year_month, total_count)) # 步骤4:计算用户的6个月/3个月使用指标,标记流失风险 def calculate_risk_metrics(monthly_records): # 按年月倒序排序,取最近6个月数据 sorted_records = sorted(monthly_records, key=lambda x: x[0], reverse=True)[:6] if len(sorted_records) < 3: # 数据不足3个月,不标记为流失风险 return (0, 0, 0, False) total_6 = sum([r[1] for r in sorted_records]) avg_6 = total_6 / len(sorted_records) # 取最近3个月数据 latest_3 = sorted_records[:3] total_3 = sum([r[1] for r in latest_3]) avg_3 = total_3 / len(latest_3) # 判断是否为流失风险用户 is_risk = avg_3 < avg_6 * 0.5 return (total_6, avg_6, avg_3, is_risk) user_metrics = monthly_usage.groupByKey() \ .mapValues(calculate_risk_metrics) \ .map(lambda x: (x[0], x[1][0], x[1][1], x[1][2], x[1][3])) # 步骤5:筛选流失风险用户,按6个月总使用量降序取前k位 top_k_users = user_metrics.filter(lambda x: x[4]) \ .sortBy(lambda x: x[1], ascending=False) \ .take(k) return top_k_users # 定义输入参数 k = 5 timestamp = datetime(2009, 4, 8) # 获取需要留存的前K高价值流失风险用户 result = k_users_to_retain(k, timestamp) # 展示结果 for record in result: print(f"用户ID: {record[0]}, 6个月总使用量: {record[1]}, 6个月平均值: {record[2]:.2f}, 3个月平均值: {record[3]:.2f}, 流失风险: {'是' if record[4] else '否'}")
修正点说明
- 优化分布式过滤:用RDD的
join操作关联用户与听歌记录,避免将用户ID收集到Driver端,提升大数量级场景下的性能。 - 补全核心逻辑:完成6个月/3个月平均值计算、流失风险标记,以及高价值用户的筛选。
- 处理数据缺失:对不足3个月数据的用户跳过风险标记,避免统计误差。
- 规范日期排序:将年月转换为整数格式,确保时间排序的准确性。
内容的提问来源于stack exchange,提问作者Yoel Ha
相关产品推荐
相关产品推荐

