如何在PySpark中按ID分组生成速度数据的直方图
问题:PySpark按ID分组生成速度直方图
我有一个包含日期时间、ID和速度字段的数据集,需要用PySpark为每个ID生成速度的直方图数据(包含区间起止点与计数)。
示例数据
df = spark.createDataFrame( [ ("2023-06-01 07:09:17", "abc", 4.5), ("2023-06-01 07:09:18", "abc", 9.1), ("2023-06-01 07:09:19", "abc", 3.2), ("2023-06-01 07:10:06", "ddc", 5.1), ("2023-06-01 07:09:07", "ddc", 3.6), ("2023-06-01 07:09:08", "ddc", 2.6) ], ["date_time", "id", "velocity"] )
现有尝试
我已经能生成所有速度值的整体直方图,代码如下:
df.filter(col("velocity").isNotNull()).rdd.histogram(list(range(0, 100, 1)))
但尝试按ID分组生成直方图时,三种方法均报错:
方法一报错
# 报错:'GroupedData' object has no attribute 'rdd' df.filter(col("velocity").isNotNull()).groupBy("id").rdd.histogram(list(range(0, 100, 1)))
方法二报错
# 报错最终为:TypeError: 'str' object is not callable df.filter(col("velocity").isNotNull()).rdd.groupBy("id").histogram(list(range(0, 100, 1)))
方法三报错
# 报错最终为:TypeError: '>' not supported between instances of 'tuple' and 'int' df.filter(col("velocity").isNotNull()).select("id", "velocity").rdd.groupByKey().histogram(list(range(0, 100, 1)))
解决方案
方法1:分组后用UDF计算直方图
通过groupBy("id")收集每个ID的所有速度值,再用自定义UDF调用histogram方法,输出和原生方法一致的结构:
from pyspark.sql import functions as F from pyspark.sql.types import ArrayType, StructType, StructField, DoubleType, IntegerType # 定义直方图输出的Schema histogram_schema = StructType([ StructField("bins", ArrayType(DoubleType()), nullable=False), StructField("counts", ArrayType(IntegerType()), nullable=False) ]) # 自定义计算直方图的UDF def compute_histogram(velocities, bins): if not velocities: return (bins, [0]*len(bins[:-1])) rdd = spark.sparkContext.parallelize(velocities) return rdd.histogram(bins) histogram_udf = F.udf(compute_histogram, histogram_schema) # 定义分箱参数 bins = list(range(0, 100, 1)) # 执行分组计算 result_df = df.filter(F.col("velocity").isNotNull()) \ .groupBy("id") \ .agg(F.collect_list("velocity").alias("velocities")) \ .withColumn("histogram", histogram_udf(F.col("velocities"), F.lit(bins))) \ .select("id", "histogram.bins", "histogram.counts") result_df.show(truncate=False)
方法2:Spark SQL区间统计(大数据场景更高效)
预先定义速度区间,直接分组统计每个区间的计数,无需UDF,性能更优:
# 定义分箱区间 bins = list(range(0, 100, 1)) # 生成区间匹配表达式 bin_expr = F.when(F.col("velocity") < bins[0], f"{bins[0]}以下") for i in range(len(bins)-1): bin_expr = bin_expr.when( (F.col("velocity") >= bins[i]) & (F.col("velocity") < bins[i+1]), f"{bins[i]}-{bins[i+1]}" ) bin_expr = bin_expr.otherwise(f"{bins[-1]}以上").alias("velocity_bin") # 分组统计每个ID的区间计数 result_df = df.filter(F.col("velocity").isNotNull()) \ .select("id", "velocity", bin_expr) \ .groupBy("id", "velocity_bin") \ .count() \ .orderBy("id", "velocity_bin") result_df.show(truncate=False)
结果说明
- 方法1输出每个ID对应的分箱数组和计数数组,和原生
rdd.histogram输出结构一致; - 方法2输出每个ID每个区间的具体计数,更直观,适合大规模数据处理。
内容的提问来源于stack exchange,提问作者CopyOfA
相关产品推荐
相关产品推荐

