先说结论:Java 绝对能写 Spark MLlib,而且在自己主导的项目里,用 Java 把协同过滤推荐系统的整条链路打通,反而比“先招一个会 Scala 的人”更可控。这篇文章是我过去一段时间做过几版内容社区推荐方案后沉淀下来的完整记录,从 Spark 集群搭建、数据清洗、ALS 模型训练、指标评估,到参数调优和上线避坑,尽量按真实交付顺序讲清楚。适合两类人看:一类是 Java 技术栈、还没碰过 Spark MLlib 的后端工程师;另一类是已经搭过 Spark 跑过 ETL、想知道怎么用 Java 调通 MLlib 做推荐模型的人。我会把原理、代码、参数和经验放在一起讲,不搞纸上谈兵。
1. 为什么我坚持用 Java 调 MLlib,而不是先招个 Scala 工程师
1.1 一个真实项目里的语言阵营问题
先说背景。当时业务线要做一个“猜你喜欢”功能,团队技术栈是纯 Java,Spark 集群已经搭好,跑着离线数仓任务。第一次评审时,团队里的第一反应几乎都是:Spark 的机器学习库不是给 Scala 准备的吗?Java 能写推荐模型吗?
我的立场很直接:如果团队里全是 Java 工程师,与其为了一个模型单独引入 Scala 人员,不如用 Java 把整条链路打通。MLlib 的 Java API 虽然啰嗦,但它是官方一等公民,从 RDD 到 DataFrame 再到 ML Pipeline 全都有 Java 接口。更重要的是,推荐模型训练完之后要落到线上服务,Java 写在线推荐引擎本来就顺理成章。
这里有一个隐藏收益:从 Spark 训练任务到线上 RPC 服务,全用 Java,CI/CD、监控、日志、配置中心都能复用团队已有的基础设施。我后来在模型上线排障时感受特别明显——直接把 Spark 训练任务挂到 Java 服务监控面板里,日志格式统一,告警规则直接复用,不用为模型单独维护一套部署栈。
1.2 常见误区:Java 写不了 Spark
结合很多人在“Spark 集群搭建”“Spark 内存”这些话题上的困惑,我列几个最常见的误区:
- 误区一:Spark 任务只能写成 Scala。实际上 Spark 的 DataFrame、ML Pipeline、RDD 全有 Java 接口,ALS 类就在
org.apache.spark.ml.recommendation.ALS,直接 new 一个实例就能用。 - 误区二:Java 代码性能差。性能瓶颈在集群计算资源,编码语言影响不大。同样一个 ALS 任务,Java 和 Scala 提交到集群上,资源消耗基本一致。
- 误区三:Java API 没有好例子。这确实是现实问题,官方文档里 Java 示例又老又少,所以我才写了这篇文章。
1.3 Java 写 MLlib 真正麻烦的点
Java 写 MLlib 真正的挑战只有一个:MLlib 的 API 大量返回Dataset<Row>,你面对的是一堆Row对象而不是强类型类。这意味着你要习惯row.getAs、row.getString、row.getDouble这类取值方式,并且自己定义Encoder来映射 POJO。
这个坎跨过去之后就顺畅了。我会在后面的代码里给出封装思路,你不需要大面积处理 Row,只需要在任务边界做一次取数转换。
2. 协同过滤算法核心:矩阵分解的直觉与 ALS 计算逻辑
2.1 从“相似的人”到用户行为矩阵
协同过滤(Collaborative Filtering)是推荐系统里最直观的一类算法。它不做内容分析,只依赖用户对物品的历史行为。最早的做法是两种:UserCF 找“和我行为相似的用户”,把这些人喜欢的物品推荐给我;ItemCF 找“和我喜欢的物品类似的物品”,把相似物品推给用户。
这两种方法在数据量小的时候很好理解,但工程上有个致命的扩展性问题:用户量、物品量一大,两两相似度矩阵就是天文数字,而且行为稀疏时相似度根本算不准。所以工业界更常用基于模型的方法,核心思路是把用户和物品都降维到一个共享的低维空间,再通过向量内积预测偏好。
这就是矩阵分解。我们把用户-物品评分矩阵 R 拆成两个低阶矩阵相乘:R ≈ U V^T。其中 U 是用户隐因子矩阵,V 是物品隐因子矩阵,用户 u 对物品 i 的预测分就等于 U[u] 和 V[i] 两个向量的点积。
2.2 ALS 为什么被 Spark 选中
直接做矩阵分解,目标函数是计算所有已知项的 (U V^T - R)² 之和,问题在于 U 和 V 都是未知的,这个优化问题是非凸的,梯度下降容易停在局部最优,而且数据规模一大根本算不动。
ALS(Alternating Least Squares,交替最小二乘法)用一个非常聪明的操作绕开了非凸问题:先固定 U,把 U 当常量,这时关于 V 的目标函数变成了一组相互独立的最小二乘问题,每个物品的向量可以独立求解,而且有闭式解;然后固定 V,再求解 U。两个步骤交替反复迭代,直到收敛。这就是“交替”两个字的由来。
ALS 特别适合分布式计算的本质原因也在这:固定 U 时,每个物品的求解只依赖“对它有行为过的用户”的向量,天然可以按物品切分并行;固定 V 时同理,按用户切分并行。这是 Spark MLlib 选择 ALS 作为推荐模型核心算法的关键。
提示:理解了这一点,后面的参数就能对号入座了。rank 是隐因子维度,决定向量空间表达能力;regParam 是正则化系数,防止过拟合;alpha 是隐式反馈的置信度权重。调参时你才知道自己在调什么。
2.3 显式反馈与隐式反馈怎么选
协同过滤的输入数据有两种形态。显式反馈是用户主动打分,比如 1 到 5 星,数据质量高但采集成本大;隐式反馈是用户的点击、浏览、加购、观看时长这类间接行为,数据量大、天然存在,但有个特点:数据缺失不一定代表“不喜欢”,可能只是“没看到”。
MLlib 的 ALS 对这两种场景分别提供了支持:显式反馈直接用默认 ALS;隐式反馈把setImplicitPrefs设为 true,并通过setAlpha控制置信度权重——行为次数越多,样本置信度越高,训练时对“未发生的行为”做负样本加权。
我做的内容社区几乎没有打分体系,核心信号就是“是否点击”和“阅读时长”,这是典型的隐式反馈场景,所以我用的是setImplicitPrefs(true)这条路线。这也直接决定了数据预处理方式,下一章展开讲。
3. 先把版本搭配搞对:工程骨架与 SparkSession 的坑
3.1 版本组合是第一个大坑
这套文章系列的编号里,这一篇对应的正是“Java+Spark MLlib 推荐系统实战与优化”的完整版本。标题里的“441”就当系列随笔编号吧。现在直接说版本问题:Spark、Java、Hadoop、Scala 四者之间有一个牵一发动全身的组合关系,配错了连 SparkSession 都起不来。
| 组件 | 推荐版本 | 说明 |
|---|---|---|
| JDK | 1.8 或 11 | Spark 3.x 在这两个版本下最稳,JDK 17 要加启动参数处理模块化限制 |
| Spark | 3.3.x | 3.2 以上对 Java 11 更友好,MLlib 功能完整 |
| Scala 编译后缀 | 2.12 | 写 Java 代码不直接接触 Scala,但 Maven 坐标后缀要和 Spark 发行版一致 |
| Hadoop | 3.3.x | 本地开发时 Windows 需要 winutils |
| Maven | 3.8+ | 依赖用 spark-core 和 spark-mllib 两个即可 |
这套组合我跑了近半年,线上集群稳定。如果你是从零搭 Spark 集群,建议先按官方 Quick Start 起一个 SparkSession,跑通一只 Demo 再回来继续。集群搭建时,executor 内存一定要按数据量预留,这个在第 7 章专门讲。
3.2 Maven 最小配置
Java 工程里引入 Spark MLlib,Maven 坐标写起来很简单,但有两个细节要注意:第一,spark-mllib会传递依赖spark-core,所以理论上只声明spark-mllib也能跑;第二,如果不小心把spark-sql、spark-streaming全引进来,依赖冲突排查起来非常痛苦,尽量按需引入。
<properties> <maven.compiler.source>1.8</maven.compiler.source> <maven.compiler.target>1.8</maven.compiler.target> <spark.version>3.3.2</spark.version> </properties> <dependencies> <dependency> <groupId>org.apache.spark</groupId> <artifactId>spark-core_2.12</artifactId> <version>${spark.version}</version> </dependency> <dependency> <groupId>org.apache.spark</groupId> <artifactId>spark-mllib_2.12</artifactId> <version>${spark.version}</version> </dependency> </dependencies>3.3 SparkSession 初始化与 Windows 本地坑
Java 初始化 SparkSession 的代码如下:
SparkSession spark = SparkSession.builder() .appName("JavaALSRecommendation") .master("yarn") .config("spark.sql.shuffle.partitions", "200") .config("spark.serializer", "org.apache.spark.serializer.KryoSerializer") .getOrCreate();本地开发时,把master改成local[*],线上提交时换成yarn。这里有一个所有 Windows 开发者都会遇到的坑:本地环境跑 Spark 必须能定位到 winutils.exe,否则会报Failed to locate the winutils binary。解决方案是下载与 Hadoop 版本匹配的 winutils 放到某目录,比如D:/hadoop/bin,然后设置环境变量HADOOP_HOME:
System.setProperty("hadoop.home.dir", "D:\\hadoop");还有一个“高级坑”:JDK 17 的模块化限制。如果非要用 JDK 17,运行前要加一串启动参数:
--add-opens=java.base/java.lang=ALL-UNNAMED --add-opens=java.base/java.lang.invoke=ALL-UNNAMED --add-opens=java.base/java.lang.reflect=ALL-UNNAMED --add-opens=java.base/java.io=ALL-UNNAMED --add-opens=java.base/java.net=ALL-UNNAMED --add-opens=java.base/java.util=ALL-UNNAMED不想折腾就退回 JDK 11,能省下至少两小时。
4. 训练数据决定模型上限:rating 格式、时间切分与隐式反馈处理
4.1 最基础的输入格式:user, item, rating
MLlib ALS 的输入数据核心是三列:用户 ID、物品 ID、评分值。评分值可以是显式打分,也可以是隐式行为的量化值。
数据类型上,两列 ID 建议一进一出都用数值型,最好直接用 int 或 long。字符串虽然也能通过 setUserCol 指定,但会多一步索引转换,而且分布式环境下字符串 shuffle 的数据量比数值型大很多,会影响训练性能。下面是读取 CSV 并转成训练 DataFrame 的 Java 代码:
Dataset<Row> raw = spark.read() .option("header", "true") .option("inferSchema", "true") .csv("hdfs:///data/user_behavior.csv"); Dataset<Row> ratings = raw.select( col("user_id").cast("int").as("userId"), col("item_id").cast("int").as("itemId"), col("behavior_value").cast("double").as("rating") );这里有个容易忽略的细节:数据集清洗后一定不能有 NaN、Infinity 或负数评分。MLlib 不会自动去脏数据,训练样本里混入异常值,训练出来的模型预测分数会非常奇怪。我第一版就是因为在行为时长列里混了几个负数,导致推荐结果的排序完全不合理,排查了很久才发现是数据问题。
4.2 时间窗口切分:推荐系统不能纯随机划分
训练集和测试集的划分,推荐场景和普通分类场景有个重要的区别:不能纯随机切。用户行为数据天然带时间属性,如果随机切分,就可能出现模型拿“未来数据”学习,再预测“过去数据”的穿越问题。严谨的做法是按时间窗口切分:用前几分钟干净的行为训练,留最近一段做评估。
数据量大之后,我推荐按用户维度做一个简单而稳的划分:对每个用户把行为时间排序,最近的日志数据留作验证,其余做训练。核心是保证同一个用户的行为不会被切成一模一样的跨时间片段:
Dataset<Row> withTime = ratings.withColumn("date", to_date(col("log_date"))); Dataset<Row> userLatest = withTime.groupBy("userId") .agg(max("date").as("max_date")); Dataset<Row> train = withTime.join(userLatest, "userId") .filter(col("date").lt(col("max_date"))); Dataset<Row> test = withTime.join(userLatest, "userId") .filter(col("date").equalTo(col("max_date")));这种“每个用户的最新一条行为做测试”的离线评估方案,比 randomSplit 更贴近线上真实行为顺序。对内容推荐、电商推荐这类用户行为时效性强的业务尤其适用。
4.3 隐式反馈的量化:让“没看到”不等于“不喜欢”
如果手上只有点击、浏览这类隐式数据,需要先把行为转成一个可训练的数值。MLlib 的隐式 ALS 对输入的理解是:rating 值越大代表偏好越强,值本身是小是小并不重要,内部会通过置信度机制加权。所以转换时不必纠结“点击算几分”,但一定要区分主要行为的权重。
一个常用的转换思路是:点击次数、收藏/加购、平均停留时长按业务权重加权求和:
rating = 点击次数 * 1 + 收藏/加购 * 3 + 平均停留分钟数 * 2转换之后做一次分位数截断,把极端值压到合理范围,比如 0.5 到 5.0,避免长尾用户行为值爆炸影响模型稳定。这一步没有唯一标准答案,完全按业务场景调节。
5. Java 版 ALS 实战:训练模型、批量 TopN 推荐与相似物品
5.1 训练一个 ALS 模型
核心调用非常短,关键是把参数设对。下面这段代码可以直接跑:
ALS als = new ALS() .setMaxIter(15) .setRank(20) .setRegParam(0.1) .setUserCol("userId") .setItemCol("itemId") .setRatingCol("rating") .setColdStartStrategy("drop"); ALSModel model = als.fit(train); model.write().save("hdfs:///model/als_model_v1");几个参数建议亲手设一次:
- setMaxIter:最大迭代次数。太小欠拟合,太大耗时且可能过拟合。一般 10 到 20 起步,观察损失收敛曲线再微调。
- setRank:隐因子维度。越大表达力越强,但计算量成倍增加,数据稀疏时反而容易过拟合。内容推荐里 20 到 50 是常见范围。
- setRegParam:L2 正则化系数。0.01 到 0.1 是常用区间,可以用网格搜索自动调。
- setColdStartStrategy:必须设置为 drop,否则预测时遇到训练集中没见过的用户或物品,模型会输出 NaN 分数,评估和线上逻辑都会被污染。这是最容易忽略的一条。
5.2 给用户做批量 TopN 推荐
模型训练完之后,要做的第一件事通常是离线批量算 TopN 候选。MLlib 提供了 recommendForUserSubset 接口,可以指定一批用户、每个用户取 N 个物品:
Dataset<Row> usersToRecommend = spark.createDataFrame(userIdList, LongType.class) .toDF("userId"); Dataset<Row> recommendations = model.recommendForUserSubset(usersToRecommend, 20); recommendations.show(false);返回的 recommendations 列是一个结构数组,格式大概是[userId, [{itemId, rating}, {itemId, rating}...]]。要把推荐结果落成一张宽表供线上 RPC 读取,需要拆开这个结构。Java 里可以直接用 explode 函数拉平:
Dataset<Row> exploded = recommendations .select(col("userId"), explode(col("recommendations")).as("rec")); Dataset<Row> online = exploded.select( col("userId"), col("rec.itemId").alias("recItemId"), col("rec.rating").alias("score") );拉平后的数据可以直接写到 HBase 或 Redis,线上推荐接口读取后直接返回。我的生产做法是每天凌晨算好用户 TopN 列表,以 userId 为 key 写入 Redis,TTL 设成两天,线上查询命中率非常高,性能远好于在线实时计算。
5.3 顺便把“看了又看”也做了
协同过滤有一个很值钱的副产品:物品向量本身。矩阵分解训练完成后,物品向量藏在 ALSModel 的 itemFactors 里。对任意物品,用它的隐因子向量和其他物品的隐因子向量做内积,按得分排序,就能得到相似物品列表。这相当于同时完成了“猜你喜欢”和“相关推荐”两件事。
Java 里取物品向量的方式很简单:
Dataset<Row> itemFactors = model.itemFactors(); // 列名:id, features(factor 隐因子向量)计算相似度可以直接用 MLlib 的 RowMatrix + columnSimilarities,或者更务实一点:把物品向量同步到在线 Redis,线上用向量内积实时算 Top20 相似物品。几十万物品规模下,这个方案响应时间完全够用。我在项目冷启动期就靠这个先用了起来,成本低,业务价值直接。
6. 模型评估不能只看 RMSE:离线指标与线上业务对齐
6.1 RMSE 怎么算:Java 里用 RegressionEvaluator
训练完模型后,最基础的评估用 MLlib 自带的回归评估器就能做:
Dataset<Row> predictions = model.transform(test) .filter(col("prediction").isNotNull()); RegressionEvaluator evaluator = new RegressionEvaluator() .setMetricName("rmse") .setLabelCol("rating") .setPredictionCol("prediction"); double rmse = evaluator.evaluate(predictions); System.out.println("RMSE = " + rmse);RMSE 计算的是预测分数和真实分数的均方根误差,越小越好。参考意义是:在经典的公开评分数据集上,最优模型 RMSE 在 0.85 到 0.90 之间。但你的业务数据不同,绝对数值没有普适意义,RMSE 更多用来做模型迭代的相对比较。
6.2 RMSE 好看,不代表用户觉得推荐得好
这是推荐系统评估里最容易踩的认知坑。RMSE 评估的是“评分预测准不准”,但线上用户感知的是“推荐列表里有没有我想要的东西”,这是两个目标。我亲身经历过:调参把 RMSE 从 1.2 降到 0.98,AB 测试点击率反而小幅下降。原因是 RMSE 优化偏保守,会让模型倾向于把分数都预测到中等区间,而推荐系统真正需要的是“把用户最喜欢的少数物品排到最前面”。
所以离线评估一定要同时关注几个面向业务的指标:
- 召回率:测试集里用户真实点击过的物品,有多少出现在推荐 TopN 里。
- 精确率:推荐列表里用户真实点击过的比例。
- 覆盖率:推荐列表覆盖了多少物品池,避免只推头部爆款。
这些指标用 DataFrame 的 join 和聚合在 Java 里就能算,不需要额外依赖。我的习惯是把 RMSE 当作模型迭代的稳定性信号,把 TopN 召回/精确率当作真正的业务指标,两者结合着看。
6.3 线上验证:离线评估只是门票
离线评估只能证明“模型没坏”,真正拍板的还是线上 AB 测试。我的实践流程是:先用时间切分算出候选模型,离线对比 RMSE 和 TopN 指标;然后小流量 AB 运行一到两周,看点击率、人均推荐点击数、转化率;指标正向且显著再全量放量。
数据量不大时,别指望模型一夜翻盘。先把线上日志埋点做好,否则连 AB 结果都无法解释。
7. 参数优化与上线避坑:网格搜索、内存、冷启动
7.1 用 CrossValidator 做网格搜索
手动调 rank、regParam、alpha 很费时费力。MLlib 标准做法是用 CrossValidator + ParamGridBuilder:
ParamGridBuilder gridBuilder = new ParamGridBuilder(); ParamGrid paramGrid = gridBuilder .addGrid(als.rank(), new int[]{10, 20, 40}) .addGrid(als.regParam(), new double[]{0.01, 0.1}) .addGrid(als.maxIter(), new int[]{10, 20}) .build(); CrossValidator validator = new CrossValidator() .setEstimator(als) .setEvaluator(evaluator) .setEstimatorParamMaps(paramGrid) .setNumFolds(3); CrossValidatorModel cvModel = validator.fit(train);这里有两点要提前说:第一,交叉验证开销很大,3 折乘以 12 种参数组合等于 36 次模型训练,小数据集没问题,大数据集一定要先采样或者缩小网格;第二,CrossValidator 自动选参的评估器默认是 RMSE,这又回到第 6 章的老话题——调参目标别全押在 RMSE 上,必要时可以用离线 TopN 指标做二次筛选。
7.2 Spark 内存三件套:Kryo、executor 内存、checkpoint
线上跑推荐训练,最常见的问题就是内存溢出和 shuffle 爆炸。按我的经验,三件套依次做好:
- 换 Kryo 序列化器:MLlib 的隐因子向量是 double 数组,Java 默认的 JavaSerializer 对象头开销很大。在 SparkSession 里设置
spark.serializer为 Kryo,内存占用能少 40% 左右。 - executor 内存与并行度匹配:训练前先看 Spark UI 里的内存水位。百万级用户、500 万物品、2 亿行为量级的数据,我在生产环境用 40 个 executor、每个 4GB 内存,
spark.sql.shuffle.partitions设成 400,作业稳定不溢写。 - 用 checkpoint 切断 RDD 血缘:ALS 训练迭代次数多,血缘链会非常长。先设置 checkpoint 目录,再对训练数据做一次 checkpoint,可以避免迭代过程中依赖链爆炸:
spark.sparkContext().setCheckpointDir("hdfs:///tmp/spark-checkpoint"); train.checkpoint();7.3 冷启动:让新用户、新物品不裸奔
模型再漂亮,线上总会遇到新用户、新物品。ALS 对没有历史行为的用户算不出可靠的隐因子向量。我的工程处理顺序是:
第一,ALS 里 coldStartStrategy 设为 drop,保证预测不吐 NaN;第二,对新用户直接走热门推荐兜底,用物品被浏览总次数热度榜;第三,如果用户已经有少量行为,比如刚收藏了几篇文章,立刻用内容相似逻辑补冷启动。这套“模型推荐 + 热门兜底 + 内容相似补充”的组合,在冷启动阶段非常管用。
7.4 重训周期与增量更新节奏
离线 ALS 重训频率不能太低,否则跟不上用户行为演化;也不能太高,否则资源扛不住。日活千万级以下的内容社区,一天全量重训一次完全够。如果行为变化很快,可以采用分段增量:先用每天的完整数据训练一个基础模型,期间每小时把新增用户行为套用 ALS 的最小二乘闭式解,快速算出新用户因子,再拼进线上推荐。
这个思路不需要改模型结构,收益却很明显。新增用户通常在几分钟内就能获得个性化推荐,而不是等第二天全量重训。
8. 个人调优经验与后续扩展
8.1 让 Java 代码组织得像 Scala Demo 一样清爽
回到开头说的 Java API“太啰嗦”的问题。我实际项目里用几个工具类把 MLlib 调用封了起来,对外暴露的方法不超过五个:loadRatings、trainModel、recommendForUsers、recommendSimilarItems、evaluateModel。调用方完全不感知 Spark 的存在。
建议你也这样封装边界:训练任务和在线服务不要互相依赖。训练任务只产出模型文件或 Redis 缓存,在线服务只读这些产物。后期无论换算法还是换框架,在线服务都不需要大面积改动。
8.2 别急着上复杂模型:先用 ALS 把基线建起来
做推荐系统最忌讳一上来就套深度学习模型。我的路径是:先用 ALS + 隐式反馈建起第一个能跑的推荐,把召回、过滤、排序、缓存链路全部打通,再用业务指标推动迭代。见过太多项目组在矩阵还没训练出来的时候就开始讨论 Transformer,结果一个月连线上流量都没接上。
ALS 作为基线,能在几分钟内给出一个还算靠谱的推荐结果。先跑起来,再谈优化,这个顺序太重要了。
8.3 跑通之后可以走的三条路
如果 ALS 在你们那里已经稳定运行,后续扩展我比较推荐三条路:
- 第一,在召回阶段做多路召回,把 ALS 召回结果和热门、新品、基于内容的召回合并,再用轻量排序模型做粗排。
- 第二,用 ALS 训练出的用户向量做用户分群,给不同分群配置不同的业务策略。
- 第三,把 ALS 训练与特征计算全部整合进 Spark SQL 流程,用统一的 SQL 任务编排离线训推链路。
个人经验总结到这里。如果你是 Java 技术栈、第一次用 Spark MLlib 做推荐系统,按本文顺序走一遍就好:原理过一遍,版本按第 3 章的组合搭,数据按第 4 章处理,训练用第 5 章代码,评估按第 6 章指标看,再根据第 7 章的坑去优化。跑通一次之后回头再看,这些东西其实没有想象中那么复杂。