PySpark ML ALSModel实现predict及recommendProducts方法咨询
PySpark ML模块ALSModel自定义predict与recommendProducts实现方案
ML包的ALS训练得到的ALSModel确实没有保留旧MLLib包中面向单条查询的接口,所有内置接口都是面向分布式DataFrame批量计算设计的,以下是完全对齐旧MLLib接口行为的可复用实现:
单条(user, product)评分预测predict方法
预测逻辑本质是用户隐向量与物品隐向量的点积计算,直接从模型存储的隐因子表提取对应向量计算即可,同时兼容冷启动场景逻辑:
from pyspark.sql import functions as F import numpy as np def predict(als_model, user_id, product_id): """ 单个用户-物品对的评分预测,对齐旧mllib ALS的predict接口行为 :param als_model: 训练好的pyspark.ml.recommendation.ALSModel实例 :param user_id: 待预测的用户ID,类型需与训练时用户ID类型一致 :param product_id: 待预测的物品ID,类型需与训练时物品ID类型一致 :return: 预测评分,冷启动场景下根据模型coldStartStrategy配置返回NaN或0 """ # 提取模型存储的用户、物品隐因子表 user_factor_df = als_model.userFactors item_factor_df = als_model.itemFactors # 过滤查询目标的隐向量 target_user = user_factor_df.filter(F.col("id") == user_id).select("features").first() target_item = item_factor_df.filter(F.col("id") == product_id).select("features").first() # 冷启动分支:用户/物品未出现在训练集中 if not target_user or not target_item: return np.nan if als_model.getColdStartStrategy() == "nan" else 0.0 # 向量点积得到预测评分 user_vec = np.array(target_user.features) item_vec = np.array(target_item.features) pred_score = float(np.dot(user_vec, item_vec)) # 隐式反馈场景下打开下方注释,将点积结果映射为0-1的偏好概率 # if als_model.getImplicitPrefs(): # pred_score = 1 / (1 + np.exp(-pred_score)) return pred_score
单用户TopN推荐recommendProducts方法
不要手动遍历全量物品计算评分做排序,直接复用模型内置的分布式优化推荐接口做单用户封装即可,性能比手动实现高2~3个数量级:
def recommendProducts(als_model, user_id, num=10): """ 给指定用户生成TopN推荐结果,对齐旧mllib ALS的recommendProducts接口行为 :param als_model: 训练好的pyspark.ml.recommendation.ALSModel实例 :param user_id: 待推荐的用户ID :param num: 返回的推荐结果条数,默认10条 :return: 列表,每个元素为(物品ID, 预测评分)元组,按评分降序排列 """ # 获取当前模型绑定的SparkSession spark = als_model.userFactors.sparkSession # 构造单用户输入数据集 target_user_df = spark.createDataFrame([(user_id,)], [als_model.getUserCol()]) # 调用模型内置的子集推荐接口,仅计算目标用户的推荐结果 rec_result_df = als_model.recommendForUserSubset(target_user_df, num) # 冷启动分支:用户不存在时返回空列表 if rec_result_df.isEmpty(): return [] # 解析结果转换为旧接口的输出格式 rec_info = rec_result_df.first().recommendations return [(row[als_model.getItemCol()], float(row.rating)) for row in rec_info]
使用注意事项
- 上述两个方法仅适合单条查询的小规模场景,批量预测/批量推荐请直接使用模型自带的
transform、recommendForAllUsers/recommendForUserSubset接口,避免逐行循环调用产生大量调度开销 - 如果你的用户/物品ID是字符串等非数值类型,不需要修改函数逻辑,保证传入的ID类型和训练时一致即可
- 若训练时设置了非默认的用户列、物品列名,函数内部会自动读取模型配置的列名,不需要手动调整
内容的提问来源于stack exchange,提问作者sun tony
相关产品推荐
相关产品推荐

