Spark MLlib 实战:构建商品推荐系统(含协同过滤)
·
Spark MLlib实战:商品推荐系统(协同过滤)
核心原理
协同过滤基于"相似用户偏好相似商品"的假设:
- 用户-商品矩阵:构建$m \times n$矩阵$R$,$R_{ij}$表示用户$i$对商品$j$的评分
- 矩阵分解:将$R$分解为低秩矩阵: $$R \approx U \cdot V^T$$ 其中$U \in \mathbb{R}^{m \times k}$为用户隐因子矩阵,$V \in \mathbb{R}^{n \times k}$为商品隐因子矩阵
- 交替最小二乘法(ALS):最小化损失函数: $$\min_{U,V} \sum_{(i,j) \in \Omega} (r_{ij} - u_i^T v_j)^2 + \lambda (|U|_F^2 + |V|_F^2)$$
实现步骤
from pyspark.sql import SparkSession
from pyspark.ml.recommendation import ALS
from pyspark.ml.evaluation import RegressionEvaluator
# 初始化Spark
spark = SparkSession.builder.appName("RecommendationDemo").getOrCreate()
# 模拟数据 (用户ID, 商品ID, 评分)
data = [(0, 10, 4.0), (0, 20, 2.0), (1, 10, 3.0),
(1, 30, 5.0), (2, 20, 1.0), (2, 30, 4.0)]
columns = ["user_id", "item_id", "rating"]
df = spark.createDataFrame(data, columns)
# 划分训练集/测试集
train, test = df.randomSplit([0.8, 0.2])
# 构建ALS模型
als = ALS(
maxIter=10, # 迭代次数
regParam=0.1, # 正则化参数
rank=5, # 隐因子数量
userCol="user_id",
itemCol="item_id",
ratingCol="rating",
coldStartStrategy="drop" # 冷启动处理
)
# 训练模型
model = als.fit(train)
# 预测测试集
predictions = model.transform(test)
# 评估模型 (RMSE)
evaluator = RegressionEvaluator(
metricName="rmse",
labelCol="rating",
predictionCol="prediction"
)
rmse = evaluator.evaluate(predictions)
print(f"模型RMSE: {rmse:.4f}")
# 生成推荐
# 为所有用户推荐TOP3商品
user_recs = model.recommendForAllUsers(3)
# 为指定用户推荐
single_user_recs = model.recommendForUserSubset(
spark.createDataFrame([(0,)], ["user_id"]), 5
)
# 显示结果
print("\n全局推荐结果:")
user_recs.show(truncate=False)
print("\n用户0的推荐:")
single_user_recs.show(truncate=False)
关键参数说明
| 参数 | 说明 | 典型值 |
|---|---|---|
rank | 隐特征维度 | 10-200 |
maxIter | 最大迭代次数 | 10-20 |
regParam | 正则化系数 | 0.01-0.1 |
alpha | 隐式反馈置信度 | 1.0-40.0 |
nonnegative | 非负约束 | True/False |
性能优化技巧
-
数据预处理:
# 处理缺失值 df = df.na.fill(0) # 标准化评分 from pyspark.ml.feature import MinMaxScaler scaler = MinMaxScaler(inputCol="rating", outputCol="scaled_rating") -
冷启动解决方案:
- 混合推荐:新用户使用基于内容的推荐
- 默认策略:推荐热门商品
from pyspark.sql.functions import count popular_items = df.groupBy("item_id").agg(count("rating").alias("count")) -
增量更新:
# 使用checkpoint定期更新模型 als.setCheckpointInterval(5) # 增量训练 updated_model = als.fit(new_data, initialModel=model)
应用场景扩展
- 隐式反馈(适用于点击数据):
als.setImplicitPrefs(True) # 启用隐式反馈 - 跨域推荐:整合用户在不同平台的行为数据
- 实时推荐:结合Spark Streaming实现实时更新
注意:实际生产环境需处理千万级用户数据时,建议:
- 使用
Parquet格式存储数据- 调整
spark.executor.memory和spark.driver.memory- 对用户/商品ID进行分桶处理
更多推荐
所有评论(0)