You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.17 22:05:20