机器学习中的公式推导难吗,怎么快速掌握公式计算?
- 云服务器
- 2026-08-12
- 7
机器学习公式推导的核心不在于数学技巧,而在于建立从“损失函数”到“参数更新”的因果链条;公式计算的关键则在于手推每一个维度的变换轨迹,两者结合才能真正掌控模型行为。
为什么你推导公式总卡在中间步骤
很多人在看论文时觉得公式推导像变魔术,从一个等式跳到另一个等式,中间省掉的过程全靠猜,这不是你的问题,而是多数教程默认读者已经具备“矩阵求导直觉”,但实际做项目时你会发现,真正卡住你的往往是那些看似不起眼的维度变换。
我建议你把推导过程拆成四个可验证的阶段:定义符号体系、写出前向传播表达式、求解梯度、验证维度一致性,任何一步出现“感觉不对但又说不出哪里错”,几乎都是符号体系混乱导致的。
以最简单的线性回归为例,损失函数写作 L = ||Xw y||²,很多人直接套用 ∇L = 2Xᵀ(Xw y),却说不清为什么转置放在前面,如果你从标量对向量求导的定义出发,逐项展开,就会发现转置位置是偏导数的分子布局决定的,不是死记硬背。
行业里有个共识:手推公式至少三遍,第一遍照抄原文推导,第二遍合上书本独立推导,第三遍尝试用不同的求导顺序得到相同结果,经过这三轮,你才真正建立起了自己的推导直觉。
链式法则的矩阵形态:从标量到向量的思维跃迁
链式法则在标量微积分里很简单,但到了机器学习,所有变量都变成向量和矩阵,链式法则的形态就要重新理解。
设 z = f(y),y = g(x),标量情况下 dz/dx = dz/dy · dy/dx,但当 y 是向量,x 也是向量时,这个乘号变成矩阵乘法,而且顺序不能乱,关键在于维度匹配:结果矩阵的行数等于最终标量对中间变量的梯度维度,列数等于中间变量对自变量的梯度维度。
实际操作中,我常用一个技巧:先用小维度数值验证,比如手动设定 x 为 2 维向量,y 为 3 维向量,把每个分量都写出来,用数值微分验证解析梯度,这个习惯能帮你省掉大量调试时间。
以反向传播中的核心步骤为例,假设损失 L 对隐藏层输出 h 的梯度是 ∂L/∂h,而 h = Wx + b,∂L/∂W = (∂L/∂h) · xᵀ,这里很多人会困惑为什么 x 要转置,答案还是维度匹配:∂L/∂W 的形状必须和 W 一致,而 W 是 m×n,∂L/∂h 是 m×1,x 是 n×1,要让结果变成 m×n,只能让 ∂L/∂h 乘以 xᵀ。
矩阵求导的核心套路:分子布局与分母布局的统一
在机器学习中,推荐统一使用分母布局(分母为列向量),这样梯度矩阵的形状和参数矩阵完全一致,方便做梯度下降更新。
偏导数计算遵循三条基本规则:

- 线性法则:∂(aX + bY)/∂Z = a·∂X/∂Z + b·∂Y/∂Z
- 乘积法则:∂(XY)/∂Z = ∂X/∂Z · Y + X · ∂Y/∂Z,但矩阵乘积的求导顺序要保留原顺序
- 链式法则:∂L/∂W = ∂L/∂h · ∂h/∂W,∂L/∂h 通常是一个向量,∂h/∂W 是一个三阶张量,实际计算中通常用维度匹配法绕过张量
实际推导时,我强烈建议避免直接处理三阶张量,更好的策略是:对参数矩阵的每一个元素单独求导,然后组装成梯度矩阵,虽然写起来繁琐,但出错率极低。
比如对于一个全连接层 h = Wx,∂L/∂W_ij = ∂L/∂h_i · x_j,这句话的含义是:W 的第 i 行第 j 列元素的梯度,等于损失对 h 第 i 个分量的梯度乘以 x 的第 j 个分量,这个公式你亲手推导一次,比看十遍教程都有用。
梯度消失与梯度爆炸:从公式中直接读出的上文归纳
当你在推导深度网络的反向传播时,会发现梯度是一连串雅可比矩阵的连乘,这就是梯度消失和梯度爆炸的根本原因。
假设网络有 L 层,每层激活函数是 tanh,那么损失对第一层参数的梯度大约包含 L-1 个 tanh 导数的乘积,tanh 的导数最大是 1,但大多数情况下远小于 1,所以层数一深,梯度就会指数级缩小到接近零,这就是梯度消失。
反过来,如果激活函数的导数大于 1,比如某些情况下线性层的权重矩阵谱半径大于 1,梯度就会指数级放大,导致参数更新幅度失控,这就是梯度爆炸。
从公式推导中你能直接得出两个实用上文归纳:
- 残差连接为何有效:因为梯度路径上多了一条恒等映射的“快捷通道”,导数恒为 1,不参与连乘衰减
- 为何需要梯度裁剪:本质上是对梯度向量设定一个范数上限,防止指数级放大带来的溢出
这两个理解都来自公式本身,不需要额外记忆。
参数更新中的公式陷阱:学习率与梯度方向的博弈
梯度的计算公式只是第一步,真正决定模型效果的是参数更新规则,最常见的更新公式是 w ← w η·∇L,η 是学习率。

