尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

大模型训练原理(五)|Loss 明明只是一个数字,凭什么能训练几十亿参数?真正看懂 Backpropagation

大模型训练原理(五)|Loss 明明只是一个数字,凭什么能训练几十亿参数?真正看懂 Backpropagation 上一课讲完 Cross Entropy交叉熵以后我们终于让模型拥有了一个评价自己预测好坏的数字。假设某一个位置真实 Token 的概率只有P(y)0.01那么L-log 0.01≈4.605很好。模型现在知道这次预测很差。但如果你真的停下来想两秒会发现这里有一个非常大的漏洞。4.605然后呢4.605 只是一个数字。而一个大语言模型里面可能有几十亿、几百亿甚至更多参数。Loss 又没有附带一张说明书告诉模型第 18,391,024 个参数加一点。第 72,381,552 个参数减一点。这个 Attention 权重改多一点。那个 MLP 权重少改一点。更前面的 Embedding 也有责任。所以真正神奇的问题根本不是模型怎么知道自己错了第四课已经解决了。真正神奇的是模型知道自己错了以后究竟怎么知道“应该改哪里”这才是 Backpropagation反向传播真正要解决的问题。而一旦这件事想明白你会突然发现所谓“训练神经网络”其实没有想象中那么玄学。整件事情可以浓缩成一句话Loss 负责告诉模型“现在有多差”Gradient 负责告诉模型“附近往哪里改Loss 会下降”。Backpropagation 做的就是把前者变成后者。先把几十亿参数忘掉我们只训练一个数字先别碰 Transformer。也先别碰矩阵。假设我们造了一个世界上最简单的模型ŷwx其中x 是输入。w 是模型参数。ŷ 是模型预测。现在给模型一个输入x2当前模型参数w3所以预测ŷ3×26但是训练数据告诉我们真实答案其实是y10模型预测 6正确答案 10。显然预测得不好。为了让问题简单一点这里暂时不用 Cross Entropy而是使用一个很常见的平方误差L½(ŷ-y)2代进去L½(6-10)28于是现在Loss8到这里和上一课的处境完全一样。模型知道自己不好。但是怎么改最关键的问题不是“Loss 是多少”而是“动一下参数会发生什么”我们偷偷试一下。现在w3如果把它稍微增大一点w3.01新的预测ŷ3.01×26.02新的 LossL½(6.02-10)2大约是7.9202原来的 Loss 是8现在变成7.9202Loss 降了。这说明什么说明在当前这个位置把 w 往大的方向推一点模型会变好。这其实已经非常接近 Gradient梯度的本质了。我们真正想知道的不是L8而是(∂ L)/(∂ w)这个东西在问如果我现在让参数 w 发生一个非常小的变化Loss 会往哪个方向变化变化有多快这就是 Derivative导数在机器学习里真正有价值的地方。它不是为了考试求导。也不是为了证明你学过高等数学。它本质上是在测量Loss 对某一个变量有多敏感。为什么 Gradient 是负数反而意味着参数应该变大我们直接把刚才的例子算完。现在L½(ŷ-y)2所以(∂ L)/(∂ŷ) ŷ-y当前ŷ6y10于是(∂ L)/(∂ŷ) 6-10-4第一次看到-4很多人会懵。负数到底是什么意思其实非常简单。它告诉你当前如果把 ŷ 稍微增大一点Loss 会下降。为什么因为正确答案是 10。你现在才预测 6。当然应该往 10 靠。所以这里的负号不是在说“预测应该变成负数。”而是在告诉我们Loss 下降的方向在另一边。但问题来了。我们不能直接修改ŷŷ 只是模型计算出来的结果。真正可以训练的是w所以我们必须继续往前找。Chain Rule神经网络真正赖以生存的一条数学规则现在整个计算过程其实是w→ŷ→L参数 w 先影响预测 ŷ。预测 ŷ 再影响 Loss。所以如果我们想知道(∂ L)/(∂ w)可以拆成两段。第一段w→ŷ因为ŷwx所以(∂ŷ)/(∂ w)x而现在x2所以(∂ŷ)/(∂ w)2第二段ŷ→L我们刚才已经算过(∂ L)/(∂ŷ)-4现在把两段连接起来(∂ L)/(∂ w) (∂ L)/(∂ŷ) × (∂ŷ)/(∂ w)所以(∂ L)/(∂ w) (-4)×2 -8于是(∂ L)/(∂ w)-8这就是参数 w 当前的 Gradient。而我们刚才用 (w3.01) 做的小实验其实也已经偷偷验证了这个结果。参数增加0.01Loss 大约减少0.0798两者相除−0.0798 / 0.01≈-7.98已经非常接近真正的导数-8这就是导数。你完全可以把它理解成把参数轻轻碰一下看看 Loss 会怎么动。只不过数学让我们不需要真的一次一次试。Chain Rule 真正厉害的地方不是公式而是“影响可以接力”刚才w→ŷ→L只有两步。如果是a→b→c→d→L怎么办完全一样。(∂ L)/(∂ a) (∂ L)/(∂ d) × (∂ d)/(∂ c) × (∂ c)/(∂ b) × (∂ b)/(∂ a)第一次看这个公式可能很烦。但其实它只表达一件事一个东西不需要直接影响 Loss。只要它能影响下一个变量下一个变量再继续影响后面的变量那么这种影响就可以一路传递下去。我更喜欢把 Chain Rule链式法则理解成一种影响力换算。比如参数 a 改 1 点会让 b 改多少b 改 1 点会让 c 改多少c 又会影响多少 Loss把这些“换算率”一路乘起来你就知道参数 a 最终对 Loss 有多大影响。到这里Backpropagation 已经出现了一半。神经网络为什么刚好特别适合 Chain Rule因为神经网络本质上就是一条非常长的函数链。例如一个极度简化的网络x→Layer1→Layer2→Layer3→Logits→Loss更真实一点Token→Embedding→Attention→MLP→Attention→MLP→…→LM Head→Logits→Softmax→LossForward Pass前向传播时信息从左往右走。输入进来。一层层计算。最后得到 Loss。但当我们想知道最前面的某个参数到底该怎么改时我们可以反过来从 Loss 出发沿着原来的计算关系一路往回求导。这就是Backpropagation反向传播。但“反向传播”这个名字其实很容易让人误会很多人第一次学的时候脑子里会出现一个画面Loss 从网络最后面出发像一股液体一样一层一层倒着流回网络前面。这个比喻只能算对了一半。真正往回传播的并不是Loss不是说最后 Loss 是 4.605那最后一层分 1.8倒数第二层分 1.2Attention 分 0.6Embedding 再承担一点。神经网络没有这种“责任分账”。真正往回传播的是(∂ L)/(∂ h)(∂ L)/(∂ z)(∂ L)/(∂ W)这些 Gradient Signal梯度信号。也就是最终 Loss 对当前这个中间变量有多敏感。所以所谓“把责任往前传”更准确的理解应该是把 Loss 对后面变量的敏感度一步一步换算成 Loss 对前面变量的敏感度。这就是 Backpropagation 的本质。Computational Graph几十亿参数为什么没有把问题复杂几十亿倍现在再看刚才那个简单模型ŷwxL½(ŷ-y)2其实可以拆成几个非常小的计算。先w× x得到ŷ再ŷ-y得到误差e最后½e2得到L这就是一个非常小的 Computational Graph计算图。关键来了。计算图里的每一个节点其实根本不需要理解整个神经网络。乘法节点只需要知道乘法怎么求导。平方节点只需要知道平方怎么求导。Softmax 节点只需要知道Softmax 怎么求导。Matrix Multiplication矩阵乘法只需要知道矩阵乘法自己的局部导数是什么。也就是说一个巨大神经网络可以被拆成大量非常简单的局部计算。Forward 时每个节点完成自己的计算。Backward 时每个节点收到后面传回来的 Gradient再结合自己的 Local Derivative局部导数算出应该继续传给前面的 Gradient。每个人只负责自己这一小段。但一段一段连起来以后最终就能算出Loss 对整个网络所有参数的 Gradient。这件事非常漂亮。为什么 Backpropagation 一定要“从后往前”这里有个很容易被忽略的问题。我们不是已经有 Chain Rule 了吗那对每一个参数单独算不就好了理论上可以。工程上会浪费得非常夸张。假设计算图里有a→b→c→L同时d→b→c→L现在你想分别算(∂ L)/(∂ a)和(∂ L)/(∂ d)两者后面的路径b→c→L其实完全一样。如果每个参数都从头单独求一遍同样的东西会被重复算无数次。Backpropagation 的做法恰恰相反。先计算(∂ L)/(∂ c)再得到(∂ L)/(∂ b)然后前面的a和d都可以复用已经得到的结果。所以 Backpropagation 不只是“使用 Chain Rule”。更完整一点应该说Backpropagation 是在 Computational Graph 上从 Loss 开始反向复用中间 Gradient高效计算所有参数偏导数的方法。这也是为什么它特别适合神经网络。神经网络的特点正是一个最终 Loss前面连着海量参数。有了 Gradient参数到底怎么动回到最开始那个例子。我们已经算出(∂ L)/(∂ w)-8如果我们要最小化 Loss就朝 Gradient 的反方向走wnew wold - η (∂ L)/(∂ w)其中η叫 Learning Rate学习率。假设η0.1那么wnew 3-0.1×(-8)得到wnew3.8新的预测ŷ3.8×27.6原来预测6现在变成7.6离正确答案10更近了。原来的 Loss8新的 Loss½(7.6-10)22.88真的下降了。所以(∂ L)/(∂ w)-8不是一个抽象数学结果。它真的告诉了我们这个参数附近往哪个方向走会让模型变好。只有一个参数叫导数几十亿参数以后就叫 Gradient如果只有Lf(w)我们关心(dL)/(dw)这是 Derivative导数。但真实模型显然有很多参数。比如Lf(w1,w2,w3,…,wn)那么每一个参数都有自己的偏导数(∂ L)/(∂ w1)(∂ L)/(∂ w2)…(∂ L)/(∂ wn)把它们放在一起∇θ L (∂ L)/(∂θ1)(∂ L)/(∂θ2)⋮(∂ L)/(∂θn)这整个东西就是Gradient梯度。所以你可以把 Gradient 理解成当前模型所有参数各自收到的一份“局部修改方向报告”。几十亿参数就有几十亿个局部偏导。到这里还只是普通神经网络现在回到大语言模型真正精彩的地方来了。第四课最后留下过一个非常漂亮的结果(∂ L)/(∂ zi) qi-yi当时我们没有真正展开它。现在终于可以看懂这条公式到底有多重要。假设模型最后得到 Vocabulary 上的 Logitsz1,z2,…,zV经过 Softmaxqi (ezi)/(Σj ezj)得到模型预测的 Probability Distribution概率分布。真实答案则可以写成 One-hot Target独热目标y1,y2,…,yVCross EntropyL-Σi yilog qi如果正确 Token 是第 k 个那么只有yk1其余yi0于是L-log qk把 Softmax 展开L -log (ezk)/(Σj ezj)可以写成L -zklogΣj ezj然后对任意一个 Logit zi 求导。最终得到(∂ L)/(∂ zi)qi-yi这一条公式值得停下来认真看。因为左边是Gradient。右边是模型预测分布减去真实目标分布。换句话说模型在 Probability Distribution 上犯的错误直接变成了 Logit 层的 Gradient。这就是第四课和第五课真正接上的地方。用一个数字例子你会马上看懂 q-y 在干什么假设 Vocabulary 只有五个 Token。模型预测q [0.70, 0.15, 0.08, 0.05, 0.02]而真实答案是第一个 Token。所以y [1, 0, 0, 0, 0]那么q-y [-0.30, 0.15, 0.08, 0.05, 0.02]先看正确 Token。它当前只有70%但 Target 是100%所以0.70-1-0.30Gradient 是负的。Gradient Descent梯度下降更新时会减去这个负数所以正确 Token 对应的 Logit 会被往上推。再看错误 Token A0.15-00.15Gradient 是正的。更新时会把它的 Logit往下压。错误 Token B0.08也往下压。但力度比 0.15 小。错误 Token D 只有0.02所以只需要很小的修正。这里出现了一个非常漂亮的结果错误 Token 当前抢走的概率越多它收到的纠正信号就越强。这不是人为写了几十条规则。它自动从Softmax Cross Entropy里面长出来了。再看一个极端例子你会更容易产生“原来如此”的感觉假设模型预测q [0.97, 0.01, 0.01, 0.01]但正确答案其实是第二个 Token。所以y [0, 1, 0, 0]那么q-y [0.97, -0.99, 0.01, 0.01]第一项0.97意思非常直接模型极度自信地把概率压在了错误 Token 上。所以它收到极强的向下修正。第二项-0.99它才是真实答案模型却只给了 1%。所以它收到极强的向上修正。第三、第四项只有0.01它们本来就没抢多少概率所以只需要轻微压低。这时候你应该能真正看懂第四课那句话Cross Entropy 不是只负责给模型“打一个分”。它真正重要的地方在于它把 Probability Distribution 上的错误变成了可以继续向网络内部传播的 Gradient Signal。但 Logits 也不是参数训练信号怎么继续往前现在我们已经得到(∂ L)/(∂ z) q-y但z只是 Logits。还不是模型真正需要训练的参数。假设最后的 LM Head语言模型输出层可以简化写成zWhb这里h 是 Transformer 最后产生的 Hidden State隐藏状态。W 是 LM Head 的权重矩阵。b 是 Bias偏置。z 是最终 Logits。现在我们已经知道(∂ L)/(∂ z) q-y于是 Chain Rule 再次登场。对于 W(∂ L)/(∂ W) (q-y)hT对于 Bias(∂ L)/(∂ b) q-y而真正关键的是(∂ L)/(∂ h) WT(q-y)为什么第三条最重要因为它告诉我们Gradient 没有停在 LM Head。LM Head 在计算自己参数 Gradient 的同时还会继续产生(∂ L)/(∂ h)告诉前面的 Transformer你刚才产生的 Hidden State也需要调整。于是训练信号继续往前。从这里开始几十亿参数真的都能收到训练信号了假设整个 Transformer 简化成h0→Block1→h1h1→Block2→h2一路到hL-1→BlockL→hL最后hL→LM Head→Logits→LossForward Pass 时从h0一路算到LossBackward Pass 时先得到(∂ L)/(∂ z)再得到(∂ L)/(∂ hL)然后进入最后一个 Transformer Block。最后一个 Block 内部可能包含Attention、MLP、Normalization、Residual Connection。每个部分都根据自己的局部导数继续应用 Chain Rule。于是得到这一层 Attention 参数的 Gradient。这一层 MLP 参数的 Gradient。同时继续得到(∂ L)/(∂ hL-1)然后再进入前一层。再前一层。继续。一直回到网络最前面。所以 Backpropagation 真正做的事情可以理解成Loss先变成(∂ L)/(∂ Logits)再变成(∂ L)/(∂ Hidden State)然后不断变成(∂ L)/(∂ Weight)以及更前面状态的(∂ L)/(∂ Hidden State)最终得到∇θ L一个地方如果有两条路Gradient 怎么办这里还有一个非常重要的规则。假设变量 x 同时走了两条计算路径x→A→L和x→B→L那么x 对 Loss 的最终影响不是只看其中一条。而是两条路径产生的 Gradient 加起来。也就是说(∂ L)/(∂ x) 路径 A 的贡献 路径 B 的贡献为什么这一点重要因为 Transformer 里到处都是分支。最典型的就是 Residual Connection残差连接yxF(x)Forward 时x 有一条路直接进入 y。另一条路经过F(x)再进入 y。Backward 时也是如此。Gradient 一条可以沿 Residual 的直接路径回来。另一条经过F回来。最后两边相加。以后我们正式讲 Residual 为什么能帮助深层网络训练时这一点会重新变得非常重要。一条 Sequence 里几千个 TokenGradient 又是怎么处理的真实语言模型当然不是一次只训练一个 Token。假设一条 Sequence 有T个参与训练的位置。每个位置都有自己的 Token LossLt -log Pθ(xt|xt)如果最终 Loss 是它们的平均L 1/T Σt1TLt那么根据导数的线性性质∇θ L 1/T Σt1T ∇θ Lt换句话说每一个 Token Position 都会产生自己的训练信号。而同一组模型参数会被整个 Sequence 反复使用。所以某一个参数最后拿到的 Gradient实际上可能同时汇集了第 1 个 Token 的贡献第 2 个 Token 的贡献第 3 个 Token 的贡献……第 4000 个 Token 的贡献。再加上 Batch 里其他 Sequence 的贡献。所以一次真实的大模型训练可以粗略理解成大量 Token Position不断产生 Error Signal↓沿 Computational Graph 反向传播↓在共享参数处不断汇总↓最后形成这一批数据对应的∇θ L这才是真正的“大规模训练信号”。为什么训练时要保存那么多中间结果现在顺便可以解释一个实际问题。Forward 的时候x→h1→h2→h3→L为什么训练程序不能算完 h1 就立刻把它彻底扔掉因为 Backward 时可能还需要它。例如ywx要计算(∂ y)/(∂ w)需要知道x很多 Activation Function激活函数在求导时也需要 Forward 阶段产生的中间 Activation激活值。所以训练时除了 Parameters参数以外还需要保存大量中间状态。这也是为什么训练大模型时显存里不只有模型参数。还会有Activations激活值、Gradients梯度以及后面第六课会出现的 Optimizer States优化器状态。这也是 Training训练和 Inference推理在资源结构上一个非常大的区别。Backpropagation 不是 Parameter Update这两个千万别混这是第五课最容易留下的错误之一。Backpropagation 做的是L→∇θ L也就是把 Gradient 算出来。它本身并没有决定参数最后怎么更新。真正的参数更新是 Optimizer优化器的工作。最简单的情况可能是θt1 θt-η∇θ L但现实中的 Optimizer 会进一步考虑很多问题。比如过去几步的 Gradient 要不要参考不同参数是不是应该用不同的有效步长Learning Rate 多大Gradient 一直震荡怎么办Weight Decay权重衰减怎么处理所以Backpropagation 回答的是往哪里走Optimizer 回答的是到底怎么走这是两件事。Gradient 也不是“这个参数的重要程度”还有一个非常常见的误解。假设某个参数(∂ L)/(∂ w) 0.000001能不能说“这个参数不重要”不能。Gradient 表达的是在当前模型、当前数据、当前 Objective目标函数、当前参数位置附近这个参数稍微变化时对当前 Loss 的一阶影响。换一条训练样本Gradient 可以变。模型更新一步Gradient 也可以变。甚至换一个 ObjectiveGradient 还会变。所以 Gradient 是局部的。动态的。和当前训练目标有关的。它绝不是一个永久的“参数价值排行榜”。为什么深度学习里总强调 Differentiable现在你应该已经能够自己回答这个问题。Backpropagation 想从L一路走到∇θ L中间就必须不断使用 Chain Rule。所以网络里的计算最好都能够提供自己的 Local Derivative。Matrix Multiplication 可以。Softmax 可以。Attention 可以。MLP 可以。Normalization 可以。于是 Gradient 可以一路往回传播。但以后你会碰到一些很麻烦的东西Sampling采样、Argmax、工具调用、代码执行、外部环境反馈。这些操作不一定还能直接放进一条漂亮的端到端可微计算图。这时候普通 Backpropagation 就开始不够用了。为什么后面还会出现Reinforcement Learning强化学习、Policy Gradient策略梯度、Reward奖励其实伏笔已经埋在这里了。到这里再重新看一次大模型训练你看到的应该已经完全不一样第一课我们只知道一条抽象流程Data→Model→Prediction→Loss→Gradient→Parameter Update第二课把 Prediction 打开Pθ(xt|xt)第三课把 Probability Distribution 打开Context→Transformer→Logits→Softmax→Probability Distribution第四课继续Probability Ground Truth→Cross Entropy→Loss这一课终于把Loss→Gradient打开了。现在整条链已经变成Text↓Token↓Context↓Transformer↓Logits↓Softmax↓Probability Distribution↓Cross Entropy↓Loss↓Backpropagation↓∇θ L到这里模型终于不只是知道“我错了。”它开始知道“如果想让下一次更好当前每一个参数附近应该往哪个方向调整。”真正理解 Backpropagation只需要抓住四件事如果这一课公式很多看完以后有一点乱可以只留下四件事。第一件Loss 不是修改方案。L4.605只能告诉你模型现在表现不好。真正能指导参数变化的是∇θ L第二件Derivative 本质上是在测量局部敏感度。(∂ L)/(∂ w)问的只是参数 w 稍微变化一点Loss 会怎么变第三件Chain Rule 让远处的参数也能知道自己对 Loss 的影响。参数不需要直接连接 Loss。只要它参与了一条最终影响 Loss 的计算路径Gradient 就可以一路算回来。第四件Backpropagation 是在 Computational Graph 上高效执行 Chain Rule。真正向后传播的不是 Loss 数字而是Gradient Signal。最值得记住的其实是 q-y如果让我从第五课只挑一条和语言模型最相关的公式我不会先让你背θt1 θt-η∇θ L而是这一条(∂ L)/(∂ z)q-y为什么因为它刚好站在两个世界的交界线上。左边(∂ L)/(∂ z)属于 Gradient 的世界。右边q-y属于 Probability Distribution 的世界。模型预测了什么真实答案是什么两者之间的差异就在这里第一次变成了可以真正进入神经网络内部的训练信号。然后这个信号经过 LM Head。经过最后一个 Transformer Block。经过前一层。再经过前一层。不断通过 Chain Rule 往前传播。最终形成∇θ L几十亿参数各自得到自己的 Gradient。所以一个只有4.605这样的 Loss 数字最终真的可以改变一个几十亿参数的大语言模型。这就是 Backpropagation。但训练到这里其实还没有结束现在我们终于拥有∇θ L看起来似乎马上就可以θ ← θ-η∇θ L然后结束。可现实马上会冒出一堆新问题。Learning Rate 到底多大为什么参数更新会来回震荡为什么要记住过去的 Gradient为什么不同参数需要不同的更新尺度Momentum动量到底在解决什么Adam 为什么会成为深度学习里如此常见的 OptimizerAdamW 又为什么要专门处理 Weight Decay所以第五课真正完成的是Loss→Gradient而第六课才真正进入Gradient→Parameter UpdateBackpropagation 已经告诉模型附近哪里是下坡。下一课真正的问题是知道下坡方向以后到底应该怎么走才能又快、又稳、还不把训练走崩这就是 Optimizer优化器。
返回列表