【问题标题】:Preserve index-string correspondence spark string indexer保留索引字符串对应的火花字符串索引器
【发布时间】:2016-02-11 17:47:02
【问题描述】:

Spark 的 StringIndexer 非常有用,但通常需要检索生成的索引值和原始字符串之间的对应关系,并且似乎应该有一种内置的方法来完成此操作。我将使用Spark documentation 中的这个简单示例进行说明:

from pyspark.ml.feature import StringIndexer

df = sqlContext.createDataFrame(
    [(0, "a"), (1, "b"), (2, "c"), (3, "a"), (4, "a"), (5, "c")],
    ["id", "category"])
indexer = StringIndexer(inputCol="category", outputCol="categoryIndex")
indexed_df = indexer.fit(df).transform(df)

这个简化的案例告诉我们:

+---+--------+-------------+
| id|category|categoryIndex|
+---+--------+-------------+
|  0|       a|          0.0|
|  1|       b|          2.0|
|  2|       c|          1.0|
|  3|       a|          0.0|
|  4|       a|          0.0|
|  5|       c|          1.0|
+---+--------+-------------+

一切都很好,但对于许多用例,我想知道我的原始字符串和索引标签之间的映射。我能想到的最简单的方法是这样的:

   In [8]: indexed.select('category','categoryIndex').distinct().show()
+--------+-------------+
|category|categoryIndex|
+--------+-------------+
|       b|          2.0|
|       c|          1.0|
|       a|          0.0|
+--------+-------------+

如果需要,我可以将其结果存储为字典或类似内容:

In [12]: mapping = {row.categoryIndex:row.category for row in
           indexed.select('category','categoryIndex').distinct().collect()}

In [13]: mapping
Out[13]: {0.0: u'a', 1.0: u'c', 2.0: u'b'}

我的问题是:由于这是一项常见的任务,而且我猜测(但当然可能是错误的)字符串索引器无论如何都会以某种方式存储此映射,有没有办法更多地完成上述任务简单地?

我的解决方案或多或少是直截了当的,但对于大型数据结构,这涉及到(也许)我可以避免的大量额外计算。想法?

【问题讨论】:

    标签: python apache-spark apache-spark-sql pyspark apache-spark-ml


    【解决方案1】:

    标签映射可以从列元数据中提取:

    meta = [
        f.metadata for f in indexed_df.schema.fields if f.name == "categoryIndex"
    ]
    meta[0]
    ## {'ml_attr': {'name': 'category', 'type': 'nominal', 'vals': ['a', 'c', 'b']}}
    

    其中ml_attr.vals 提供位置和标签之间的映射:

    dict(enumerate(meta[0]["ml_attr"]["vals"]))
    ## {0: 'a', 1: 'c', 2: 'b'}
    

    Spark 1.6+

    您可以使用IndexToString 将数值转换为标签。这将使用如上所示的列元数据。

    from pyspark.ml.feature import IndexToString
    
    idx_to_string = IndexToString(
        inputCol="categoryIndex", outputCol="categoryValue")
    
    idx_to_string.transform(indexed_df).drop("id").distinct().show()
    ## +--------+-------------+-------------+
    ## |category|categoryIndex|categoryValue|
    ## +--------+-------------+-------------+
    ## |       b|          2.0|            b|
    ## |       a|          0.0|            a|
    ## |       c|          1.0|            c|
    ## +--------+-------------+-------------+
    

    火花

    这是一种肮脏的技巧,但您可以简单地从 Java 索引器中提取标签,如下所示:

    from pyspark.ml.feature import StringIndexerModel
    
    # A simple monkey patch so we don't have to _call_java later 
    def labels(self):
        return self._call_java("labels")
    
    StringIndexerModel.labels = labels
    
    # Fit indexer model
    indexer = StringIndexer(inputCol="category", outputCol="categoryIndex").fit(df)
    
    # Extract mapping
    mapping = dict(enumerate(indexer.labels()))
    mapping
    ## {0: 'a', 1: 'c', 2: 'b'}
    

    【讨论】:

    • 只是 enumerate(indexer.labels()) 不保证相同的顺序,因为 stringIndexer 默认使用频率来索引类别
    • Pyspark 1.6+ 解决方案与简单的indexed_df.drop('id').distinct().show() 有何不同...categorycategoryValue 是相同的
    猜你喜欢
    • 2017-01-27
    • 2017-07-22
    • 1970-01-01
    • 1970-01-01
    • 2010-10-27
    • 2014-11-06
    • 2015-11-17
    • 2015-04-23
    • 1970-01-01
    相关资源
    最近更新 更多