如何在PySpark中统计各数值区间内的元素数量
在PySpark中统计整数列的区间元素数量
针对你提出的“统计指定区间内整数元素数量”的需求,下面提供几种实用的实现方式,涵盖固定区间、动态区间以及可选的Window函数用法:
方法一:固定区间用CASE WHEN直接匹配
如果区间数量固定、范围已知,直接通过CASE WHEN给每个数值打上对应区间标签,再分组统计即可。
代码示例
假设你的DataFrame名为df,整数列是value:
from pyspark.sql import SparkSession from pyspark.sql.functions import col, count, when # 初始化Spark会话 spark = SparkSession.builder.appName("RangeCountDemo").getOrCreate() # 构造示例数据 sample_data = [(1,), (10,), (101,), (150,), (200,), (250,)] df = spark.createDataFrame(sample_data, ["value"]) # 标记每个值所属的区间 range_labeled_df = df.withColumn( "range", when((col("value") >= 0) & (col("value") <= 100), "0-100") .when((col("value") > 100) & (col("value") <= 200), "101-200") .when((col("value") > 200) & (col("value") <= 300), "201-300") # 按需添加更多区间 ) # 分组统计区间内元素数量 result = range_labeled_df.groupBy("range").agg(count("value").alias("count")) result.show()
输出
+---------+-----+ | range|count| +---------+-----+ | 0-100 | 2| |101-200 | 3| |201-300 | 1| +---------+-----+
方法二:动态区间列表的处理(灵活适配可变区间)
如果区间是动态生成的(比如从配置文件读取),可以把区间列表转换成DataFrame,通过范围Join匹配数值和区间,再统计数量。
代码示例
from pyspark.sql.functions import lit, concat # 定义你的动态区间列表 intervals = [(0, 100), (100, 200), (200, 300)] # 转换为带区间标签的DataFrame interval_df = spark.createDataFrame(intervals, ["lower", "upper"]).withColumn( "range", concat(lit(col("lower") + 1), lit("-"), col("upper")) # 生成示例要求的区间标签,比如(100,200)对应101-200 ) # 范围Join:匹配value落在(lower, upper]区间的记录 joined_df = df.join( interval_df, (col("value") > col("lower")) & (col("value") <= col("upper")), "inner" ) # 分组统计 result = joined_df.groupBy("range").agg(count("value").alias("count")) result.show()
说明
- 这里的匹配逻辑是
value > lower且value <= upper,和你给出的示例规则一致(比如200归到101-200区间) - 若你的区间是左闭右开,只需调整Join条件即可
方法三:用Window函数(保留原始数据的同时统计)
如果需要保留原始数据,同时显示每条记录所属区间的总数量,可以结合Window函数实现:
from pyspark.sql.window import Window # 先标记区间(同方法一) range_labeled_df = df.withColumn( "range", when((col("value") >= 0) & (col("value") <= 100), "0-100") .when((col("value") > 100) & (col("value") <= 200), "101-200") .when((col("value") > 200) & (col("value") <= 300), "201-300") ) # 按区间定义Window分区 window_spec = Window.partitionBy("range") # 计算每个区间的元素数量 result = range_labeled_df.withColumn("count", count("value").over(window_spec)).distinct() result.show()
输出
+-----+---------+-----+ |value| range|count| +-----+---------+-----+ | 1| 0-100 | 2| | 10| 0-100 | 2| | 101|101-200 | 3| | 150|101-200 | 3| | 200|101-200 | 3| | 250|201-300 | 1| +-----+---------+-----+
如果只需要区间和计数的聚合结果,前两种方法的效率更高。
内容的提问来源于stack exchange,提问作者Aman
相关产品推荐
相关产品推荐

