Angel 中的因子分解机(FM)算法:原理、参数配置与分布式训练实战

发布时间:2026/10/7 9:55:37
Angel 中的因子分解机(FM)算法:原理、参数配置与分布式训练实战 人工智能机器学习分布式训练图计算后端【免费下载链接】angelA Flexible and Powerful Parameter Server for large-scale machine learning项目地址https://gitcode.com/gh_mirrors/an/angel点击查看免费下载因子分解机Factorization Machine简称 FM由 Steffen Rendle 提出是一种基于矩阵分解的机器学习算法可对任意的实值向量进行预测。本文以 Angel 仓库中的 docs/algo/fm_on_angel.md 为核心结合源码与配置文件完整讲解 FM 的数学模型、Angel 上的wide embedding实现结构、数据格式、全部训练参数以及端到端的提交命令帮助读者在 Angel 参数服务器架构上快速跑通 FM 训练与预测。1. FM 算法介绍1.1 Factorization Model因子分解模型FM 的核心假设是特征之间的交互项可以用两个低维隐向量因子向量的点积来刻画从而在高度稀疏的数据上也能够可靠地估计二阶交叉特征。其模型形式为ŷ(x) b Σ(i1..n) w_i·x_i Σ(i1..n) Σ(ji1..n) v_i, v_j·x_i·x_j其中v_i, v_j是两个 k 维因子向量embedding 向量的点积v_i, v_j Σ(f1..k) v_i,f · v_j,f模型参数为w_0 ∈ R全局偏置项 bw ∈ R^n一阶权重向量wide 部分V ∈ R^(n×k)二阶交叉项对应的因子矩阵embedding 部分。其中v_i表示用 k 个因子特征来表征特征 i 的隐向量k 是决定因子分解表达能力的超参数即本文后续的ml.fm.rank。相比直接对每一对特征单独学习交互权重FM 用 k 维隐向量共享参数因此参数量从 O(n²) 降为 O(n·k)且稀疏特征也能通过共享学习得到稳定的估计——这正是 FM 擅长处理高度稀疏数据场景、并保持线性计算复杂度的原因。1.2 Factorization Machines as PredictorsFM 的预测任务FM 可以用于多种预测任务回归ŷ直接作为预测值使用优化准则为最小化平方误差分类使用ŷ的符号正负作为分类预测结果参数通过合页损失函数Hinge Loss或对数损失/逻辑回归Log Loss进行估计。Angel 中 FM 默认使用LogLoss见下文源码佐证。2. FM on Angel模型结构、训练流程与预测输出2.1 FM 模型结构wide embedding在 Angel 中FM 算法的模型由两部分组成wide典型的线性模型对应 FM 公式中的一阶项b Σ w_i·x_iembedding因子分解部分对应二阶交叉项ΣΣ v_i,v_j·x_i·x_j。最终的输出结果为 wide 与 embedding 两部分之和。这一点可以从模型实现源码中得到直接印证在 angel-ps/mllib/src/main/scala/com/tencent/angel/ml/classification/FactorizationMachines.scala 的buildNetwork()方法中模型被构建为SimpleInputLayer(input, 1, new Identity(), optimizer)—— 即 wide 线性部分Embedding(embedding, numFields * numFactors, numFactors, ...)—— 即 embedding 因子矩阵维度为field 数 × rankBiInnerSumCross(innerSumPooling, embedding)—— 计算所有 field 两两交叉项内积的求和SumPooling(sumPooling, 1, Array(wide, innerSumCross))—— 将 wide 输出与交叉项输出相加得到最终预测SimpleLossLayer(simpleLossLayer, join, lossFunc)—— 默认损失函数为new LogLoss()用于二分类训练。其中二阶交叉项的求和实现位于 angel-ps/mllib/src/main/scala/com/tencent/angel/ml/core/network/layers/linear/BiInnerSumCross.scala它利用恒等式Σ(ij) v_i,v_j ½·(‖Σv‖² − Σ‖v_i‖²)将两两交叉计算从 O(field²) 降为 O(field·k)在calOutput()中对每个 field 的向量累加sum、累加平方和square_sum再以(sum·sum − square_sum) / 2一次性求出全部二阶交互项之和这是 FM 保持线性计算复杂度的关键实现。此外仓库还提供了 FM 的 JSON 网络配置样例 angel-ps/mllib/src/test/jsons/fm.json其中完整声明了上述网络结构wide、embedding、BiInnerSumCross、SumPooling、simplelosslayer可作为可视化理解与自定义改写的参考。2.2 FM 训练过程worker 与 PS 的协作Angel 使用梯度下降方法迭代训练 FM 模型每一轮迭代中 worker 与 PS 的分工如下worker每次迭代从 PS 上拉取 wide 和 embedding 矩阵到本地计算对应的梯度更新值然后 push 回 PSPS汇总所有 worker 推送的梯度更新值并取平均通过优化器计算新的 wide 和 embedding 模型参数并进行更新。这一拉取-计算梯度-推送-汇总更新的流程是 Angel 参数服务器架构下所有模型包括 FM训练的标准模式相关入口可参见 angel-ps/mllib/src/main/scala/com/tencent/angel/ml/core/graphsubmit/GraphRunner.scala 及其配套的GraphTrainTask、GraphLearner。2.3 FM 预测输出格式FM 预测结果的输出格式为rowID,pred,prob,label字段含义rowID样本所在的行 ID从 0 开始计数pred样本的预测结果值即 ŷprob样本相对该预测结果的概率label预测样本被分到的类别当预测结果值pred大于 0 时 label 为 1小于 0 时 label 为 -1。3. 数据格式说明Angel 的 FM 支持libsvm与dummy两种数据格式通过参数ml.data.type指定。更完整的格式规范可参考 docs/algo/data_format.md。libsvm 格式每行一个样本字段以空格分隔格式为label index1:value1 index2:value2 ...特征 index 从 1 开始计数1 1:1 214:1 233:1 234:1dummy 格式每行一个样本label index1 index2 index3 ...特征 index 从 0 开始计数仅列出特征值为 1 的 index其余特征值为 01 1 214 233 234补充说明源自 docs/algo/data_format.md训练数据的 label 为样本标签二分类为 {0, 1} 或正负 1预测数据的 label 为样本 index若输入数据分隔符不是空格可用ml.data.splitor指定例如ml.data.splitor,标签转换可用ml.data.label.trans.class见参数表。4. 参数说明含默认值与取值范围下表完整列出了 FM 训练涉及的核心参数。其中默认值均取自 angel-ps/mllib/src/main/scala/com/tencent/angel/ml/core/conf/MLConf.scala 中的常量定义可直接在-D命令行参数或 JSON 配置中覆盖。参数名含义默认值 / 取值ml.epoch.num迭代轮数默认 30ml.feature.index.range特征索引范围应不小于最大特征 index默认 -1需设置ml.model.size特征维数默认 -1需设置ml.data.validate.ratio验证集采样率默认 0.05ml.data.type数据类型libsvm默认/dummyml.learn.rate学习率默认 0.5ml.opt.decay.class.name学习率衰减类默认StandardDecay可选ConstantLearningRate、WarmRestarts等见 decayer 目录ml.opt.decay.on.batch是否对每个 mini batch 衰减默认 falseml.opt.decay.alpha学习率衰减参数 alpha默认 0.001ml.opt.decay.beta学习率衰减参数 beta默认 0.001ml.opt.decay.intervals学习率衰减参数 intervals默认 100ml.reg.l2L2 正则项系数默认 0.0action.type任务类型训练用train预测用predict另支持inctrainml.fm.field.num输入数据领域field的个数默认 -1需按数据设置ml.fm.rankembedding 中 vector因子向量的长度即超参数 k默认 8ml.inputlayer.optimizer优化器类型默认Momentum可选adam、ftrl、momentum等ml.embedding.optimizerembedding 层优化器默认同ml.inputlayer.optimizerml.data.label.trans.class是否对标签进行转换默认NoTrans可选ZeroOneTrans转为 0-1、PosNegTrans转为正负 1、AddOneTrans加 1、SubOneTrans减 1ml.data.label.trans.threshold标签转换阈值默认 0ZeroOneTrans、PosNegTrans需配合使用大于阈值的为 1ml.data.posneg.ratio正负样本重采样比例默认 -1对正负样本相差较大如 5 倍以上的场景有效要点解读ml.fm.field.num与ml.fm.rank直接决定 embedding 矩阵的尺寸。在 FactorizationMachines.scala 中Embedding层输出维度为numFields * numFactors即field.num × rank二者分别对应 MLConf.scala 中的ML_FIELD_NUMml.fm.field.num默认 -1与ML_RANK_NUMml.fm.rank默认 8。field.num表示样本被划分出的特征域个数例如 one-hot 后的特征分组数一般取特征维度除以每域独热维度。优化器ml.inputlayer.optimizer对应 wide 层的优化器ml.embedding.optimizer可单独指定 embedding 层的优化器。可选实现包括Adam、FTRL、Momentum、AdaGrad、AdaDelta、SGD等见 angel-ps/mllib/src/main/scala/com/tencent/angel/ml/core/optimizer 目录。FTRL 常用于大规模稀疏场景Adam/Momentum 则更通用。学习率衰减衰减类实现位于 angel-ps/mllib/src/main/scala/com/tencent/angel/ml/core/optimizer/decayer 目录包含ConstantLearningRate、StandardDecay、CorrectionDecay、WarmRestarts、StepSizeScheduler等。5. 提交命令与运行示例5.1 命令行方式-D 参数可以通过下面的命令提交 FM 训练任务需先按 docs/deploy/source_compile.md 完成编译得到bin/angel-submit脚本$input_path、$featureNum为按数据实际情况设置的环境变量../../bin/angel-submit \ -Dml.epoch.num20 \ -Dangel.app.submit.classcom.tencent.angel.ml.core.graphsubmit.GraphRunner \ -Dml.model.class.namecom.tencent.angel.ml.classification.FactorizationMachines \ -Dml.feature.index.range$featureNum \ -Dml.model.size$featureNum \ -Dml.data.validate.ratio0.1 \ -Dml.data.typelibsvm \ -Dml.learn.rate0.1 \ -Dml.reg.l20.03 \ -Daction.typetrain \ -Dml.fm.field.num11 \ -Dml.fm.rank8 \ -Dml.inputlayer.optimizerftrl \ -Dangel.train.data.path$input_path \ -Dangel.workergroup.number20 \ -Dangel.worker.memory.mb20000 \ -Dangel.worker.task.number1 \ -Dangel.ps.number20 \ -Dangel.ps.memory.mb10000 \ -Dangel.task.data.storage.levelmemory \ -Dangel.job.nameangel_l1命令中各参数的作用说明算法入口angel.app.submit.class固定为GraphRunnerml.model.class.name指定 FM 模型类com.tencent.angel.ml.classification.FactorizationMachines数据与特征angel.train.data.path指向训练数据HDFS 路径ml.feature.index.range与ml.model.size需等于特征维数ml.data.validate.ratio划分验证集比例学习与优化ml.learn.rate、ml.reg.l2、ml.inputlayer.optimizerftrl分别控制学习率、L2 正则与优化器FM 结构ml.fm.field.num11、ml.fm.rank8决定 embedding 矩阵尺寸集群资源angel.workergroup.numberworker 数、angel.ps.numberPS 数、angel.worker.memory.mb/angel.ps.memory.mb内存、angel.worker.task.number每 worker 任务数、angel.task.data.storage.level数据存储级别等。上述任务调度相关参数与 docs/deploy/run_on_yarn.md 中描述的 YARN 提交方式一致在本地测试环境可参考 docs/deploy/local_run.md。5.2 JSON 配置方式除-D命令行参数外FM 也支持通过 JSON 配置声明数据、训练、模型与网络结构仓库自带的 angel-ps/mllib/src/test/jsons/fm.json 即为完整示例dummy 格式、13 个 field、rank8、momentum 优化器、WarmRestarts 衰减。它清晰展示了 wide/embedding/BiInnerSumCross/SumPooling/损失层的逐层定义便于读者理解 FM 网络的组装方式并在此基础上自定义网络。5.3 预测任务将action.type改为predict并指定模型路径与预测数据路径angel.predict.data.path即可加载训练产出的模型进行预测输出即为上文 2.3 节所述的rowID,pred,prob,label四列格式。6. 深入阅读算法原始文档docs/algo/fm_on_angel.mdFM 模型实现源码angel-ps/mllib/src/main/scala/com/tencent/angel/ml/classification/FactorizationMachines.scala二阶交叉求和算子angel-ps/mllib/src/main/scala/com/tencent/angel/ml/core/network/layers/linear/BiInnerSumCross.scala参数常量定义angel-ps/mllib/src/main/scala/com/tencent/angel/ml/core/conf/MLConf.scala优化器与衰减器angel-ps/mllib/src/main/scala/com/tencent/angel/ml/core/optimizer数据格式规范docs/algo/data_format.mdFM 网络 JSON 样例angel-ps/mllib/src/test/jsons/fm.json赞分享人工智能机器学习分布式训练图计算后端【免费下载链接】angelA Flexible and Powerful Parameter Server for large-scale machine learning项目地址https://gitcode.com/gh_mirrors/an/angel点击查看免费下载相关推荐如何一键保存网页精华Notesnook网页剪藏器 Web Clipper 完整使用指南如何一键保存网页精华Notesnook网页剪藏器 Web Clipper 完整使用指南 Notesnook Web Clipper网页剪藏器 是 Note人工智能机器学习分布式训练图计算后端Apache MXNet 参数初始化器 mxnet.initializer 完全指南从常量初始化到 Xavier/MSRA 与自定义注册Apache MXNet 参数初始化器 mxnet.initializer 完全指南从常量初始化到 Xavier/MSRA 与自定义注册 导读 本文围绕 Ap人工智能机器学习分布式训练图计算后端Angel 图计算之 Motif 计数33 种有向图范式特征的算法原理、参数配置与 Spark on Angel 实战Angel 图计算之 Motif 计数33 种有向图范式特征的算法原理、参数配置与 Spark on Angel 实战 本篇技术指南以 Angel 开源仓库中人工智能机器学习分布式训练图计算后端上一篇Netty 4.1 心跳服务与断线重连实战IdleStateHandler 空闲检测与 ChannelFutureListener 自动重连下一篇如何快速修复 DLL 报错VisualCppRedist AIO 一个安装包装齐所有 Visual C 运行库创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考