Spark MLlib实战:商品推荐系统(协同过滤)

核心原理

协同过滤基于"相似用户偏好相似商品"的假设:

  1. 用户-商品矩阵:构建$m \times n$矩阵$R$,$R_{ij}$表示用户$i$对商品$j$的评分
  2. 矩阵分解:将$R$分解为低秩矩阵: $$R \approx U \cdot V^T$$ 其中$U \in \mathbb{R}^{m \times k}$为用户隐因子矩阵,$V \in \mathbb{R}^{n \times k}$为商品隐因子矩阵
  3. 交替最小二乘法(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
性能优化技巧
  1. 数据预处理

    # 处理缺失值
    df = df.na.fill(0) 
    
    # 标准化评分
    from pyspark.ml.feature import MinMaxScaler
    scaler = MinMaxScaler(inputCol="rating", outputCol="scaled_rating")
    

  2. 冷启动解决方案

    • 混合推荐:新用户使用基于内容的推荐
    • 默认策略:推荐热门商品
    from pyspark.sql.functions import count
    popular_items = df.groupBy("item_id").agg(count("rating").alias("count"))
    

  3. 增量更新

    # 使用checkpoint定期更新模型
    als.setCheckpointInterval(5) 
    # 增量训练
    updated_model = als.fit(new_data, initialModel=model)
    

应用场景扩展
  1. 隐式反馈(适用于点击数据):
    als.setImplicitPrefs(True)  # 启用隐式反馈
    

  2. 跨域推荐:整合用户在不同平台的行为数据
  3. 实时推荐:结合Spark Streaming实现实时更新

注意:实际生产环境需处理千万级用户数据时,建议:

  1. 使用Parquet格式存储数据
  2. 调整spark.executor.memoryspark.driver.memory
  3. 对用户/商品ID进行分桶处理

更多推荐