【发布时间】:2018-05-16 22:23:33
【问题描述】:
我有一个问题。我有一个带有几列的 spark 数据框,如下所示:
id 颜色
1 红、蓝、黑
2 红、绿
3 蓝色、黄色、绿色
...
我还有一个地图文件,看起来像:
红色,0
蓝色,1
绿色,2
黑色,3
黄色,4
我需要做的是将颜色名称映射成不同的id,比如将“红、蓝、黑”映射成[1,1,0,1,0]的数组。 我这样写代码:
def mapColor(label_string:String):Array[Int]={
var labels = label_string.split(",")
var index_array = new Array[Int](COLOR_LENGTH)
for (label<-labels){
if(COLOR_MAP.contains(label)){
index_array(COLOR_MAP(label))=1
}
else{
//dictionary does not contain the label, the last index set to be one
index_array(COLOR_LENGTH-1)=1
}
}
index_array
}
COLOR_LENGTH 是字典的长度,COLOR_MAP 是包含字符串->id 关系的字典。
我这样调用这个函数:
val color_function = udf(mapColor:(String)=>Array[Int])
sql.withColumn("color_idx",color_function(col("Color")))
由于我有多个列需要这个操作,但不同的列需要不同的字典。目前,我为每一列复制了这个函数(只需更改字典和长度信息)。但是代码看起来很乏味。有没有什么方法,可以把长度和字典传给映射函数,比如
def map(label_string:String,map:Map[String,Integer],len:Int):Array[Int]
但是我应该如何在 spark 数据框中调用这个函数呢?由于我无法在声明中传递参数
val color_function = udf(mapColor:(String)=>Array[Int])
【问题讨论】:
标签: scala apache-spark user-defined-functions