如何用Scala为多列Spark DataFrame新增存储所有列值的ArrayType列
实现代码
你可以直接用Spark内置的array函数完成列转数组的需求,完整Scala实现如下:
import org.apache.spark.sql.SparkSession import org.apache.spark.sql.functions._ import org.apache.spark.sql.types.{DoubleType, StructField, StructType, Row} object ColumnToArrayExample { def main(args: Array[String]): Unit = { // 初始化SparkSession val spark = SparkSession.builder() .appName("Col2Array") .master("local[*]") .getOrCreate() // -------------------------- // 步骤1:生成测试用DataFrame(如果已有DF可跳过此段) // 生成col_1到col_2048的列名序列 val colNames = (1 to 2048).map(idx => s"col_$idx") // 生成2行测试数据,每行都是随机Double值 val testRows = Seq.fill(2)(Row.fromSeq(Seq.fill(2048)(scala.util.Random.nextDouble()))) // 构造表结构 val schema = StructType(colNames.map(colName => StructField(colName, DoubleType, nullable = false))) var df = spark.createDataFrame(spark.sparkContext.parallelize(testRows), schema) // -------------------------- // 步骤2:新增array_col数组列 // 获取所有目标列的Column对象,按col_1到col_2048排序 val targetColumns = colNames.map(col(_)) // 用array函数合并多列为数组列,:_*是Scala语法,将序列展开为可变参数 df = df.withColumn("array_col", array(targetColumns: _*)) // 验证输出结果 df.select("array_col").show(2, truncate = false) spark.stop() } }
适配已有DataFrame的场景
如果你是从外部数据源读取的已有DataFrame,不需要自己生成列名,可以用如下方式自动提取所有col_xxx格式的列:
val targetColumns = df.columns // 过滤所有符合col_数字格式的列 .filter(colName => colName.matches("col_\\d+")) // 按列名后缀的数字排序,保证数组元素顺序和col_1到col_2048一致 .sortBy(colName => colName.split("_")(1).toInt) .map(col(_)) val resultDF = df.withColumn("array_col", array(targetColumns: _*))
注意事项
- 参与合并的所有列类型必须一致,如果存在类型差异,需要先用
cast函数统一转换为相同类型后再合并,否则会触发类型不匹配报错 - 生成的
array_col默认类型为ArrayType(元素类型, 包含空值 = false),符合需求中的ArrayType要求
内容的提问来源于stack exchange,提问作者Alain ux
相关产品推荐
相关产品推荐

