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

PySpark结合pandas_on_spark实现Cosine Similarity匹配函数的优化问题

PySpark分布式实现get_matches_df函数解决方案

问题背景

我正在尝试基于PySpark、pandas_on_spark(即koalas)实现名称匹配的余弦相似度方法中的get_matches_df函数,核心优化需求是避免调用DataFrame的toPandas()方法防止驱动节点过载,实现可扩展的分布式运行逻辑,支持批量处理、pandas_udf或接收单个向量加两个DataFrame的普通UDF实现方式,参考实现示例如下:

>>> psdf = ps.DataFrame({'a': [1,2,3], 'b':[4,5,6]})
>>> def pandas_plus(pdf):
...     return pdf[pdf.a > 1]  # allow arbitrary length
...
>>> psdf.pandas_on_spark.apply_batch(pandas_plus)

目前已完成自定义tfidfvectorizer、cosine值缩放、PySpark稀疏矩阵生成器等模块开发,仅get_matches_df函数待优化。该函数当前的本地实现逻辑如下,可接受驱动端加载全量数据的运行模式,优先需求为分布式可扩展方案:

def get_matches_df(sparse_matrix, name_vector, top=100):
    non_zeros = sparse_matrix.nonzero()
    
    sparserows = non_zeros[0]
    sparsecols = non_zeros[1]
    
    if top:
        nr_matches = top
    else:
        nr_matches = sparsecols.size
    
    left_side = np.empty([nr_matches], dtype=object)
    right_side = np.empty([nr_matches], dtype=object)
    similairity = np.zeros(nr_matches)
    
    for index in range(0, nr_matches):
        left_side[index] = name_vector[sparserows[index]]
        right_side[index] = name_vector[sparsecols[index]]
        similairity[index] = sparse_matrix.data[index]
    
    return pd.DataFrame({'left_side': left_side,
                          'right_side': right_side,
                           'similairity': similairity})

注:原本地实现中similairity为拼写错误,以下优化方案统一修正为正确拼写similarity

优化方案

方案1:纯PySpark原生分布式实现(无驱动端压力)

核心思路是将稀疏矩阵非零元素、名称向量全部转为Spark DataFrame,通过分布式关联替代本地循环,全程不在驱动端加载全量数据:

import pyspark.sql.functions as F
from pyspark.sql.types import StructType, StructField, IntegerType, FloatType, StringType

def distributed_get_matches_df(sparse_matrix, name_vector, top=100):
    # 1. 提取稀疏矩阵非零三元组 (行索引, 列索引, 相似度)
    non_zeros = sparse_matrix.nonzero()
    rows = non_zeros[0]
    cols = non_zeros[1]
    scores = sparse_matrix.data

    # 2. 生成匹配对Spark DataFrame
    match_schema = StructType([
        StructField("row_idx", IntegerType(), nullable=False),
        StructField("col_idx", IntegerType(), nullable=False),
        StructField("similarity", FloatType(), nullable=False)
    ])
    match_df = spark.createDataFrame(zip(rows, cols, scores), schema=match_schema)
    
    # 3. 按相似度降序取TopN(如果指定top参数)
    if top:
        match_df = match_df.orderBy(F.col("similarity").desc()).limit(top)
    
    # 4. 生成名称向量Spark DataFrame
    name_schema = StructType([
        StructField("idx", IntegerType(), nullable=False),
        StructField("name", StringType(), nullable=False)
    ])
    name_df = spark.createDataFrame(enumerate(name_vector), schema=name_schema)
    
    # 5. 两次关联拿到左右两侧名称
    result_df = match_df.join(
        name_df, match_df.row_idx == name_df.idx, "left"
    ).withColumnRenamed("name", "left_side").drop("idx")\
    .join(
        name_df, match_df.col_idx == name_df.idx, "left"
    ).withColumnRenamed("name", "right_side").drop("idx", "row_idx", "col_idx")
    
    return result_df

方案2:pandas_on_spark 批量处理实现(适配示例apply_batch写法)

如果偏好pandas风格的API,可使用pandas_on_spark的批量处理能力,逻辑和原有代码几乎一致,自动分布式运行:

import pandas as pd
import numpy as np
import pyspark.pandas as ps

def pandas_on_spark_get_matches_df(sparse_matrix, name_vector, top=100):
    # 提取稀疏矩阵非零三元组
    non_zeros = sparse_matrix.nonzero()
    rows = non_zeros[0]
    cols = non_zeros[1]
    scores = sparse_matrix.data
    
    # 转为pandas_on_spark DataFrame
    match_psdf = ps.DataFrame({
        "row_idx": rows,
        "col_idx": cols,
        "similarity": scores
    })
    
    # 定义批量处理函数
    def process_batch(pdf: pd.DataFrame) -> pd.DataFrame:
        # 名称向量提前广播,批量内直接索引
        name_arr = np.array(name_vector)
        pdf["left_side"] = pdf["row_idx"].apply(lambda x: name_arr[x])
        pdf["right_side"] = pdf["col_idx"].apply(lambda x: name_arr[x])
        return pdf[["left_side", "right_side", "similarity"]]
    
    # 取TopN后批量处理
    if top:
        match_psdf = match_psdf.sort_values("similarity", ascending=False).head(top)
    result_psdf = match_psdf.pandas_on_spark.apply_batch(process_batch)
    
    return result_psdf

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.23 22:45:07