首页
学习
活动
专区
圈层
工具
发布
社区首页 >问答首页 >如何根据余弦相似度得到最相似的N个项目?

如何根据余弦相似度得到最相似的N个项目?
EN

Stack Overflow用户
提问于 2019-04-24 08:49:03
回答 1查看 1.7K关注 0票数 0

我有一个图像数据集( id,url,功能),我在这些数据集上执行了所有图像之间的余弦相似性。其结果是一个具有以下结构的:

代码语言:javascript
复制
>>> 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列表。

我想做的是:

  1. 将"vec“替换为顶部N个最相似项的索引(udf)的列表/数组
  2. 将该列表/数组替换为urls的列表/数组(udf)

我停留在第一步,因为我似乎无法将我的"vec“值转换为一个数组/列表来查找前10个值。

代码语言:javascript
复制
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值的类型时,它将返回

代码语言:javascript
复制
net.razorvine.pickle.objects.ClassDictConstructor@2db673eb

您知道这是什么类型吗?我如何操作它以使我能够转换vec?

记者:我也愿意接受任何其他的解决方案,这将是更好的给出的问题!

EN

回答 1

Stack Overflow用户

发布于 2019-04-24 09:43:12

想出了1的解决方案,看起来很管用。

代码语言:javascript
复制
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"))
票数 0
EN
页面原文内容由Stack Overflow提供。腾讯云小微IT领域专用引擎提供翻译支持
原文链接:

https://stackoverflow.com/questions/55825977

复制
相关文章

相似问题

领券
问题归档专栏文章快讯文章归档关键词归档开发者手册归档开发者手册 Section 归档