如何在Pyspark.ml Word2vec中同时获取context向量和target向量
PySpark Word2Vec获取target向量解决方案
PySpark官方封装的Word2VecModel默认仅开放输入层的context向量公共访问接口,输出层的target向量没有直接的公共API可以调用,可通过访问模型底层私有Java对象的方式提取对应权重,具体操作如下:
训练阶段提取target向量
训练完成后可通过模型内置的_java_obj访问底层Scala实现的权重接口:
from pyspark.ml.feature import Word2Vec, Word2VecModel import pickle # 你的原有训练逻辑 wv = Word2Vec(settings) model = wv.fit(doc) # 1. 获取词汇表(target向量顺序与词汇表顺序完全对应) vocab = model.getVectors().select("word").rdd.flatMap(lambda x: x).collect() vector_size = model.getVectorSize() vocab_size = len(vocab) # 2. 从Java模型对象中提取输出层target向量 target_vectors_java = model._java_obj.wordVectors().word2vecModel().getOutputWeights() # 转换为Python可处理的二维数组 target_vectors = [ target_vectors_java[i*vector_size : (i+1)*vector_size] for i in range(vocab_size) ] # 转成字典方便按词汇查询 target_vector_dict = dict(zip(vocab, target_vectors))
注:以上接口基于Spark 3.x版本测试,如遇调用报错可打印dir(model._java_obj.wordVectors())查看对应版本的可用方法,调整属性调用路径即可。
保存与加载target向量
官方的Word2VecModel.save()方法不会序列化输出层权重,需要单独持久化target向量:
保存target向量
with open("target_vectors.pkl", "wb") as f: pickle.dump(target_vector_dict, f)
加载target向量
with open("target_vectors.pkl", "rb") as f: loaded_target_vectors = pickle.load(f)
注意事项
- 以上方法使用了PySpark私有API,跨版本兼容性不做保证,生产环境使用前需做好对应Spark版本的适配测试
- 如需更稳定的双向量访问能力,可考虑改用gensim库的Word2Vec实现,原生支持同时访问context和target两类向量
内容的提问来源于stack exchange,提问作者Mahmood
相关产品推荐
相关产品推荐

