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

用Polars原生命令替代UDF优化用户余弦相似度计算性能

问题背景

免责声明

  • 本问题为某Stack Overflow问题的补充,应用户要求详述场景。
  • 11月29日补充:已测试两种使用explode()的解决方案,但在300万行全量数据集上内存占用激增,仅能在样本数据集验证。

数据集与加载

输入数据集为ml-latest中的ratings.csv(约300万行),包含8万部电影及33万用户的评分,文件大小891MB。通过Polars加载:

movie_ratings = pl.read_csv(os.path.join(application_path + data_directory, "ratings.csv"))

现有实现与瓶颈

目标是计算目标用户与其他所有用户的余弦相似度,需先将评分表转换为单用户一行的结构。现有代码存在两处UDF性能瓶颈:

  1. 生成用户元数据:处理25万行耗时约2分15秒,核心逻辑为分组后通过UDF计算共同电影及数量。
  2. 计算相似度得分:通过3个UDF提取共同电影评分、目标用户对应评分并计算余弦相似度,处理19万用户耗时约27秒。

现有UDF函数:

def get_common_movie_ratings(row) -> pl.List(pl.Float64):
    common_movies = row['common_movies']
    user_ratings = row['user_ratings']
    ratings_for_common_movies = [user_ratings[list(row['user_movies']).index(movie)] for movie in common_movies]
    return ratings_for_common_movies

def get_target_movie_ratings(row, target_user_movies:np.ndarray, target_user_ratings:np.ndarray) -> pl.List(pl.Float64):
    common_movies = row['common_movies']
    target_user_common_ratings = [target_user_ratings[list(target_user_movies).index(movie)] for movie in common_movies]
    return target_user_common_ratings

def compute_cosine(row)->pl.Float64:
    array1 = row["common_movie_ratings"]
    array2 = row["target_user_common_movie_ratings"]
    magnitude1 = norm(array1)
    magnitude2 = norm(array2)
    if magnitude1 != 0 or magnitude2 != 0: #avoid division with 0 norms/magnitudes
        score: float = np.dot(array1, array2) / (norm(array1) * norm(array2))
    else:
        score: float = 0.0
    return score

性能基准与核心问题

  • 单个用户计算耗时约4分钟,遍历33万用户总耗时极高
  • 计算时内存占用3-5GB
  • 生成user_metadata为主要性能瓶颈

核心需求:将上述3个UDF函数转换为Polars原生命令,优化计算性能。


解决方案:Polars原生指令替代UDF

步骤1:预处理用户评分数据(替代用户元数据生成UDF)

通过Polars原生分组聚合生成用户基础数据,同时预计算评分模长以避免重复计算:

import polars as pl
import numpy as np
from numpy.linalg import norm

# 分组聚合生成每个用户的电影列表、评分列表,以及评分模长
user_profiles = movie_ratings.group_by("userId").agg(
    pl.col("movieId").alias("user_movies"),
    pl.col("rating").alias("user_ratings"),
    pl.col("rating").map_elements(lambda x: norm(x), return_dtype=pl.Float64).alias("user_norm")
)

步骤2:提取目标用户数据

指定目标用户ID,提取其电影列表、评分列表和模长:

target_user_id = 12345  # 替换为实际目标用户ID

# 提取目标用户的核心数据
target_data = user_profiles.filter(pl.col("userId") == target_user_id).select(
    pl.col("user_movies").alias("target_movies"),
    pl.col("user_ratings").alias("target_ratings"),
    pl.col("user_norm").alias("target_norm")
).row(0)

target_movies, target_ratings, target_norm = target_data

步骤3:原生提取共同电影及对应评分(替代两个评分提取UDF)

利用Polars数组原生方法获取共同电影,并直接提取对应评分,全程无需Python循环:

# 计算共同电影,并提取双方对应评分
user_with_common = user_profiles.with_columns(
    # 获取当前用户与目标用户的共同电影
    pl.col("user_movies").array.intersection(target_movies).alias("common_movies"),
    # 原生提取当前用户共同电影的评分
    pl.col("user_ratings").array.eval(
        pl.element().take(pl.col("user_movies").array.index(pl.col("common_movies")))
    ).alias("common_movie_ratings"),
    # 原生提取目标用户共同电影的评分
    pl.lit(target_ratings).array.eval(
        pl.element().take(pl.lit(target_movies).array.index(pl.col("common_movies")))
    ).alias("target_user_common_movie_ratings")
)

步骤4:原生计算余弦相似度(替代compute_cosine UDF)

通过Polars原生数组点积和条件判断计算相似度,避免UDF的性能损耗:

# 计算余弦相似度
similarity_scores = user_with_common.with_columns(
    # 计算共同评分的点积
    pl.col("common_movie_ratings").array.dot(pl.col("target_user_common_movie_ratings")).alias("dot_product"),
    # 计算模长乘积(分母)
    (pl.col("user_norm") * target_norm).alias("norm_product")
).with_columns(
    # 处理模长为0的情况,计算最终相似度
    pl.when(pl.col("norm_product") != 0)
    .then(pl.col("dot_product") / pl.col("norm_product"))
    .otherwise(0.0)
    .alias("cosine_similarity")
).select("userId", "cosine_similarity", "common_movies")

优化说明

  • 内存优化:全程使用数组操作,无需explode展开数据,大幅降低内存占用
  • 性能提升:所有操作均为Polars原生向量运算,比Python UDF/循环快数倍
  • 减少重复计算:预计算用户评分模长,避免多次调用norm函数

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 03:17:10