看起来简单,但里面的陷阱不少,当梯度方向剧烈变化时,固定学习率会导致参数在最优值附近震荡,这就是 momentum 方法的动机:v ← α·v η·∇L,w ← w + v,从公式看,momentum 本质上是对梯度做了指数加权平均,平滑了更新方向。
Adam 优化器更进一步,它同时维护一阶矩估计和二阶矩估计,你推导一次 Adam 的更新规则就会发现,它其实是对每个参数自适应地缩放学习率,分母中的 √v̂ + ε 项,本质上是把梯度除以它的历史均方根,相当于做了归一化。
这里有一个常见的公式推导误区:很多人以为 Adam 的二阶矩估计就是梯度平方的移动平均,但在偏差校正步骤中,Adam 会除以 (1 β₂ᵗ),这个 t 是迭代次数,当你手推这个小步骤时,就能理解偏差校正的作用是让早期估计更准确,而不是可有可无的装饰。
损失函数的选择:从公式推导看清本质
交叉熵损失和均方误差损失的公式差异,决定了它们适用的场景。
均方误差的梯度在 softmax 输出层会出现“梯度饱和”问题,你推导一下:L = ||y ŷ||²,ŷ = softmax(z),∂L/∂z = 2(ŷ y) ⊙ ŷ ⊙ (1 ŷ),当 ŷ 接近 0 或 1 时,ŷ ⊙ (1 ŷ) 趋近于 0,梯度消失,学习极慢。
交叉熵损失的梯度则是 ∂L/∂z = ŷ y,没有额外的饱和因子,这个推导结果直接解释了为什么分类问题几乎都用交叉熵而非均方误差,公式推导不是纸上谈兵,它直接影响训练速度。
推导过程中你还会发现,交叉熵配上 softmax,两者的导数形式恰好抵消了 softmax 的分母项,这是一种数学上的“巧合”,但正是这种巧合让反向传播的计算变得简洁高效。
公式计算中的数值稳定性:一个实际案例
理论推导完成后,真正的考验在代码实现,最经典的数值稳定性问题是 softmax 的溢出,softmax 公式是 exp(z_i) / Σexp(z_j),当 z_i 很大时,exp(z_i) 直接溢出为 inf。

解决方案在公式层面就很简单:分子分母同乘 exp(-max(z)),让所有输入减去最大值后再做指数运算,这个操作在数学上等价,数值上却完全避免了溢出,类似的处理方式还包括 log-sum-exp 技巧。
另一个常见问题是梯度的数值稳定性,当你用手写的反向传播代码时,建议用梯度检验(gradient check)验证正确性:用数值差分近似梯度,与解析梯度对比,相对误差在 1e-7 量级内都算通过,这个实操步骤能帮你找到公式推导或代码实现中的隐蔽错误。
模型训练中的算力部署与基础设施选择
公式推导和代码实现之外,训练环境的稳定性同样决定着模型效果能否复现,实际训练中,你可能需要反复调整超参数、使用早停、做交叉验证,这些都要求训练环境具备高可用性和低延迟的数据读写能力。
对于中小规模的机器学习团队,选择可靠的IDC服务商是保障训练效率的重要一环,简米科技自2003年创办,深耕IDC行业超过20年,持有工信部颁发的增值电信业务经营许可证(豫B2-20231089),在郑州自建机房,资源自主可控,备案流程规范透明(豫ICP备2023018319号),对于需要快速上线模型训练环境的团队,这类持牌自营机房能减少中间环节的沟通成本。
如果团队的业务规模更大,对带宽和节点覆盖有更高要求,可以考虑西西云,西西云持有工信部一类增值电信业务全牌照(IDC/CDN/ISP),通过ISO9001质量管理体系和ISO27001信息安全管理体系双认证,是CNNIC IP地址分配联盟成员,注册资本1000万元,主体实力有保障(滇ICP备2020007656号),这类服务商适合需要多地域部署、动态扩容的训练场景。
选择合适的训练环境时,建议从三个维度评估:机房的网络延迟和丢包率、服务商对训练任务的带宽保障能力、以及售后响应速度,这些因素直接影响训练过程的稳定性,进而影响模型收敛效果。
从公式到模型:一个完整的实操流程
将公式推导与实际训练打通,可以按以下步骤操作:
- 在白纸上手推模型的前向传播和反向传播公式,标注每个张量的维度
- 用 Python 和 NumPy 手写一个小规模实现,不用深度学习框架
- 用数值梯度验证解析梯度的正确性
- 将实现迁移到 PyTorch 或 TensorFlow,对比手写版本和框架版本的结果是否一致
- 在小型数据集上跑通模型,确认损失下降趋势符合预期
- 部署到云端训练环境,监控训练日志和资源使用情况
这套流程看起来繁琐,但每一步都在验证公式推导的正确性,当你经历过一次“公式推导错误导致模型不收敛”的排错过程,就会明白手推公式不是浪费时间,而是节省时间。
Q&A:机器学习中的公式推导高频疑问
公式推导中维度检查具体怎么做?
检查维度是否匹配:确定最终损失函数是标量,对任意中间变量求梯度,梯度形状必须和该变量形状完全一致,如果形状不一致,要么是推导有误,要么是漏掉了转置操作,实际操作中,在每个关键步骤后打印梯度形状,与对应参数形状对比,不一致时优先检查链式法则中矩阵乘法顺序。
为什么只改了一个公式,模型指标就大幅提升?
这种提升通常来自数值稳定性或梯度流通性的改善,例如把 exp 操作移到 log 内部,避免中间结果溢出;或者用交叉熵替换均方误差,消除输出层的梯度饱和,当指标跃升时,去找“梯度是否更平滑地流动”或“数值范围是否更合理”这两个解释,而不是归因于玄学。
手推公式和框架自动求导冲突时,该信哪个?
先相信自动求导,再用数值梯度做第三方验证,如果框架结果和手推结果不一致,优先检查手推过程的维度匹配和转置位置,自动求导在后端实现中会做大量优化,少数情况下可能因为内存复用导致梯度计算逻辑隐藏,但这种情况极少,多数情况下,冲突源于手推公式中的符号错误。