PySpark中如何获取CountVectorizer的词汇表及对应语料词频?
获取CountVectorizer词汇表及对应总词频(Spark 2.1)
在Spark 2.1版本中,CountVectorizerModel并没有直接提供获取词汇对应总词频的API,但我们可以通过对原始训练数据做简单的聚合统计来实现需求,下面是两种实用的方法:
方法一:直接统计所有词汇的词频并生成字典
这种方法适用于你需要所有出现在训练数据中且被CountVectorizer选中的词汇的词频:
- 首先将原始的数组列展开为单个词汇的行:
from pyspark.sql.functions import explode, col # 把raw列的数组拆分成每行一个单词 exploded_df = df.select(explode(col("raw")).alias("word"))
- 分组统计每个单词的总出现次数:
# 聚合得到每个词的总词频,collect转为本地行列表 word_count_rows = exploded_df.groupBy("word").count().collect() # 转换为目标字典格式 voc_counts = {row["word"]: row["count"] for row in word_count_rows}
执行后voc_counts就会得到你想要的结果:{'a': 3, 'b': 3, 'c': 2}
方法二:严格匹配CountVectorizerModel的词汇表
如果你的CountVectorizer设置了minDF等过滤参数,只想保留模型最终选中的词汇(也就是model.vocabulary里的词汇),可以在统计后做一次筛选:
# 先获取模型的词汇表 vocabulary = model.vocabulary # 统计所有词的词频 word_count_dict = {row["word"]: row["count"] for row in exploded_df.groupBy("word").count().collect()} # 只保留词汇表中存在的词及其词频 voc_counts = {word: word_count_dict[word] for word in vocabulary}
补充说明
CountVectorizer的词汇表是基于训练数据中词汇的出现频率排序的(默认按词频降序),并且会过滤掉出现次数低于minDF阈值的词汇,所以我们统计训练数据的词频,完全对应模型词汇表的词频逻辑。
内容的提问来源于stack exchange,提问作者plalanne
相关产品推荐
相关产品推荐

