【问题标题】:Spark struct represented by OneHotEncoder由 OneHotEncoder 表示的 Spark 结构体
【发布时间】:2018-07-24 20:16:19
【问题描述】:

我有一个包含两列的数据框,

+---+-------+
| id|  fruit|
+---+-------+
|  0|  apple|
|  1| banana|
|  2|coconut|
|  1| banana|
|  2|coconut|
+---+-------+

我还有一个包含所有项目的通用列表,

fruitList: Seq[String] = WrappedArray(apple, coconut, banana)

现在我想在数据框中创建一个新列,其中包含 1、0 的数组,其中 1 表示该项目存在,如果该行不存在该项目,则为 0。

期望的输出

    +---+-----------+
    | id|  fruitlist|
    +---+-----------+
    |  0|  [1,0,0]  |
    |  1| [0,1,0]   |
    |  2|[0,0,1]    |
    |  1| [0,1,0]   |
    |  2|[0,0,1]    |
    +---+-----------+

这是我尝试过的,

import org.apache.spark.ml.feature.{OneHotEncoder, StringIndexer}

val df = spark.createDataFrame(Seq(
  (0, "apple"),
  (1, "banana"),
  (2, "coconut"),
  (1, "banana"),
  (2, "coconut")
)).toDF("id", "fruit")

df.show
import org.apache.spark.sql.functions._
val fruitList = df.select(collect_set("fruit")).first().getAs[Seq[String]](0)
print(fruitList)

我尝试使用 OneHotEncoder 解决这个问题,但转换为密集向量后的结果是这样的,这不是我需要的。

    +---+-------+----------+-------------+---------+
| id|  fruit|fruitIndex|     fruitVec|       vd|
+---+-------+----------+-------------+---------+
|  0|  apple|       2.0|    (2,[],[])|[0.0,0.0]|
|  1| banana|       1.0|(2,[1],[1.0])|[0.0,1.0]|
|  2|coconut|       0.0|(2,[0],[1.0])|[1.0,0.0]|
|  1| banana|       1.0|(2,[1],[1.0])|[0.0,1.0]|
|  2|coconut|       0.0|(2,[0],[1.0])|[1.0,0.0]|
+---+-------+----------+-------------+---------+

【问题讨论】:

  • 我不确定我是否理解你的最后一句话。您显示的最后一个表格中的结果是稀疏的。
  • @eliasah,我已经用正确的密集向量表示编辑了这个问题。
  • 您的向量长度为​​ 2,因为您在 OneHotEncoder 上使用了默认参数 dropLast=True。但无论如何,它不会保持您想要的顺序(因为 StringIndexer 按出现顺序排列项目)。使用@Ramesh Maharjan 回答
  • @Arius,谢谢你的提示,设置 encoder.setDropLast(false) 将显示完整的向量元素。onehotestimator 会保持顺序吗?
  • @Masterbuilder 是的 OneHotEncoder 会,但之前的 StringIndexer 会破坏它。

标签: apache-spark apache-spark-mllib apache-spark-ml


【解决方案1】:

如果你有一个集合

val fruitList: Seq[String] = Array("apple", "coconut", "banana")

然后你可以使用 inbuilt functionsudf function

内置函数(数组、when 和 lit)

import org.apache.spark.sql.functions._
df.withColumn("fruitList", array(fruitList.map(x => when(lit(x) === col("fruit"),1).otherwise(0)): _*)).show(false)

udf 函数

import org.apache.spark.sql.functions._
def containedUdf = udf((fruit: String) => fruitList.map(x => if(x == fruit) 1 else 0))

df.withColumn("fruitList", containedUdf(col("fruit"))).show(false)

这应该给你

+---+-------+---------+
|id |fruit  |fruitList|
+---+-------+---------+
|0  |apple  |[1, 0, 0]|
|1  |banana |[0, 0, 1]|
|2  |coconut|[0, 1, 0]|
|1  |banana |[0, 0, 1]|
|2  |coconut|[0, 1, 0]|
+---+-------+---------+

udf 函数易于理解且直截了当,可以处理原始数据类型,但如果可以使用优化和快速的内置函数来完成相同的任务,则应避免使用

希望回答对你有帮助

【讨论】:

  • 您介意吗,请向我解释一下 Ramesh ?我不确定这里发生了什么。
  • 你不明白什么? @eliasah。
  • 我不明白这个问题以及与您的答案的关系。我不质疑你在这里所说的。没什么好争论的。
  • 我只是试图帮助 OP 生成预期的输出。就这样。在他的输出中(他试图为匹配字符串生成 1,为不匹配生成 0)。我希望它清楚。
  • 感谢您向我解释 OP 想要什么。 :)
猜你喜欢
  • 2017-07-06
  • 1970-01-01
  • 2017-08-10
  • 2016-03-14
  • 2020-01-30
  • 1970-01-01
  • 1970-01-01
  • 2019-03-05
  • 1970-01-01
相关资源
最近更新 更多