【问题标题】:One-hot encoding in pyspark with Multiple 1's in a rowpyspark 中的单热编码,连续多个 1
【发布时间】:2019-04-06 01:38:28
【问题描述】:

我有一个 Python 数据框 final_df 如下:

这些行有重复的ID 值。如何使用 pyspark 获得如下的 one-hot 编码输出?

我已将其转换为 spark 数据框:

spark_df = sqlContext.createDataFrame(final_df)

然后在CONCEPTS列中收集唯一值如下:

types = spark_df.select("CONCEPTS").distinct().rdd.flatMap(lambda x: x).collect()

但是当我调用以下命令时:

types_expr = [F.when((F.col("CONCEPTS") == ty), 1).otherwise(0).alias(ty) for ty in types]
df = spark_df.select("ID", *types_expr)
df.show()

我得到以下信息:

与此类似的其他问题的解决方案不会为一行产生多个 1。

【问题讨论】:

    标签: python pyspark one-hot-encoding


    【解决方案1】:

    您可以使用 GroupedData 类的 pivot 函数,因为您只使用 1 和 0。示例代码:

    l =[( 115        ,'A' ),
    ( 116        , 'B' ),
    ( 118        , 'C' ),
    ( 121        , 'D' ),
    ( 125        , 'E' ),
    ( 127        , 'F' ),
    ( 127        , 'G' ),
    ( 127        , 'H' ),
    ( 136        , 'I' ),
    ( 136        , 'J' )]
    
    df = spark.createDataFrame(l, ['id','concepts'])
    df.groupBy('id').pivot('concepts').count().show()
    

    将导致以下数据框:

    +---+----+----+----+----+----+----+----+----+----+----+   
    | id|   A|   B|   C|   D|   E|   F|   G|   H|   I|   J| 
    +---+----+----+----+----+----+----+----+----+----+----+ 
    |136|null|null|null|null|null|null|null|null|   1|   1| 
    |116|null|   1|null|null|null|null|null|null|null|null| 
    |115|   1|null|null|null|null|null|null|null|null|null| 
    |127|null|null|null|null|null|   1|   1|   1|null|null| 
    |118|null|null|   1|null|null|null|null|null|null|null| 
    |125|null|null|null|null|   1|null|null|null|null|null| 
    |121|null|null|null|   1|null|null|null|null|null|null| 
    +---+----+----+----+----+----+----+----+----+----+----+
    

    如果需要,用fill-函数替换空值。


    cmets 中的某个人询问如何使用 pandas 来做到这一点。方法基本相同,但需要的函数是pivot_table

    import pandas as pd
    import numpy as np
    
    l =[( 115        ,'A' ),
    ( 116        , 'B' ),
    ( 118        , 'C' ),
    ( 121        , 'D' ),
    ( 125        , 'E' ),
    ( 127        , 'F' ),
    ( 127        , 'G' ),
    ( 127        , 'H' ),
    ( 136        , 'I' ),
    ( 136        , 'J' )]
    
    df = pd.DataFrame(l,columns=['id','concepts'] )
    df.pivot_table(index='id', columns='concepts', aggfunc=len)
    

    输出:

    concepts    A    B    C    D    E    F    G    H    I    J
    id                                                        
    115       1.0  NaN  NaN  NaN  NaN  NaN  NaN  NaN  NaN  NaN
    116       NaN  1.0  NaN  NaN  NaN  NaN  NaN  NaN  NaN  NaN
    118       NaN  NaN  1.0  NaN  NaN  NaN  NaN  NaN  NaN  NaN
    121       NaN  NaN  NaN  1.0  NaN  NaN  NaN  NaN  NaN  NaN
    125       NaN  NaN  NaN  NaN  1.0  NaN  NaN  NaN  NaN  NaN
    127       NaN  NaN  NaN  NaN  NaN  1.0  1.0  1.0  NaN  NaN
    136       NaN  NaN  NaN  NaN  NaN  NaN  NaN  NaN  1.0  1.0
    

    【讨论】:

    • 谢谢@cronoik。我从来没有想过数据透视表。这是一个更简洁的解决方案。
    • 如何在 python 中使用 pandas?谢谢
    • @Neo:我为 pandas 添加了一个解决方案。请在以后提出您自己的问题。
    • @cronoik 谢谢
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 2023-03-18
    • 2019-09-06
    • 1970-01-01
    • 2021-08-14
    • 1970-01-01
    • 1970-01-01
    • 2021-11-02
    相关资源
    最近更新 更多