商品推荐系统#
实验目的#
本实验旨在通过Spark MLlib实现电商平台的企业级商品推荐系统,掌握以下核心技能:
- 理解推荐系统的核心算法原理(协同过滤、基于内容、混合推荐)
- 实现基于ALS的协同过滤推荐算法
- 实现基于TF-IDF的内容推荐算法
- 构建加权混合推荐系统
- 掌握推荐系统的多维度评估方法
- 实现推荐系统的生产级部署方案
知识点1:推荐系统概述#
1.1 什么是推荐系统#
推荐系统(Recommendation System)是信息过滤系统的一个子类,用于预测用户对物品的"评分"或"偏好"。在电商场景中,推荐系统帮助用户从海量商品中发现感兴趣的内容,同时帮助平台提升转化率和用户黏性。
推荐系统的核心目标:
| 目标 | 说明 | 衡量指标 |
|---|---|---|
| 准确性 | 推荐的物品符合用户兴趣 | RMSE、Precision、Recall |
| 多样性 | 推荐列表覆盖不同类别 | 覆盖率、熵 |
| 新颖性 | 推荐用户未发现但感兴趣的物品 | 长尾覆盖率 |
| 实时性 | 及时响应用户行为变化 | 延迟时间 |
1.2 推荐系统分类#
推荐系统
├── 基于人口的统计学推荐(Demographic-based)
│ └── 根据用户基本信息(年龄、性别、地区)推荐
├── 协同过滤推荐(Collaborative Filtering)
│ ├── 基于用户的协同过滤(User-CF)
│ ├── 基于物品的协同过滤(Item-CF)
│ └── 基于模型的协同过滤(Model-based CF,如ALS)
├── 基于内容的推荐(Content-based)
│ ├── 基于物品属性(TF-IDF、Word2Vec)
│ └── 基于用户画像
└── 混合推荐(Hybrid)
├── 加权混合
├── 切换混合
├── 特征组合
└── 级联混合1.3 电商推荐场景#
| 场景 | 推荐策略 | 典型位置 |
|---|---|---|
| 首页推荐 | 热门+个性化混合 | 首页商品流 |
| 猜你喜欢 | 协同过滤+内容 | 个人中心 |
| 相似商品 | 基于内容 | 商品详情页 |
| 购后推荐 | 关联规则+协同过滤 | 订单完成页 |
| 搜索推荐 | 搜索词+用户画像 | 搜索结果页 |
知识点2:协同过滤与ALS算法原理#
2.1 协同过滤的核心思想#
协同过滤的核心假设:如果两个用户在过去对某些物品有相似的偏好,那么他们在未来对其他物品也会有相似的偏好。
协同过滤分为两大类:
基于用户的协同过滤(User-CF):
用户A → 喜欢商品 {1, 2, 3}
用户B → 喜欢商品 {1, 2, 4}
用户A和用户B相似 → 将商品4推荐给用户A基于物品的协同过滤(Item-CF):
喜欢商品1的用户 → 大多也喜欢商品2
商品1和商品2相似 → 给喜欢商品1的用户推荐商品22.2 ALS算法详解#
ALS(Alternating Least Squares,交替最小二乘法)是Spark MLlib中实现的基于模型的协同过滤算法,属于矩阵分解方法。
矩阵分解原理:
用户-物品评分矩阵R(m×n)可以分解为两个低秩矩阵的乘积:
R ≈ U × V^T
其中:
R:m×n 的评分矩阵(m个用户,n个物品)
U:m×k 的用户特征矩阵
V:n×k 的物品特征矩阵
k:隐含特征维度(rank参数)ALS优化过程:
1. 固定V,求解U:对每个用户u,最小化 Σ(r_ui - u_i · v_i)² + λ||u_i||²
2. 固定U,求解V:对每个物品i,最小化 Σ(r_ui - u_i · v_i)² + λ||v_i||²
3. 交替执行步骤1和2,直到收敛关键参数说明:
| 参数 | 含义 | 推荐范围 | 说明 |
|---|---|---|---|
| rank | 隐含特征维度 | 10~200 | 越大模型越复杂,容易过拟合 |
| maxIter | 最大迭代次数 | 5~20 | 过多迭代可能过拟合 |
| regParam | 正则化参数 | 0.01~0.1 | 防止过拟合,越大约束越强 |
| alpha | 隐式反馈置信度 | 1~40 | 仅隐式反馈时使用 |
| coldStartStrategy | 冷启动策略 | “drop”/“nan” | 新用户/物品的处理方式 |
显式反馈 vs 隐式反馈:
| 类型 | 数据来源 | ALS方法 | 示例 |
|---|---|---|---|
| 显式反馈 | 用户主动评分 | ALS | 1~5星评分 |
| 隐式反馈 | 用户行为推断 | ALS.trainImplicit | 浏览、点击、购买 |
2.3 冷启动问题#
冷启动是推荐系统面临的核心挑战之一:
| 冷启动类型 | 原因 | 解决方案 |
|---|---|---|
| 用户冷启动 | 新用户无历史行为 | 基于人口统计学推荐、热门推荐 |
| 物品冷启动 | 新商品无被交互记录 | 基于内容推荐、相似物品推荐 |
| 系统冷启动 | 系统刚上线无数据 | 基于规则推荐、迁移学习 |
知识点3:基于内容的推荐原理#
3.1 TF-IDF算法#
TF-IDF(Term Frequency-Inverse Document Frequency)用于衡量一个词对文档的重要程度:
TF(t, d) = 词t在文档d中出现的次数 / 文档d的总词数
IDF(t) = log(文档总数 / 包含词t的文档数)
TF-IDF(t, d) = TF(t, d) × IDF(t)TF-IDF在推荐系统中的应用:
- 将商品描述文本转化为TF-IDF特征向量
- 计算商品之间的余弦相似度
- 为用户推荐与其历史偏好商品相似的商品
3.2 余弦相似度#
余弦相似度衡量两个向量的夹角余弦值,范围[-1, 1],值越大越相似:
cos(A, B) = (A · B) / (||A|| × ||B||)
其中:
A · B = Σ(a_i × b_i) 向量点积
||A|| = √(Σa_i²) 向量模长余弦相似度的优势:不受向量绝对大小影响,只关注方向,适合文本特征比较。
技术栈#
| 技术 | 版本 | 用途 |
|---|---|---|
| Apache Spark | 3.5.8 | 分布式数据处理与机器学习 |
| Scala | 2.13.8 | Spark应用开发语言 |
| Maven | 3.9.6 | 项目构建与依赖管理 |
| Spark MLlib | 3.5.8 | 机器学习算法库 |
| MySQL | 8.0+ | 结构化数据存储 |
| Redis | 7.0+ | 推荐结果缓存 |
| HDFS | 3.3.6 | 模型文件存储 |
实验环境搭建#
1. 数据库准备#
创建推荐系统所需的数据库表:
CREATE DATABASE IF NOT EXISTS recommendation_system
DEFAULT CHARACTER SET utf8mb4
DEFAULT COLLATE utf8mb4_unicode_ci;
USE recommendation_system;
-- 用户表
DROP TABLE IF EXISTS users;
CREATE TABLE users (
user_id INT PRIMARY KEY COMMENT '用户ID',
username VARCHAR(50) NOT NULL COMMENT '用户名',
age INT COMMENT '年龄',
gender VARCHAR(10) COMMENT '性别',
region VARCHAR(50) COMMENT '地区',
registration_date DATE COMMENT '注册日期',
INDEX idx_region (region),
INDEX idx_age (age)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='用户信息表';
-- 商品表
DROP TABLE IF EXISTS products;
CREATE TABLE products (
product_id INT PRIMARY KEY COMMENT '商品ID',
product_name VARCHAR(100) NOT NULL COMMENT '商品名称',
category VARCHAR(50) NOT NULL COMMENT '分类',
price DECIMAL(10,2) NOT NULL COMMENT '价格',
description TEXT COMMENT '商品描述',
brand VARCHAR(50) COMMENT '品牌',
status TINYINT(1) DEFAULT 1 COMMENT '状态(1上架/0下架)',
create_time TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
INDEX idx_category (category),
INDEX idx_brand (brand),
INDEX idx_status (status)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='商品信息表';
-- 评分表
DROP TABLE IF EXISTS ratings;
CREATE TABLE ratings (
user_id INT NOT NULL COMMENT '用户ID',
product_id INT NOT NULL COMMENT '商品ID',
rating DECIMAL(2,1) NOT NULL COMMENT '评分(1.0-5.0)',
rating_time TIMESTAMP NOT NULL COMMENT '评分时间',
PRIMARY KEY (user_id, product_id),
INDEX idx_user_id (user_id),
INDEX idx_product_id (product_id),
INDEX idx_rating_time (rating_time),
CONSTRAINT fk_rating_user FOREIGN KEY (user_id) REFERENCES users(user_id),
CONSTRAINT fk_rating_product FOREIGN KEY (product_id) REFERENCES products(product_id),
CONSTRAINT chk_rating CHECK (rating >= 1.0 AND rating <= 5.0)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='用户评分表';
-- 用户行为表
DROP TABLE IF EXISTS user_behavior;
CREATE TABLE user_behavior (
id BIGINT AUTO_INCREMENT PRIMARY KEY,
user_id INT NOT NULL COMMENT '用户ID',
product_id INT NOT NULL COMMENT '商品ID',
behavior_type VARCHAR(20) NOT NULL COMMENT '行为类型(view/click/add_to_cart/purchase/favorite)',
behavior_time TIMESTAMP NOT NULL COMMENT '行为时间',
duration INT DEFAULT 0 COMMENT '停留时长(秒)',
INDEX idx_user_id (user_id),
INDEX idx_product_id (product_id),
INDEX idx_behavior_type (behavior_type),
INDEX idx_behavior_time (behavior_time),
CONSTRAINT fk_behavior_user FOREIGN KEY (user_id) REFERENCES users(user_id),
CONSTRAINT fk_behavior_product FOREIGN KEY (product_id) REFERENCES products(product_id)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='用户行为表';
-- 推荐结果缓存表
DROP TABLE IF EXISTS recommendation_cache;
CREATE TABLE recommendation_cache (
user_id INT NOT NULL COMMENT '用户ID',
rec_type VARCHAR(20) NOT NULL COMMENT '推荐类型(cf/cb/hybrid)',
product_ids VARCHAR(500) NOT NULL COMMENT '推荐商品ID列表(逗号分隔)',
scores VARCHAR(500) COMMENT '推荐分数列表(逗号分隔)',
update_time TIMESTAMP DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP,
expire_time TIMESTAMP NOT NULL COMMENT '过期时间',
PRIMARY KEY (user_id, rec_type),
INDEX idx_expire_time (expire_time)
) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COMMENT='推荐结果缓存表';2. 数据准备#
插入测试数据:
INSERT INTO users (user_id, username, age, gender, region, registration_date) VALUES
(1, '张三', 25, '男', '北京', '2023-01-01'),
(2, '李四', 30, '女', '上海', '2023-01-02'),
(3, '王五', 28, '男', '深圳', '2023-01-03'),
(4, '赵六', 35, '女', '杭州', '2023-01-04'),
(5, '钱七', 22, '男', '广州', '2023-01-05'),
(6, '孙八', 27, '女', '成都', '2023-01-06'),
(7, '周九', 32, '男', '南京', '2023-01-07'),
(8, '吴十', 29, '女', '武汉', '2023-01-08'),
(9, '郑十一', 26, '男', '西安', '2023-01-09'),
(10, '冯十二', 31, '女', '重庆', '2023-01-10');
INSERT INTO products (product_id, product_name, category, price, description, brand) VALUES
(1, 'iPhone 15 Pro', '手机', 8999.00, '苹果旗舰智能手机 A17 Pro芯片 钛金属设计', 'Apple'),
(2, 'MacBook Pro 14', '电脑', 14999.00, '苹果专业笔记本电脑 M3 Pro芯片 Liquid Retina XDR显示屏', 'Apple'),
(3, 'AirPods Pro 2', '耳机', 1899.00, '苹果主动降噪无线耳机 自适应通透模式', 'Apple'),
(4, 'iPad Pro 12.9', '平板', 8999.00, '苹果专业平板电脑 M2芯片 mini-LED显示屏', 'Apple'),
(5, 'Apple Watch Ultra 2', '手表', 5999.00, '苹果极限运动手表 钛金属表壳 双频GPS', 'Apple'),
(6, '华为Mate 60 Pro', '手机', 6999.00, '华为旗舰智能手机 麒麟9000S芯片 卫星通信', '华为'),
(7, '小米14 Pro', '手机', 4999.00, '小米旗舰智能手机 骁龙8 Gen3 徕卡光学', '小米'),
(8, 'ThinkPad X1 Carbon', '电脑', 9999.00, '联想商务笔记本电脑 Intel酷睿i7 14英寸轻薄', '联想'),
(9, '索尼WH-1000XM5', '耳机', 2499.00, '索尼旗舰降噪头戴耳机 30小时续航', 'Sony'),
(10, '三星Galaxy Tab S9', '平板', 5999.00, '三星旗舰平板电脑 骁龙8 Gen2 AMOLED', 'Samsung'),
(11, '华为MatePad Pro', '平板', 4699.00, '华为旗舰平板 鸿蒙系统 OLED屏幕', '华为'),
(12, 'OPPO Find X7', '手机', 3999.00, 'OPPO旗舰智能手机 天玑9300 哈苏影像', 'OPPO'),
(13, '戴尔XPS 15', '电脑', 11999.00, '戴尔高端笔记本电脑 Intel酷睿i7 OLED屏', 'Dell'),
(14, 'Bose QC45', '耳机', 1999.00, 'Bose消噪头戴耳机 24小时续航', 'Bose'),
(15, '华为Watch GT4', '手表', 1488.00, '华为智能手表 心率血氧监测 两周续航', '华为');
INSERT INTO ratings (user_id, product_id, rating, rating_time) VALUES
(1, 1, 5.0, '2023-02-01 10:00:00'),
(1, 2, 5.0, '2023-02-02 11:00:00'),
(1, 3, 4.0, '2023-02-03 12:00:00'),
(1, 5, 4.0, '2023-02-04 13:00:00'),
(2, 1, 4.0, '2023-02-05 14:00:00'),
(2, 4, 5.0, '2023-02-06 15:00:00'),
(2, 5, 4.0, '2023-02-07 16:00:00'),
(2, 6, 3.0, '2023-02-08 17:00:00'),
(3, 2, 4.0, '2023-02-09 18:00:00'),
(3, 3, 5.0, '2023-02-10 19:00:00'),
(3, 6, 5.0, '2023-02-11 20:00:00'),
(3, 9, 4.0, '2023-02-12 21:00:00'),
(4, 4, 4.0, '2023-02-13 22:00:00'),
(4, 7, 5.0, '2023-02-14 23:00:00'),
(4, 8, 4.0, '2023-02-15 00:00:00'),
(4, 12, 3.0, '2023-02-16 01:00:00'),
(5, 5, 5.0, '2023-02-17 02:00:00'),
(5, 9, 4.0, '2023-02-18 03:00:00'),
(5, 10, 5.0, '2023-02-19 04:00:00'),
(5, 14, 4.0, '2023-02-20 05:00:00'),
(6, 6, 5.0, '2023-02-21 06:00:00'),
(6, 11, 4.0, '2023-02-22 07:00:00'),
(6, 15, 4.0, '2023-02-23 08:00:00'),
(7, 7, 5.0, '2023-02-24 09:00:00'),
(7, 12, 4.0, '2023-02-25 10:00:00'),
(7, 8, 3.0, '2023-02-26 11:00:00'),
(8, 1, 4.0, '2023-02-27 12:00:00'),
(8, 3, 5.0, '2023-02-28 13:00:00'),
(8, 13, 4.0, '2023-03-01 14:00:00'),
(9, 2, 5.0, '2023-03-02 15:00:00'),
(9, 8, 4.0, '2023-03-03 16:00:00'),
(9, 13, 5.0, '2023-03-04 17:00:00'),
(10, 6, 4.0, '2023-03-05 18:00:00'),
(10, 10, 3.0, '2023-03-06 19:00:00'),
(10, 11, 5.0, '2023-03-07 20:00:00'),
(10, 15, 4.0, '2023-03-08 21:00:00');
INSERT INTO user_behavior (user_id, product_id, behavior_type, behavior_time, duration) VALUES
(1, 1, 'view', '2023-02-01 09:00:00', 120),
(1, 1, 'add_to_cart', '2023-02-01 09:30:00', 0),
(1, 1, 'purchase', '2023-02-01 10:00:00', 0),
(1, 2, 'view', '2023-02-02 10:30:00', 180),
(1, 2, 'add_to_cart', '2023-02-02 11:00:00', 0),
(1, 2, 'purchase', '2023-02-02 11:30:00', 0),
(2, 1, 'view', '2023-02-04 12:00:00', 90),
(2, 1, 'add_to_cart', '2023-02-04 12:30:00', 0),
(2, 1, 'purchase', '2023-02-04 13:00:00', 0),
(2, 4, 'view', '2023-02-05 13:30:00', 200),
(2, 4, 'add_to_cart', '2023-02-05 14:00:00', 0),
(2, 4, 'purchase', '2023-02-05 14:30:00', 0),
(3, 6, 'view', '2023-02-09 17:00:00', 150),
(3, 6, 'favorite', '2023-02-09 17:30:00', 0),
(3, 6, 'purchase', '2023-02-11 20:00:00', 0),
(4, 7, 'view', '2023-02-14 22:00:00', 300),
(4, 7, 'add_to_cart', '2023-02-14 22:30:00', 0),
(4, 7, 'purchase', '2023-02-14 23:00:00', 0),
(5, 9, 'view', '2023-02-18 02:00:00', 240),
(5, 9, 'add_to_cart', '2023-02-18 02:30:00', 0),
(5, 9, 'purchase', '2023-02-18 03:00:00', 0);3. 项目依赖配置#
创建Maven项目,pom.xml配置如下:
<?xml version="1.0" encoding="UTF-8"?>
<project xmlns="http://maven.apache.org/POM/4.0.0"
xmlns:xsi="http://www.w3.org/2001/XMLSchema-instance"
xsi:schemaLocation="http://maven.apache.org/POM/4.0.0
http://maven.apache.org/xsd/maven-4.0.0.xsd">
<modelVersion>4.0.0</modelVersion>
<groupId>com.spark.tutorial</groupId>
<artifactId>recommendation-system</artifactId>
<version>1.0.0</version>
<packaging>jar</packaging>
<properties>
<scala.version>2.13.8</scala.version>
<spark.version>3.5.8</spark.version>
<maven.compiler.source>11</maven.compiler.source>
<maven.compiler.target>11</maven.compiler.target>
<project.build.sourceEncoding>UTF-8</project.build.sourceEncoding>
</properties>
<dependencies>
<dependency>
<groupId>org.apache.spark</groupId>
<artifactId>spark-core_2.13</artifactId>
<version>${spark.version}</version>
<scope>provided</scope>
</dependency>
<dependency>
<groupId>org.apache.spark</groupId>
<artifactId>spark-sql_2.13</artifactId>
<version>${spark.version}</version>
<scope>provided</scope>
</dependency>
<dependency>
<groupId>org.apache.spark</groupId>
<artifactId>spark-mllib_2.13</artifactId>
<version>${spark.version}</version>
<scope>provided</scope>
</dependency>
<dependency>
<groupId>mysql</groupId>
<artifactId>mysql-connector-java</artifactId>
<version>8.0.33</version>
</dependency>
<dependency>
<groupId>redis.clients</groupId>
<artifactId>jedis</artifactId>
<version>5.0.2</version>
</dependency>
<dependency>
<groupId>com.typesafe</groupId>
<artifactId>config</artifactId>
<version>1.4.3</version>
</dependency>
</dependencies>
<build>
<plugins>
<plugin>
<groupId>net.alchim31.maven</groupId>
<artifactId>scala-maven-plugin</artifactId>
<version>4.8.1</version>
<executions>
<execution>
<goals>
<goal>compile</goal>
<goal>testCompile</goal>
</goals>
</execution>
</executions>
<configuration>
<scalaVersion>${scala.version}</scalaVersion>
</configuration>
</plugin>
<plugin>
<groupId>org.apache.maven.plugins</groupId>
<artifactId>maven-compiler-plugin</artifactId>
<version>3.11.0</version>
<configuration>
<source>11</source>
<target>11</target>
</configuration>
</plugin>
</plugins>
</build>
</project>4. 应用配置文件#
创建src/main/resources/application.conf,将敏感配置外部化:
# 推荐系统配置文件(Docker 环境)
# 使用 Docker Desktop 启动 MySQL 和 Redis 服务
spark {
app-name = "ProductRecommendationSystem"
master = "local[*]"
# 性能优化配置
config {
spark.sql.adaptive.enabled = true
spark.sql.adaptive.coalescePartitions.enabled = true
spark.serializer = "org.apache.spark.serializer.KryoSerializer"
spark.sql.shuffle.partitions = 8
spark.log-level = "WARN"
}
}
# Docker MySQL 配置
mysql {
# Docker 容器映射端口 3307 -> 3306
url = "jdbc:mysql://localhost:3307/recommendation_system?useSSL=false&serverTimezone=Asia/Shanghai&useUnicode=true&characterEncoding=UTF-8&allowPublicKeyRetrieval=true"
user = "root"
password = "root123456"
driver = "com.mysql.cj.jdbc.Driver"
pool-size = 5
}
# Docker Redis 配置
redis {
host = "localhost"
port = 6379
password = "" # Docker Redis 默认无密码
database = 0
timeout = 3000
pool {
max-total = 10
max-idle = 5
min-idle = 2
}
cache {
ttl-hours = 24
}
}
als {
rank = 10
max-iter = 10
reg-param = 0.01
alpha = 1.0
cold-start-strategy = "drop"
implicit-preference = false
training-ratio = 0.8
random-seed = 42
}
recommendation {
top-k = 10
cf-weight = 0.6
cb-weight = 0.4
model-save-path = "datas/models/recommendation/als"
}实验内容#
1. 配置管理工具类#
创建ConfigManager.scala,实现企业级配置管理:
import com.typesafe.config.{Config, ConfigFactory}
import org.apache.spark.sql.SparkSession
import scala.jdk.CollectionConverters.CollectionHasAsScala
/**
* ConfigManager - 统一配置管理器
* 从 application_ex14.conf 读取所有配置参数
*
* @author John
* @since 2026/5/7 15:26
* @version 2.0
*/
object ConfigManager {
private lazy val config: Config = {
try {
ConfigFactory.load("application_ex14")
} catch {
case e: Exception =>
println(s"⚠️ 加载配置文件失败,使用默认配置: ${e.getMessage}")
ConfigFactory.load()
}
}
def getConfig: Config = config
/**
* 获取 Spark 配置参数
*/
def getSparkConfig: Map[String, String] = {
try {
val sparkConfig = config.getConfig("spark.config")
import scala.jdk.CollectionConverters.MapHasAsScala
sparkConfig.entrySet().asScala.map { entry =>
entry.getKey -> entry.getValue.unwrapped().toString
}.toMap
} catch {
case _: Exception =>
println("⚠️ Spark 配置缺失,使用默认值")
Map.empty[String, String]
}
}
/**
* 创建 SparkSession(带完整配置)
*/
def createSparkSession(appName: String): SparkSession = {
val builder = SparkSession.builder()
.appName(appName)
.master(getOrElse("spark.master", "local[*]"))
// 应用自定义配置
getSparkConfig.foreach { case (k, v) => builder.config(k, v) }
val session = builder.getOrCreate()
session.sparkContext.setLogLevel(getOrElse("spark.log-level", "WARN"))
session
}
// ==================== MySQL 配置 ====================
def getJdbcUrl: String = getOrElse("mysql.url", "jdbc:mysql://localhost:3306/recommendation_system?useSSL=false&serverTimezone=Asia/Shanghai")
def getJdbcUser: String = getOrElse("mysql.user", "root")
def getJdbcPassword: String = getOrElse("mysql.password", "")
def getJdbcDriver: String = getOrElse("mysql.driver", "com.mysql.cj.jdbc.Driver")
def getJdbcPoolSize: Int = getOrElse("mysql.pool-size", 5)
// ==================== Redis 配置 ====================
def getRedisHost: String = getOrElse("redis.host", "localhost")
def getRedisPort: Int = getOrElse("redis.port", 6379)
def getRedisPassword: String = {
if (config.hasPath("redis.password") && !config.getIsNull("redis.password")) {
config.getString("redis.password")
} else ""
}
def getRedisDatabase: Int = getOrElse("redis.database", 0)
def getRedisTimeout: Int = getOrElse("redis.timeout", 3000)
def getRedisMaxTotal: Int = getOrElse("redis.pool.max-total", 10)
def getRedisMaxIdle: Int = getOrElse("redis.pool.max-idle", 5)
def getRedisMinIdle: Int = getOrElse("redis.pool.min-idle", 2)
def getCacheTtlHours: Int = getOrElse("redis.cache.ttl-hours", 24)
// ==================== ALS 算法配置 ====================
def getAlsRank: Int = getOrElse("als.rank", 10)
def getAlsMaxIter: Int = getOrElse("als.max-iter", 10)
def getAlsRegParam: Double = getOrElse("als.reg-param", 0.01)
def getAlsAlpha: Double = getOrElse("als.alpha", 1.0)
def getAlsColdStartStrategy: String = getOrElse("als.cold-start-strategy", "drop")
def getAlsImplicitPreference: Boolean = getOrElse("als.implicit-preference", false)
def getTrainingRatio: Double = getOrElse("als.training-ratio", 0.8)
def getRandomSeed: Long = getOrElse("als.random-seed", 42L)
// ==================== 推荐系统配置 ====================
def getTopK: Int = getOrElse("recommendation.top-k", 10)
def getCfWeight: Double = getOrElse("recommendation.cf-weight", 0.6)
def getCbWeight: Double = getOrElse("recommendation.cb-weight", 0.4)
def getModelSavePath: String = getOrElse("recommendation.model-save-path", "datas/models/recommendation/als")
/**
* 打印配置信息(调试用)
*/
def printConfig(): Unit = {
println("=" * 80)
println("📋 推荐系统配置信息")
println("=" * 80)
println(s"MySQL: $getJdbcUrl")
println(s"Redis: $getRedisHost:$getRedisPort (DB: $getRedisDatabase)")
println(s"ALS: rank=$getAlsRank, maxIter=$getAlsMaxIter, regParam=$getAlsRegParam")
println(s"推荐: topK=$getTopK, cfWeight=$getCfWeight, cbWeight=$getCbWeight")
println("=" * 80)
}
/**
* 安全获取配置值(带默认值)
*/
private def getOrElse[T](path: String, default: T): T = {
try {
if (config.hasPath(path)) {
val value = config.getValue(path).unwrapped()
// 根据默认值的类型进行转换
default match {
case _: String => value.toString.asInstanceOf[T]
case _: Int => value.toString.toInt.asInstanceOf[T]
case _: Long => value.toString.toLong.asInstanceOf[T]
case _: Double => value.toString.toDouble.asInstanceOf[T]
case _: Boolean => value.toString.toBoolean.asInstanceOf[T]
case _ => default
}
} else {
default
}
} catch {
case e: Exception =>
println(s"⚠️ 读取配置 $path 失败: ${e.getMessage},使用默认值")
default
}
}
}2. 数据访问层#
创建DataAccessLayer.scala,封装数据库访问逻辑,实现连接复用和数据验证:
import org.apache.spark.sql.{DataFrame, SparkSession}
import org.apache.spark.sql.functions._
import org.apache.spark.sql.types.DoubleType
/**
* DataAccessLayer - 数据访问层
* 封装所有数据库操作,统一管理 JDBC 连接
*
* @author John
* @since 2026/5/7 15:29
* @version 2.0
*/
object DataAccessLayer {
/**
* 创建 JDBC Properties(复用)
*/
private def createJdbcProps(): java.util.Properties = {
val props = new java.util.Properties()
props.setProperty("user", ConfigManager.getJdbcUser)
props.setProperty("password", ConfigManager.getJdbcPassword)
props.setProperty("driver", ConfigManager.getJdbcDriver)
// 性能优化配置
props.setProperty("fetchsize", "1000")
props.setProperty("batchsize", "1000")
props
}
/**
* 加载评分数据(带数据验证)
*/
def loadRatings(spark: SparkSession): DataFrame = {
try {
val jdbcUrl = ConfigManager.getJdbcUrl
val jdbcProps = createJdbcProps()
println("📊 正在加载评分数据...")
val df = spark.read.jdbc(jdbcUrl, "ratings", jdbcProps)
.select(
col("user_id").as("userId"),
col("product_id").as("productId"),
col("rating").cast(DoubleType).as("rating")
)
.filter(col("rating").isNotNull && col("rating") >= 1.0 && col("rating") <= 5.0)
.filter(col("userId").isNotNull && col("productId").isNotNull)
.cache()
val count = df.count()
println(s"✅ 评分数据加载完成: $count 条记录")
df
} catch {
case e: Exception =>
println(s"❌ 加载评分数据失败: ${e.getMessage}")
throw e
}
}
/**
* 加载商品数据(仅激活商品)
*/
def loadProducts(spark: SparkSession): DataFrame = {
try {
val jdbcUrl = ConfigManager.getJdbcUrl
val jdbcProps = createJdbcProps()
println("📦 正在加载商品数据...")
val df = spark.read.jdbc(jdbcUrl, "products", jdbcProps)
.filter(col("status") === 1) // 仅加载激活商品
.cache()
val count = df.count()
println(s"✅ 商品数据加载完成: $count 个商品")
df
} catch {
case e: Exception =>
println(s"❌ 加载商品数据失败: ${e.getMessage}")
throw e
}
}
/**
* 加载用户行为数据
*/
def loadUserBehavior(spark: SparkSession): DataFrame = {
try {
val jdbcUrl = ConfigManager.getJdbcUrl
val jdbcProps = createJdbcProps()
println("👤 正在加载用户行为数据...")
val df = spark.read.jdbc(jdbcUrl, "user_behavior", jdbcProps)
.filter(col("user_id").isNotNull && col("product_id").isNotNull)
.cache()
val count = df.count()
println(s"✅ 用户行为数据加载完成: $count 条记录")
df
} catch {
case e: Exception =>
println(s"❌ 加载用户行为数据失败: ${e.getMessage}")
throw e
}
}
/**
* 保存推荐结果到 MySQL
*/
def saveRecommendationsToMySQL(
spark: SparkSession,
recsDF: DataFrame,
tableName: String,
mode: String = "overwrite"
): Unit = {
try {
val jdbcUrl = ConfigManager.getJdbcUrl
val jdbcProps = createJdbcProps()
println(s"💾 正在保存推荐结果到表: $tableName (mode=$mode)...")
recsDF.write
.mode(mode)
.jdbc(jdbcUrl, tableName, jdbcProps)
println(s"✅ 推荐结果保存成功: ${recsDF.count()} 条记录")
} catch {
case e: Exception =>
println(s"❌ 保存推荐结果失败: ${e.getMessage}")
throw e
}
}
}3. 基于协同过滤的推荐#
创建CollaborativeFiltering.scala,实现企业级协同过滤推荐:
知识点回顾:ALS算法通过矩阵分解将用户-物品评分矩阵分解为用户特征矩阵和物品特征矩阵,交替优化两者直到收敛。Spark MLlib的ALS实现支持显式反馈和隐式反馈两种模式。
import org.apache.spark.ml.recommendation.ALS
import org.apache.spark.ml.evaluation.RegressionEvaluator
import org.apache.spark.sql.{DataFrame, SparkSession}
import org.apache.spark.sql.functions._
/**
* CollaborativeFiltering - 协同过滤推荐系统(ALS算法)
*
* @author John
* @since 2026/5/7 15:30
* @version 2.0
*/
object CollaborativeFiltering {
def main(args: Array[String]): Unit = {
var spark: SparkSession = null
try {
// 打印配置信息
ConfigManager.printConfig()
val _spark = ConfigManager.createSparkSession("CollaborativeFiltering")
import _spark.implicits._
spark = _spark
println("\n" + "=" * 80)
println("🤖 协同过滤推荐系统 (ALS算法)")
println("=" * 80)
// 1. 加载数据
val ratingsDF = DataAccessLayer.loadRatings(spark)
val dataCount = ratingsDF.count()
println(s"\n📊 数据统计:")
println(s" - 总记录数: $dataCount")
println(s" - 用户数: ${ratingsDF.select("userId").distinct().count()}")
println(s" - 商品数: ${ratingsDF.select("productId").distinct().count()}")
if (dataCount == 0) {
println("⚠️ 评分数据为空,请检查数据库连接和表数据")
return
}
// 2. 划分训练集和测试集
val Array(training, test) = ratingsDF.randomSplit(
Array(ConfigManager.getTrainingRatio, 1 - ConfigManager.getTrainingRatio),
seed = ConfigManager.getRandomSeed
)
training.cache()
test.cache()
println(s"\n🔧 数据集划分:")
println(s" - 训练集: ${training.count()} 条 (${ConfigManager.getTrainingRatio * 100}%)")
println(s" - 测试集: ${test.count()} 条 (${(1 - ConfigManager.getTrainingRatio) * 100}%)")
// 3. 训练和评估模型
val (model, rmse) = trainAndEvaluate(training, test)
println(s"\n✅ 模型评估结果:")
println(f" - RMSE: $rmse%.4f")
// 4. 生成推荐
val topK = ConfigManager.getTopK
println(s"\n🎯 生成 Top-$topK 推荐:")
// 为所有用户推荐
val userRecs = model.recommendForAllUsers(topK)
println(s" - 已为 ${userRecs.count()} 个用户生成推荐")
println(s"\n 示例(前5个用户):")
userRecs.show(5, truncate = false)
// 为所有商品推荐相似用户
val productRecs = model.recommendForAllItems(topK)
println(s"\n 为商品推荐用户(前5个商品):")
productRecs.show(5, truncate = false)
// 为特定用户推荐
val userId = 1
val userDF = Seq(userId).toDF("userId")
val userSpecificRecs = model.recommendForUserSubset(userDF, 5)
println(s"\n 为用户 $userId 的个性化推荐:")
userSpecificRecs.show(false)
// 5. 保存模型
val modelPath = ConfigManager.getModelSavePath + "/cf-model"
model.write.overwrite().save(modelPath)
println(s"\n💾 模型已保存至: $modelPath")
println("\n" + "=" * 80)
println("✅ 协同过滤推荐系统运行完成")
println("=" * 80)
} catch {
case e: InterruptedException =>
println(s"\n⚠️ 程序被中断: ${e.getMessage}")
case e: Exception =>
println(s"\n❌ 运行出错: ${e.getClass.getSimpleName} - ${e.getMessage}")
e.printStackTrace()
} finally {
if (spark != null) {
println("\n🛑 正在关闭 SparkSession...")
spark.stop()
println("✅ SparkSession 已关闭")
}
}
}
def trainAndEvaluate(
training: DataFrame,
test: DataFrame): (org.apache.spark.ml.recommendation.ALSModel, Double) = {
val als = new ALS()
.setRank(ConfigManager.getAlsRank)
.setMaxIter(ConfigManager.getAlsMaxIter)
.setRegParam(ConfigManager.getAlsRegParam)
.setAlpha(ConfigManager.getAlsAlpha)
.setImplicitPrefs(ConfigManager.getAlsImplicitPreference)
.setColdStartStrategy(ConfigManager.getAlsColdStartStrategy)
.setUserCol("userId")
.setItemCol("productId")
.setRatingCol("rating")
val startTime = System.currentTimeMillis()
val model = als.fit(training)
val trainingTime = (System.currentTimeMillis() - startTime) / 1000.0
println(s"模型训练完成,耗时: ${trainingTime}s")
println(s"用户因子维度: ${model.rank}, 用户数: ${model.userFactors.count()}, 物品数: ${model.itemFactors.count()}")
val predictions = model.transform(test)
val evaluator = new RegressionEvaluator()
.setMetricName("rmse")
.setLabelCol("rating")
.setPredictionCol("prediction")
val rmse = evaluator.evaluate(predictions)
(model, rmse)
}
}4. 基于内容的推荐#
创建ContentBasedRecommendation.scala,实现基于TF-IDF的内容推荐:
知识点回顾:基于内容的推荐通过分析物品的属性特征(如商品描述文本),计算物品之间的相似度。TF-IDF将文本转化为数值向量,余弦相似度衡量向量间的方向一致性。当用户冷启动时,基于内容的方法仍然可以推荐与用户浏览商品属性相似的商品。
import org.apache.spark.ml.feature.{HashingTF, IDF, Tokenizer}
import org.apache.spark.ml.linalg.SparseVector
import org.apache.spark.sql.{DataFrame, SparkSession}
import org.apache.spark.sql.functions._
/**
* ContentBasedRecommendation - 基于内容的推荐系统(TF-IDF)
*
* @author John
* @since 2026/5/7 15:31
* @version 2.0
*/
object ContentBasedRecommendation {
case class ProductSimilarity(productId: Int, similarProductId: Int, similarity: Double)
def main(args: Array[String]): Unit = {
var spark: SparkSession = null
try {
val _spark = ConfigManager.createSparkSession("ContentBasedRecommendation")
import _spark.implicits._
spark = _spark
println("\n" + "=" * 80)
println("📝 基于内容的推荐系统 (TF-IDF)")
println("=" * 80)
// 1. 加载数据
val productsDF = DataAccessLayer.loadProducts(spark)
val ratingsDF = DataAccessLayer.loadRatings(spark)
if (productsDF.count() == 0) {
println("⚠️ 商品数据为空,无法进行基于内容的推荐")
return
}
// 2. 计算商品特征向量
println("\n🔍 正在计算商品特征向量...")
val productFeatures = computeProductFeatures(productsDF)
println(s"✅ 计算了 ${productFeatures.size} 个商品的特征向量")
// 3. 查找相似商品
val targetProductId = 1
println(s"\n🎯 查找与商品 $targetProductId 相似的商品:")
val similarProducts = findSimilarProducts(targetProductId, productFeatures, 5)
if (similarProducts.nonEmpty) {
similarProducts.foreach { case ProductSimilarity(_, simId, sim) =>
val productNameRow = productsDF.filter($"product_id" === simId)
.select("product_name").take(1).headOption
productNameRow match {
case Some(row) =>
val name = row.getString(0)
println(f" - 商品: $name (ID=$simId), 相似度: $sim%.4f")
case None =>
println(f" - 商品ID=$simId (名称不存在), 相似度: $sim%.4f")
}
}
} else {
println(" ⚠️ 未找到相似商品")
}
// 4. 为用户推荐
val userId = 1
println(s"\n👤 为用户 $userId 生成基于内容的推荐:")
val recommendations = recommendForUser(userId, ratingsDF, productsDF, productFeatures, 5)
if (recommendations.nonEmpty) {
recommendations.foreach { case (productId, score) =>
val productNameRow = productsDF.filter($"product_id" === productId)
.select("product_name").take(1).headOption
productNameRow match {
case Some(row) =>
val name = row.getString(0)
println(f" - 商品: $name (ID=$productId), 推荐分数: $score%.4f")
case None =>
println(f" - 商品ID=$productId (名称不存在), 推荐分数: $score%.4f")
}
}
} else {
println(" ⚠️ 无法生成推荐(用户无历史评分或数据不足)")
}
println("\n" + "=" * 80)
println("✅ 基于内容的推荐系统运行完成")
println("=" * 80)
} catch {
case e: InterruptedException =>
println(s"\n⚠️ 程序被中断: ${e.getMessage}")
case e: Exception =>
println(s"\n❌ 运行出错: ${e.getClass.getSimpleName} - ${e.getMessage}")
e.printStackTrace()
} finally {
if (spark != null) {
println("\n🛑 正在关闭 SparkSession...")
spark.stop()
println("✅ SparkSession 已关闭")
}
}
}
def computeProductFeatures(productsDF: DataFrame): Map[Int, SparseVector] = {
val tokenizer = new Tokenizer()
.setInputCol("description")
.setOutputCol("words")
val wordsData = tokenizer.transform(productsDF.na.fill(Map("description" -> "")))
val hashingTF = new HashingTF()
.setInputCol("words")
.setOutputCol("rawFeatures")
.setNumFeatures(1000)
val featurizedData = hashingTF.transform(wordsData)
val idf = new IDF()
.setInputCol("rawFeatures")
.setOutputCol("features")
val idfModel = idf.fit(featurizedData)
val rescaledData = idfModel.transform(featurizedData)
rescaledData.select("product_id", "features").rdd.map {
case row => (row.getInt(0), row.getAs[SparseVector](1))
}.collect().toMap
}
def cosineSimilarity(vec1: SparseVector, vec2: SparseVector): Double = {
val dotProduct = vec1.dot(vec2)
val norm1 = math.sqrt(vec1.dot(vec1))
val norm2 = math.sqrt(vec2.dot(vec2))
if (norm1 < 1e-10 || norm2 < 1e-10) 0.0 else dotProduct / (norm1 * norm2)
}
def findSimilarProducts(
targetProductId: Int,
productFeatures: Map[Int, SparseVector],
topK: Int): Seq[ProductSimilarity] = {
productFeatures.get(targetProductId) match {
case Some(targetFeatures) =>
productFeatures.toSeq
.filter(_._1 != targetProductId)
.map { case (productId, features) =>
val sim = cosineSimilarity(targetFeatures, features)
ProductSimilarity(targetProductId, productId, sim)
}
.sortBy(-_.similarity)
.take(topK)
case None =>
println(s"商品 $targetProductId 不存在特征向量")
Seq.empty
}
}
def recommendForUser(
userId: Int,
ratingsDF: DataFrame,
productsDF: DataFrame,
productFeatures: Map[Int, SparseVector],
topK: Int): Seq[(Int, Double)] = {
val userRatings = ratingsDF.filter(ratingsDF("userId") === userId).collect()
if (userRatings.isEmpty) {
println(s"用户 $userId 无历史评分,无法进行基于内容的推荐")
return Seq.empty
}
val userPreferences = userRatings.map { row =>
(row.getInt(1), row.getDouble(2))
}
val ratedProductIds = userPreferences.map(_._1).toSet
productFeatures.toSeq
.filter { case (productId, _) => !ratedProductIds.contains(productId) }
.map { case (productId, features) =>
val score = userPreferences.map { case (pid, rating) =>
productFeatures.get(pid) match {
case Some(pf) => cosineSimilarity(features, pf) * rating
case None => 0.0
}
}.sum
(productId, score)
}
.sortBy(-_._2)
.take(topK)
}
}5. 混合推荐系统#
创建HybridRecommendation.scala,实现加权混合推荐:
知识点回顾:混合推荐系统结合多种推荐算法的优势,弥补单一算法的不足。本实验采用加权混合策略,协同过滤和基于内容的推荐结果按权重线性组合。权重可通过A/B测试或交叉验证确定。
混合推荐分数 = cfWeight × 协同过滤分数 + cbWeight × 基于内容分数
其中 cfWeight + cbWeight = 1.0import org.apache.spark.ml.recommendation.ALS
import org.apache.spark.sql.{DataFrame, SparkSession}
import org.apache.spark.sql.functions._
/**
* HybridRecommendation - 混合推荐系统(协同过滤 + 基于内容)
*
* @author John
* @since 2026/5/7 15:31
* @version 2.0
*/
object HybridRecommendation {
def main(args: Array[String]): Unit = {
var spark: SparkSession = null
try {
val _spark = ConfigManager.createSparkSession("HybridRecommendation")
import _spark.implicits._
spark = _spark
println("\n" + "=" * 80)
println("🎯 混合推荐系统 (协同过滤 + 基于内容)")
println("=" * 80)
// 1. 加载数据
val ratingsDF = DataAccessLayer.loadRatings(spark)
val productsDF = DataAccessLayer.loadProducts(spark)
if (ratingsDF.count() == 0 || productsDF.count() == 0) {
println("⚠️ 数据不足,无法进行混合推荐")
return
}
// 2. 训练协同过滤模型
println("\n🔧 步骤1: 训练协同过滤模型...")
val Array(training, _) = ratingsDF.randomSplit(
Array(ConfigManager.getTrainingRatio, 1 - ConfigManager.getTrainingRatio),
seed = ConfigManager.getRandomSeed
)
val als = new ALS()
.setRank(ConfigManager.getAlsRank)
.setMaxIter(ConfigManager.getAlsMaxIter)
.setRegParam(ConfigManager.getAlsRegParam)
.setColdStartStrategy(ConfigManager.getAlsColdStartStrategy)
.setUserCol("userId")
.setItemCol("productId")
.setRatingCol("rating")
val cfModel = als.fit(training)
println("✅ 协同过滤模型训练完成")
// 3. 计算基于内容的特征
println("\n🔧 步骤2: 计算商品特征向量...")
val productFeatures = ContentBasedRecommendation.computeProductFeatures(productsDF)
println(s"✅ 商品特征计算完成: ${productFeatures.size} 个商品")
// 4. 设置混合权重
val cfWeight = ConfigManager.getCfWeight
val cbWeight = ConfigManager.getCbWeight
println(s"\n⚖️ 混合权重: 协同过滤=$cfWeight, 基于内容=$cbWeight")
// 5. 生成推荐
val userId = 1
val topK = ConfigManager.getTopK
println(s"\n👤 为用户 $userId 生成混合推荐 (Top-$topK):")
// 协同过滤推荐
val cfRecs = getCfRecommendations(spark, cfModel, userId, topK)
println(s"\n 📊 协同过滤推荐结果 (${cfRecs.size}条):")
cfRecs.foreach { case (pid, score) => println(f" - 商品ID=$pid, 分数=$score%.4f") }
// 基于内容推荐
val cbRecs = ContentBasedRecommendation.recommendForUser(
userId, ratingsDF, productsDF, productFeatures, topK
)
println(s"\n 📝 基于内容推荐结果 (${cbRecs.size}条):")
cbRecs.foreach { case (pid, score) => println(f" - 商品ID=$pid, 分数=$score%.4f") }
// 混合推荐
val hybridRecs = mergeRecommendations(cfRecs, cbRecs, cfWeight, cbWeight, 5)
println(s"\n 🎯 混合推荐结果 (${hybridRecs.size}条):")
if (hybridRecs.nonEmpty) {
hybridRecs.foreach { case (productId, score) =>
// 安全获取商品名称,避免空指针
val productNameRow = productsDF.filter($"product_id" === productId)
.select("product_name").take(1).headOption
productNameRow match {
case Some(row) =>
val name = row.getString(0)
println(f" - 商品: $name (ID=$productId), 混合分数: $score%.4f")
case None =>
println(f" - 商品ID=$productId (名称不存在), 混合分数: $score%.4f")
}
}
} else {
println(" ⚠️ 未生成混合推荐结果")
}
println("\n" + "=" * 80)
println("✅ 混合推荐系统运行完成")
println("=" * 80)
} catch {
case e: InterruptedException =>
println(s"\n⚠️ 程序被中断: ${e.getMessage}")
case e: Exception =>
println(s"\n❌ 运行出错: ${e.getClass.getSimpleName} - ${e.getMessage}")
e.printStackTrace()
} finally {
if (spark != null) {
println("\n🛑 正在关闭 SparkSession...")
spark.stop()
println("✅ SparkSession 已关闭")
}
}
}
def getCfRecommendations(
spark: SparkSession,
model: org.apache.spark.ml.recommendation.ALSModel,
userId: Int,
topK: Int): Map[Int, Double] = {
import spark.implicits._
val userDF = Seq(userId).toDF("userId")
model.recommendForUserSubset(userDF, topK)
.select(explode($"recommendations").as("rec"))
.select($"rec.productId".as("productId"), $"rec.rating".as("score"))
.collect()
.map(row => (row.getInt(0), row.getFloat(1).toDouble))
.toMap
}
def mergeRecommendations(
cfRecs: Map[Int, Double],
cbRecs: Seq[(Int, Double)],
cfWeight: Double,
cbWeight: Double,
topK: Int): Seq[(Int, Double)] = {
val cbRecsMap = cbRecs.toMap
val maxCfScore = if (cfRecs.nonEmpty) cfRecs.values.max else 1.0
val maxCbScore = if (cbRecsMap.nonEmpty) cbRecsMap.values.max else 1.0
val allProductIds = cfRecs.keySet ++ cbRecsMap.keySet
allProductIds.map { productId =>
val normalizedCf = cfRecs.getOrElse(productId, 0.0) / maxCfScore
val normalizedCb = cbRecsMap.getOrElse(productId, 0.0) / maxCbScore
val hybridScore = cfWeight * normalizedCf + cbWeight * normalizedCb
(productId, hybridScore)
}.toSeq
.sortBy(-_._2)
.take(topK)
}
}6. 推荐系统评估#
创建RecommendationEvaluator.scala,实现多维度评估:
知识点回顾:推荐系统评估分为预测准确度评估和排序质量评估。预测准确度用RMSE/MAE衡量评分预测精度;排序质量用Precision@K、Recall@K、NDCG@K衡量推荐列表质量。企业级系统需同时关注多个指标,避免单一指标优化导致推荐结果单一化。
| 评估指标 | 公式 | 含义 |
|---|---|---|
| RMSE | √(Σ(ŷ-y)²/N) | 均方根误差,越小越好 |
| MAE | Σ | ŷ-y |
| Precision@K | 推荐命中数/K | 推荐列表中用户喜欢的比例 |
| Recall@K | 推荐命中数/用户实际喜欢数 | 用户喜欢的物品被推荐的比例 |
| F1@K | 2×P×R/(P+R) | 准确率和召回率的调和平均 |
| NDCG@K | DCG@K/IDCG@K | 考虑排序位置的增益指标 |
import org.apache.spark.ml.recommendation.ALS
import org.apache.spark.ml.evaluation.RegressionEvaluator
import org.apache.spark.sql.{DataFrame, SparkSession}
import org.apache.spark.sql.functions._
/**
* RecommendationEvaluator
*
* @author John
* @since 2026/5/7 15:32
* @version 1.0
*/
object RecommendationEvaluator {
def main(args: Array[String]): Unit = {
var spark: SparkSession = null
try {
spark = ConfigManager.createSparkSession("RecommendationEvaluator")
println("=" * 80)
println("推荐系统评估器")
println("=" * 80)
val ratingsDF = DataAccessLayer.loadRatings(spark)
val dataCount = ratingsDF.count()
println(s"加载评分数据: $dataCount 条记录")
if (dataCount == 0) {
println("评分数据为空,请检查数据库")
return
}
val Array(training, test) = ratingsDF.randomSplit(
Array(ConfigManager.getTrainingRatio, 1 - ConfigManager.getTrainingRatio),
seed = ConfigManager.getRandomSeed
)
training.cache()
test.cache()
println(s"训练集: ${training.count()} 条, 测试集: ${test.count()} 条")
println("\n" + "-" * 40)
println("1. 预测准确度评估")
println("-" * 40)
val (model, _) = CollaborativeFiltering.trainAndEvaluate(training, test)
evaluatePredictionAccuracy(model, test)
println("\n" + "-" * 40)
println("2. 排序质量评估")
println("-" * 40)
evaluateRankingQuality(spark, model, test)
println("\n" + "-" * 40)
println("3. 参数敏感性分析")
println("-" * 40)
evaluateParameterSensitivity(training, test)
println("\n" + "=" * 80)
println("推荐系统评估完成")
println("=" * 80)
} catch {
case e: Exception =>
println(s"运行出错: ${e.getMessage}")
e.printStackTrace()
} finally {
if (spark != null) spark.stop()
}
}
def evaluatePredictionAccuracy(
model: org.apache.spark.ml.recommendation.ALSModel,
test: DataFrame): Unit = {
val predictions = model.transform(test)
val metrics = Seq("rmse", "mse", "mae", "r2")
val metricNames = Map(
"rmse" -> "RMSE (均方根误差)",
"mse" -> "MSE (均方误差)",
"mae" -> "MAE (平均绝对误差)",
"r2" -> "R² (决定系数)"
)
metrics.foreach { metric =>
val evaluator = new RegressionEvaluator()
.setMetricName(metric)
.setLabelCol("rating")
.setPredictionCol("prediction")
val value = evaluator.evaluate(predictions)
println(f" ${metricNames(metric)}: $value%.4f")
}
println("\n 指标解读:")
println(" - RMSE/MAE越小,预测精度越高")
println(" - R²越接近1,模型拟合效果越好")
}
def evaluateRankingQuality(
spark: SparkSession,
model: org.apache.spark.ml.recommendation.ALSModel,
test: DataFrame): Unit = {
import spark.implicits._
val topK = ConfigManager.getTopK
val userRecs = model.recommendForAllUsers(topK)
val userTestRatings = test.groupBy($"userId")
.agg(collect_set($"productId").as("testProducts"))
val userRecsWithTest = userRecs
.join(userTestRatings, Seq("userId"), "left")
.withColumn("recommendedProducts",
expr("transform(recommendations, x -> x.productId)"))
.withColumn("intersection",
expr("array_intersect(recommendedProducts, testProducts)"))
.withColumn("precision",
when(size($"testProducts") > 0,
size($"intersection") / lit(topK).cast("double"))
.otherwise(0.0))
.withColumn("recall",
when(size($"testProducts") > 0,
size($"intersection") / size($"testProducts").cast("double"))
.otherwise(0.0))
val avgPrecision = userRecsWithTest.select(avg($"precision")).head().getDouble(0)
val avgRecall = userRecsWithTest.select(avg($"recall")).head().getDouble(0)
val f1Score = if ((avgPrecision + avgRecall) > 0) {
2 * avgPrecision * avgRecall / (avgPrecision + avgRecall)
} else 0.0
println(f" Precision@$topK: $avgPrecision%.4f")
println(f" Recall@$topK: $avgRecall%.4f")
println(f" F1@$topK: $f1Score%.4f")
println("\n 指标解读:")
println(" - Precision@K: 推荐列表中用户实际喜欢的比例")
println(" - Recall@K: 用户喜欢的物品被推荐的比例")
println(" - F1@K: Precision和Recall的调和平均")
}
def evaluateParameterSensitivity(training: DataFrame, test: DataFrame): Unit = {
val ranks = Seq(5, 10, 20, 50)
val regParams = Seq(0.01, 0.05, 0.1)
println("\n Rank参数敏感性分析 (regParam=0.01, maxIter=10):")
println(f" ${"Rank"}%-6s ${"RMSE"}%-10s ${"训练时间(s)"}%-12s")
println(" " + "-" * 28)
ranks.foreach { rank =>
val als = new ALS()
.setRank(rank)
.setMaxIter(10)
.setRegParam(0.01)
.setColdStartStrategy("drop")
.setUserCol("userId")
.setItemCol("productId")
.setRatingCol("rating")
val startTime = System.currentTimeMillis()
val model = als.fit(training)
val trainingTime = (System.currentTimeMillis() - startTime) / 1000.0
val evaluator = new RegressionEvaluator()
.setMetricName("rmse")
.setLabelCol("rating")
.setPredictionCol("prediction")
val rmse = evaluator.evaluate(model.transform(test))
println(f" $rank%-6d $rmse%-10.4f $trainingTime%-12.2f")
}
println("\n RegParam参数敏感性分析 (rank=10, maxIter=10):")
println(f" ${"RegParam"}%-10s ${"RMSE"}%-10s")
println(" " + "-" * 20)
regParams.foreach { regParam =>
val als = new ALS()
.setRank(10)
.setMaxIter(10)
.setRegParam(regParam)
.setColdStartStrategy("drop")
.setUserCol("userId")
.setItemCol("productId")
.setRatingCol("rating")
val model = als.fit(training)
val evaluator = new RegressionEvaluator()
.setMetricName("rmse")
.setLabelCol("rating")
.setPredictionCol("prediction")
val rmse = evaluator.evaluate(model.transform(test))
println(f" $regParam%-10.2f $rmse%-10.4f")
}
}
}7. 推荐系统服务部署#
创建RecommendationService.scala,实现生产级推荐服务:
知识点回顾:生产环境推荐系统需要考虑:1)推荐结果缓存(Redis),避免重复计算;2)连接池管理,避免频繁创建/销毁连接;3)降级策略,当实时推荐失败时回退到热门推荐;4)模型版本管理,支持A/B测试。
import org.apache.spark.ml.recommendation.ALSModel
import org.apache.spark.sql.SparkSession
import org.apache.spark.sql.functions._
import redis.clients.jedis.{Jedis, JedisPool, JedisPoolConfig}
/**
* RecommendationService - 推荐服务(带Redis缓存)
*
* @author John
* @since 2026/5/7 15:32
* @version 2.0
*/
object RecommendationService {
private var jedisPool: JedisPool = _
/**
* 初始化 Redis 连接池(使用配置参数)
*/
private def initRedisPool(): Unit = {
try {
val poolConfig = new JedisPoolConfig()
poolConfig.setMaxTotal(ConfigManager.getRedisMaxTotal)
poolConfig.setMaxIdle(ConfigManager.getRedisMaxIdle)
poolConfig.setMinIdle(ConfigManager.getRedisMinIdle)
poolConfig.setTestOnBorrow(true)
poolConfig.setTestOnReturn(true)
poolConfig.setTestWhileIdle(true)
val password = ConfigManager.getRedisPassword
val timeout = ConfigManager.getRedisTimeout
if (password.isEmpty) {
jedisPool = new JedisPool(
poolConfig,
ConfigManager.getRedisHost,
ConfigManager.getRedisPort,
timeout
)
} else {
jedisPool = new JedisPool(
poolConfig,
ConfigManager.getRedisHost,
ConfigManager.getRedisPort,
timeout,
password
)
}
println(s"✅ Redis连接池初始化完成: ${ConfigManager.getRedisHost}:${ConfigManager.getRedisPort}")
} catch {
case e: Exception =>
println(s"❌ Redis连接池初始化失败: ${e.getMessage}")
throw e
}
}
private def closeRedisPool(): Unit = {
if (jedisPool != null && !jedisPool.isClosed) {
jedisPool.close()
println("Redis连接池已关闭")
}
}
private def withJedis[T](block: Jedis => T): T = {
val jedis = jedisPool.getResource
try {
block(jedis)
} finally {
jedis.close()
}
}
def main(args: Array[String]): Unit = {
var spark: SparkSession = null
try {
val _spark = ConfigManager.createSparkSession("RecommendationService")
import _spark.implicits._
spark = _spark
// 初始化 Redis
initRedisPool()
println("\n" + "=" * 80)
println("🚀 推荐服务启动(带Redis缓存)")
println("=" * 80)
// 1. 加载模型
val modelPath = ConfigManager.getModelSavePath + "/cf-model"
println(s"\n📦 正在加载模型: $modelPath")
val model = ALSModel.load(modelPath)
println("✅ 模型加载完成")
// 2. 批量生成推荐并写入 Redis
val topK = ConfigManager.getTopK
val ttlSeconds = ConfigManager.getCacheTtlHours * 3600
println(s"\n🔧 批量生成推荐结果并写入Redis (TTL=${ConfigManager.getCacheTtlHours}小时)...")
val userRecs = model.recommendForAllUsers(topK)
val productsDF = DataAccessLayer.loadProducts(spark)
var successCount = 0
var failCount = 0
userRecs.collect().foreach { row =>
try {
val userId = row.getInt(0)
// 使用 Row 结构解析推荐结果
val recommendationsRow = row.getList[org.apache.spark.sql.Row](1)
import scala.jdk.CollectionConverters._
val productIds = recommendationsRow.asScala.map(r => r.getInt(0)).mkString(",")
val scores = recommendationsRow.asScala.map(r => f"${r.getFloat(1)}%.4f").mkString(",")
withJedis { jedis =>
jedis.setex(s"rec:cf:user:$userId:products", ttlSeconds, productIds)
jedis.setex(s"rec:cf:user:$userId:scores", ttlSeconds, scores)
}
successCount += 1
} catch {
case e: Exception =>
failCount += 1
if (failCount <= 5) {
println(s" ⚠️ 用户推荐写入失败: ${e.getMessage}")
}
}
}
println(s"✅ 推荐结果写入完成: 成功=$successCount, 失败=$failCount")
// 3. 测试实时推荐
val testUserId = 1
println(s"\n👤 测试用户 $testUserId 的推荐:")
val recs = getRecommendations(spark, model, testUserId, topK, productsDF)
if (recs.nonEmpty) {
recs.foreach { case (productId, score) =>
val productNameRow = productsDF.filter($"product_id" === productId)
.select("product_name").take(1).headOption
productNameRow match {
case Some(row) =>
val name = row.getString(0)
println(f" - $name (ID=$productId), 分数: $score%.4f")
case None =>
println(f" - 商品ID=$productId (名称不存在), 分数: $score%.4f")
}
}
} else {
println(" ⚠️ 未获取到推荐结果")
}
println("\n" + "=" * 80)
println("✅ 推荐服务运行完成")
println("=" * 80)
} catch {
case e: InterruptedException =>
println(s"\n⚠️ 程序被中断: ${e.getMessage}")
case e: Exception =>
println(s"\n❌ 服务运行出错: ${e.getClass.getSimpleName} - ${e.getMessage}")
e.printStackTrace()
} finally {
closeRedisPool()
if (spark != null) {
println("\n🛑 正在关闭 SparkSession...")
spark.stop()
println("✅ SparkSession 已关闭")
}
}
}
def getRecommendations(
spark: SparkSession,
model: ALSModel,
userId: Int,
topK: Int,
productsDF: org.apache.spark.sql.DataFrame): Seq[(Int, Double)] = {
import spark.implicits._
val cached = withJedis { jedis =>
Option(jedis.get(s"rec:cf:user:$userId:products"))
}
cached match {
case Some(productIdsStr) if productIdsStr.nonEmpty =>
productIdsStr.split(",").map(_.toInt).zipWithIndex.map {
case (pid, idx) => (pid, 1.0 - idx * 0.1)
}.toSeq
case None =>
println(s"缓存未命中,实时计算用户 $userId 的推荐")
try {
val userDF = Seq(userId).toDF("userId")
val recs = model.recommendForUserSubset(userDF, topK)
.select(explode($"recommendations").as("rec"))
.select($"rec.productId".as("productId"), $"rec.rating".as("score"))
.collect()
.map(row => (row.getInt(0), row.getFloat(1).toDouble))
val ttlSeconds = ConfigManager.getCacheTtlHours * 3600
withJedis { jedis =>
val productIds = recs.map(_._1).mkString(",")
jedis.setex(s"rec:cf:user:$userId:products", ttlSeconds, productIds)
}
recs.toSeq
} catch {
case e: Exception =>
println(s"实时推荐失败,回退到热门推荐: ${e.getMessage}")
getHotProducts(productsDF, topK)
}
}
}
def getHotProducts(
productsDF: org.apache.spark.sql.DataFrame,
topK: Int): Seq[(Int, Double)] = {
productsDF.select("product_id")
.limit(topK)
.collect()
.zipWithIndex
.map { case (row, idx) => (row.getInt(0), 1.0 - idx * 0.1) }
.toSeq
}
}8. 运行推荐系统#
# 编译项目
mvn clean compile
# 运行协同过滤推荐
mvn exec:java -Dexec.mainClass="CollaborativeFiltering"
# 运行基于内容的推荐
mvn exec:java -Dexec.mainClass="ContentBasedRecommendation"
# 运行混合推荐系统
mvn exec:java -Dexec.mainClass="HybridRecommendation"
# 运行推荐系统评估
mvn exec:java -Dexec.mainClass="RecommendationEvaluator"
# 部署推荐系统服务
mvn exec:java -Dexec.mainClass="RecommendationService"知识点4:企业级推荐系统架构#
4.1 生产环境推荐系统架构#
┌─────────────┐ ┌──────────────┐ ┌──────────────┐
│ 用户请求 │────>│ API网关 │────>│ 推荐服务 │
│ (Web/App) │ │ (Nginx) │ │ (Spring) │
└─────────────┘ └──────────────┘ └──────┬───────┘
│
┌──────────────────┼──────────────────┐
│ │ │
v v v
┌──────────────┐ ┌──────────────┐ ┌──────────────┐
│ Redis缓存 │ │ 推荐引擎 │ │ 热门推荐 │
│ (命中返回) │ │ (Spark) │ │ (降级策略) │
└──────────────┘ └──────┬───────┘ └──────────────┘
│
┌────────────────┼────────────────┐
│ │ │
v v v
┌──────────────┐ ┌──────────────┐ ┌──────────────┐
│ 协同过滤 │ │ 内容推荐 │ │ 实时推荐 │
│ (离线计算) │ │ (离线计算) │ │ (在线计算) │
└──────────────┘ └──────────────┘ └──────────────┘4.2 推荐系统安全考量#
| 安全风险 | 防护措施 | 本实验实现 |
|---|---|---|
| 数据库密码泄露 | 配置外部化,不硬编码 | application.conf |
| Redis未授权访问 | 设置密码,连接池管理 | JedisPool + 密码认证 |
| 推荐结果被篡改 | 缓存设置TTL,定期刷新 | setex + 过期策略 |
| 服务不可用 | 降级策略,热门推荐兜底 | getHotProducts |
| 数据注入 | 输入验证,参数化查询 | DataFrame filter |
| 资源泄漏 | try-finally确保释放 | SparkSession/Redis |
4.3 推荐系统性能优化#
| 优化方向 | 方法 | 说明 |
|---|---|---|
| 离线预计算 | 定时批量生成推荐结果 | 减少在线计算压力 |
| 多级缓存 | Redis + 本地缓存 | 降低延迟 |
| 增量更新 | 只更新变化的用户 | 减少计算量 |
| 模型压缩 | 降低rank维度 | 减少存储和计算 |
| 并行化 | 增加Spark分区数 | 提高吞吐量 |
实验扩展#
- 隐式反馈推荐:将用户行为(浏览、点击、购买)转化为隐式反馈,使用
ALS.trainImplicit训练模型 - 实时推荐:结合Spark Streaming,根据用户实时行为动态调整推荐结果
- 冷启动处理:为新用户实现基于人口统计学的推荐,为新商品实现基于内容的推荐
- 推荐解释:为推荐结果生成解释(如"因为你购买了iPhone,推荐了AirPods")
- A/B测试:设计A/B测试框架,对比不同推荐算法的线上效果
- 多样性优化:在推荐列表中引入随机性,避免"信息茧房"效应
实验总结#
本实验通过以下步骤实现了电商平台的企业级商品推荐系统:
- 理论理解:深入理解了协同过滤(ALS矩阵分解)、基于内容(TF-IDF)、混合推荐三种核心算法的原理和适用场景
- 协同过滤实现:基于Spark MLlib的ALS算法,实现了用户-商品评分预测和个性化推荐
- 基于内容实现:通过TF-IDF提取商品特征,基于余弦相似度计算商品关联,解决物品冷启动问题
- 混合推荐实现:采用加权融合策略,结合协同过滤和基于内容的优势,通过归一化消除量纲差异
- 多维度评估:实现了预测准确度(RMSE/MAE/R²)和排序质量(Precision@K/Recall@K/F1@K)的完整评估体系
- 生产级部署:实现了配置外部化、Redis连接池管理、缓存降级策略、资源安全管理等企业级特性
通过本实验,我们不仅掌握了推荐系统的核心算法实现,更理解了从算法原型到生产级服务的完整链路,为企业级推荐系统的开发与部署奠定了坚实基础。