如何在Spark Scala中按d列数组指定的列拆分字符串为数组?
根据指定列列表拆分DataFrame对应列的字符串(Spark Scala实现)
现有如下Spark DataFrame,需要根据d列数组中指定的列名,将对应列的逗号分隔字符串拆分为数组(或按需展开)。例如第二行d列的值为["b", "c"],则需拆分b列和c列的字符串。
DataFrame示例:
+----------+--------+-------------------+------+ | a| b| c| d| +----------+--------+-------------------+------+ | India, US| 4,5,6|apple, banana, pear|[a, b]| |Canada, LA|13,14,15| berry, strawberry|[b, c]| | UK, US|22,23,24| carrot, plum|[a, c]| +----------+--------+-------------------+------+
创建该DataFrame的代码:
import spark.implicits._ val data = Seq( ("India, US", "4,5,6", "apple, banana, pear", Array("a", "b")), ("Canada, LA", "13,14,15", "berry, strawberry", Array("b", "c")), ("UK, US", "22,23,24", "carrot, plum", Array("a", "c")) ) val df = data.toDF("a", "b", "c", "d") df.show()
实现方案
方案1:将指定列拆分为数组(保留原结构)
通过条件判断动态处理每一列:如果列名存在于d列的数组中,就将该列的字符串按逗号拆分转为数组,否则保留原列值。
import org.apache.spark.sql.functions._ // 定义所有待处理的列名 val targetColumns = Seq("a", "b", "c") // 生成每一列的处理逻辑 val processedCols = targetColumns.map(colName => { when(array_contains(col("d"), colName), split(col(colName), ",\\s*")) .otherwise(col(colName)) .alias(colName) }) ++ Seq(col("d")) // 保留原d列 val resultDf = df.select(processedCols: _*) resultDf.show(truncate = false)
执行结果:
+----------------+-----------+---------------------------+------+ |a |b |c |d | +----------------+-----------+---------------------------+------+ |[India, US] |[4, 5, 6] |apple, banana, pear |[a, b]| |Canada, LA |[13, 14, 15]|[berry, strawberry] |[b, c]| |[UK, US] |22,23,24 |[carrot, plum] |[a, c]| +----------------+-----------+---------------------------+------+
方案2:将拆分后的数组展开为多行(可选)
如果需要把拆分后的数组展开成多行,可结合flatMap实现(假设所有拆分后的数组长度一致):
// 先执行方案1的列处理逻辑 val processedDf = df.select(processedCols: _*) // 展开拆分后的数组为多行 val explodedDf = processedDf.flatMap(row => { val targetColNames = row.getAs[Array[String]]("d") // 获取所有需要展开的列的数组值 val explodedValues = targetColNames.map(colName => row.getAs[Seq[String]](colName)) // 按索引遍历生成新行 explodedValues.head.indices.map(idx => { val newRowValues = df.columns.map(colName => { if (targetColNames.contains(colName)) explodedValues(targetColNames.indexOf(colName))(idx) else row.getAs[Any](colName) }) Row.fromSeq(newRowValues) }) }).toDF(df.columns: _*) explodedDf.show()
执行结果:
+--------+---+-------------------+------+ | a| b| c| d| +--------+---+-------------------+------+ | India| 4|apple, banana, pear|[a, b]| | US| 5|apple, banana, pear|[a, b]| | Canada,| 13| berry, strawberry|[b, c]| | LA| 14| berry, strawberry|[b, c]| | UK|22,| carrot, plum|[a, c]| | US|23,| carrot, plum|[a, c]| +--------+---+-------------------+------+
关键说明
split(col(colName), ",\\s*")使用正则表达式拆分,自动去除逗号后的空格;array_contains用于判断当前列是否属于d列指定的待拆分列表;- 方案2中若拆分后的数组长度不一致,可根据实际需求调整遍历逻辑。
内容的提问来源于stack exchange,提问作者Aravind Yarram
相关产品推荐
相关产品推荐

