You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PySpark中使用Pandas UDF计算向量余弦相似度报错的解决方法

Pandas UDF实现向量余弦相似度的正确方式

问题原因

你之前的Pandas UDF报错是因为原cosine_similarity函数是针对单个向量对设计的,但Pandas UDF接收的是pd.Series对象(每个元素是一个向量列表)。直接将整个Series传入原函数会让numpy把Series当作单一向量处理,引发TypeError: only size-1 arrays can be converted to Python scalars错误。

正确实现方式

方法1:逐元素调用原函数(简单直观)

保留原有的单向量计算逻辑,通过遍历Series中的每一对向量完成计算:

import pandas as pd
import numpy as np
from pyspark.sql.functions import pandas_udf
from pyspark.sql.types import FloatType

# 原有单向量对余弦相似度计算函数
def cosine_similarity(vec1, vec2):
    return float(np.dot(vec1, vec2) / (np.linalg.norm(vec1) * np.linalg.norm(vec2)))

# Pandas UDF实现
@pandas_udf(returnType=FloatType())
def cosine_similarity_udf(vec1: pd.Series, vec2: pd.Series) -> pd.Series:
    # 遍历每一对向量,应用计算函数
    return pd.Series([cosine_similarity(v1, v2) for v1, v2 in zip(vec1, vec2)])

方法2:向量化计算(性能更优)

利用numpy的向量化操作批量处理,避免Python循环,适合大数据量场景:

import pandas as pd
import numpy as np
from pyspark.sql.functions import pandas_udf
from pyspark.sql.types import FloatType

@pandas_udf(returnType=FloatType())
def cosine_similarity_udf(vec1: pd.Series, vec2: pd.Series) -> pd.Series:
    # 将Series中的向量列表转换为二维numpy数组
    arr1 = np.array(vec1.tolist())
    arr2 = np.array(vec2.tolist())
    
    # 批量计算点积
    dot_product = (arr1 * arr2).sum(axis=1)
    # 批量计算范数乘积
    norm_product = np.linalg.norm(arr1, axis=1) * np.linalg.norm(arr2, axis=1)
    
    # 计算余弦相似度,添加极小值避免除以0报错
    similarity = dot_product / np.maximum(norm_product, 1e-10)
    
    return pd.Series(similarity.astype(float))

测试验证

使用你提供的测试DataFrame验证功能:

# 创建测试数据
df = spark.createDataFrame([("A", [1, 2, 3], [3, 4, 5]), ("B", [5, 6, 7], [7, 8, 9] )], ("name", "vec1", "vec2"))

# 应用UDF并查看结果
df.withColumn("cosine_similarity", cosine_similarity_udf("vec1", "vec2")).show()

运行后会输出包含余弦相似度列的结果,和原Spark UDF的计算结果一致。

内容的提问来源于stack exchange,提问作者Haritha Thilakarathne

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.23 09:22:21