Scala实现按行数拆分Spark DataFrame(每批最多20000行)
Spark DataFrame 按固定行数拆分实现方案
方案一:基于行号分组拆分(支持顺序保留)
适合需要严格按全局顺序拆分的场景,通过添加行号并计算分组ID,遍历生成小DataFrame。
代码示例(Scala)
import org.apache.spark.sql.functions.{row_number, floor, col} import org.apache.spark.sql.expressions.Window // 假设原DataFrame为originalDF // 1. 添加全局行号(替换"排序字段"为你需要的排序列,如时间、ID等) val windowSpec = Window.orderBy("排序字段") val dfWithRowNum = originalDF.withColumn("row_num", row_number().over(windowSpec)) // 2. 计算分组ID,每20000行一组 val dfWithGroup = dfWithRowNum.withColumn("group_id", floor((col("row_num") - 1) / 20000)) // 3. 提取所有分组ID并生成小DataFrame val groupIds = dfWithGroup.select("group_id").distinct().collect().map(_.getInt(0)).sorted val smallDFs = groupIds.map(id => dfWithGroup.filter(col("group_id") === id).drop("row_num", "group_id") )
无顺序要求的优化版
如果不需要保留全局顺序,用monotonically_increasing_id()替代全局排序,避免shuffle开销:
import org.apache.spark.sql.functions.{monotonically_increasing_id, floor, col} val dfWithId = originalDF.withColumn("unique_id", monotonically_increasing_id()) val dfWithGroup = dfWithId.withColumn("group_id", floor(col("unique_id") / 20000)) val groupIds = dfWithGroup.select("group_id").distinct().collect().map(_.getLong(0)).sorted val smallDFs = groupIds.map(id => dfWithGroup.filter(col("group_id") === id).drop("unique_id", "group_id") )
方案二:自定义分区器拆分(分布式友好)
适合需要分布式处理拆分后数据的场景,通过自定义分区器按行数分配分区,再将每个分区转为小DataFrame。
代码示例(Scala)
import org.apache.spark.Partitioner // 自定义按行数分区的分区器 class RowCountPartitioner(rowsPerPartition: Int, totalRows: Long) extends Partitioner { override def numPartitions: Int = math.ceil(totalRows.toDouble / rowsPerPartition).toInt override def getPartition(key: Any): Int = { val rowId = key.asInstanceOf[Long] math.min(rowId / rowsPerPartition, numPartitions - 1).toInt } } // 使用自定义分区器拆分 val totalRows = originalDF.count() val dfWithId = originalDF.withColumn("unique_id", monotonically_increasing_id()) // 将DataFrame转为带键的RDD,按自定义分区器分区 val partitionedRDD = dfWithId.rdd.keyBy(_.getAs[Long]("unique_id")) .partitionBy(new RowCountPartitioner(20000, totalRows)) .map(_._2) // 将每个分区转为小DataFrame val smallDFsFromRDD = partitionedRDD.glom().map(rdd => spark.createDataFrame(rdd, originalDF.schema))
关键注意事项
- 避免将全量数据拉到Driver端:上述示例中
collect()仅用于获取分组ID(最多50个,100万/2万),不会造成内存压力。 - 写入文件替代生成小DF:如果最终目的是输出文件,可直接用
dfWithGroup.write.partitionBy("group_id").parquet("/输出路径"),无需单独生成小DataFrame。 - Python版本适配:将Scala API替换为PySpark对应方法即可,如
pyspark.sql.functions.row_number、Window等用法一致。
内容的提问来源于stack exchange,提问作者vicmartin
相关产品推荐
相关产品推荐

