Spark DataFrame分区字节计算方法触发时机及cache/checkpoint优化问询
问题解答
一、触发DataFrame计算的节点说明
首先明确:Spark中只有行动(Action)算子会触发实际任务执行,count()属于典型的行动算子,你列出的三个节点里,节点1和节点3会触发计算,节点2不会,具体原因:
- 节点1:
dropDuplicates.count中的count是行动算子,会触发第一次作业,扫描DataFrame的指定分区列去重后计数 - 节点2:这里只是对已经计算好的三个Long类型数值做大小判断和返回,属于纯内存运算,没有Spark算子执行,不会触发作业
- 节点3:
df.count是第二次调用行动算子,会触发第二次全表扫描作业,统计DataFrame的总行数
二、方法优化方案
你当前的实现最坏情况下会对DataFrame执行2次全量扫描,在DataFrame血缘复杂、数据量大时会有严重的重复计算开销,完全可以通过优化减少重复计算:
1. 最优优化:合并2次Action为1次
不用依赖cache/checkpoint,直接用agg算子一次计算两个统计值,只触发1次作业,性能提升一倍:
private def getExpectedPartitionBytes(df: DataFrame, partitionNames: Seq[String] = Seq()) (implicit spark: SparkSession): Long = { // 一次Action同时计算分区去重数量、全表行数 val statRow = df.agg( countDistinct(partitionNames.map(col): _*).alias("partition_cnt"), count("*").alias("total_cnt") ).head() val partitionsCount = statRow.getLong(0) val totalCount = statRow.getLong(1) // 该逻辑从优化计划元数据取统计值,不会触发作业 val expectedTotalBytes = df.queryExecution.optimizedPlan.stats(spark.sessionState.conf) .sizeInBytes.toLong val expectedPartitionBytes = expectedTotalBytes / partitionsCount val maxExpectedPartitionBytes = df.dtypes.filter(t => !partitionNames.contains(t._1)).map(_._2).map { case "StringType" => 10 case "ByteType" => 1 case "ShortType" => 2 case "IntegerType" => 4 case "LongType" => 8 case "FloatType" => 4 case "DoubleType" => 8 case "TimestampType" => 6 case _ => 2 }.sum * totalCount / partitionsCount if (expectedTotalBytes > 0 && expectedPartitionBytes <= maxExpectedPartitionBytes) { expectedPartitionBytes } else { maxExpectedPartitionBytes } }
2. cache/checkpoint适用场景说明
如果你的DataFrame还会在当前方法外部被其他逻辑复用,可以在调用getExpectedPartitionBytes前对df执行df.cache(),所有逻辑执行完成后手动调用df.unpersist()释放缓存即可,不需要在方法内部加缓存逻辑避免污染原DataFrame的缓存状态。
如果这个方法是单次调用,上述合并Action的方案已经足够高效,不需要额外加缓存。checkpoint仅在DataFrame血缘极其复杂、单次计算开销极高时使用,普通场景下不需要,缓存已经足够满足需求。
内容的提问来源于stack exchange,提问作者Mardaunt
相关产品推荐
相关产品推荐

