Spark SQL按分区后用random()排序抽取样本报错的解决咨询
Spark SQL窗口函数ORDER BY中使用random()报错问题解答
问题描述
我尝试使用Spark SQL按region_id和marketplace_id对客户样本集进行分区,并在每个分区结果中随机选取100,000名用户。想请教:能否在Spark SQL的partition by子句对应的order by中使用random()?以下使用random()的代码始终报错,但移除random()后可正常运行,求解答!
报错代码与正常代码对比
customer_panel_s3_location = f"s3://my-bucket/region_id={region_id}/marketplace_id={marketplace_id}/" customer_panel_table = spark.read.parquet(customer_panel_s3_location) customer_panel_table.createOrReplaceTempView("customer_panel") dataset_date = '2023-03-16' customer_sample_size = 100000 partition_customers_by = 'region_id, marketplace_id' # 以下代码执行报错 customer_panel_df = spark.sql(f""" SELECT * FROM ( SELECT * , row_number() over (partition by {partition_customers_by} order by random()) AS rn FROM customer_panel AS c WHERE c.target_date < CAST('{dataset_date}' AS DATE) AND c.target_date >= date_sub(CAST('{dataset_date}' AS DATE), 7) ) t WHERE t.rn <= bigint({customer_sample_size}) """) # 移除random()后可正常运行 customer_panel_df = spark.sql(f""" SELECT * , row_number() over (partition by {partition_customers_by} order by {partition_customers_by}) AS rn FROM customer_panel AS c WHERE c.target_date < CAST('{dataset_date}' AS DATE) AND c.target_date >= date_sub(CAST('{dataset_date}' AS DATE), 7) """) print(f"Row count of {table_name}: {customer_panel_df.count():,}")
问题原因与解决方案
核心原因
Spark SQL原生不支持random()函数,生成随机数需要使用内置的rand()函数(Spark 2.x及以上版本均支持)。你使用random()会触发函数不存在的报错,而移除后改用分区字段排序,语法合法所以能正常运行。
修正后的代码
将order by random()替换为order by rand()即可解决问题,如果需要固定随机种子保证抽样结果可复现,可以给rand()传入一个整数参数(比如rand(123)):
customer_panel_df = spark.sql(f""" SELECT * FROM ( SELECT * , row_number() over (partition by {partition_customers_by} order by rand()) AS rn FROM customer_panel AS c WHERE c.target_date < CAST('{dataset_date}' AS DATE) AND c.target_date >= date_sub(CAST('{dataset_date}' AS DATE), 7) ) t WHERE t.rn <= bigint({customer_sample_size}) """)
额外优化建议
如果数据集规模较大,使用窗口函数抽样的性能可能不如直接使用Spark的sampleBy方法,该方法可以按指定分区字段直接抽样,示例如下:
# 先获取每个分区的大致数据量,计算抽样比例 partition_counts = customer_panel_table.groupBy(partition_customers_by).count().collect() fractions = { (row.region_id, row.marketplace_id): customer_sample_size / row.count for row in partition_counts } # 按分区抽样 customer_panel_df = customer_panel_table.sampleBy(partition_customers_by, fractions=fractions, seed=123)
内容的提问来源于stack exchange,提问作者user1330974
相关产品推荐
相关产品推荐

