我有一个图像数据集( id,url,功能),我在这些数据集上执行了所有图像之间的余弦相似性。其结果是一个具有以下结构的:
>>> cos_df.printSchema()
root
|-- id: integer (nullable = true)
|-- url: string (nullable = true)
|-- vec: vector (nullable = true)vec是一个包含余弦相似性( DenseVector )结果的列。我要做的是创建一个列"similar_urls“或更新" vec”,并根据vec列值为每一行输入最类似的N个项。
例如,如果我取id = 26,我想在" vec“中查找顶级N项的索引( id和indexes是相同的),并将vec的值替换为顶级N项的url列表。
我想做的是:
我停留在第一步,因为我似乎无法将我的"vec“值转换为一个数组/列表来查找前10个值。
from pyspark.sql.functions import udf
def convert_to_array(vec):
return type(vec)
test_udf = udf(convert_to_array, StringType())
cos_df = cos_df.withColumn("vec", test_udf("vec"))当我试图查看vec值的类型时,它将返回
net.razorvine.pickle.objects.ClassDictConstructor@2db673eb您知道这是什么类型吗?我如何操作它以使我能够转换vec?
记者:我也愿意接受任何其他的解决方案,这将是更好的给出的问题!
发布于 2019-04-24 09:43:12
想出了1的解决方案,看起来很管用。
from pyspark.sql.functions import udf
def convert_to_array(vec):
vec_list = vec.tolist()
sorted_top = sorted(range(len(vec_list)), key=lambda i: vec_list[i], reverse=True)[1:16]
return sorted_top
test_udf = udf(convert_to_array, ArrayType(IntegerType()))
cos_df = cos_df.withColumn("similar_url", test_udf("vec"))https://stackoverflow.com/questions/55825977
复制相似问题