大模型推理前向传播全解:从QKV计算到KV Cache优化
2026/9/8 19:34:16 网站建设 项目流程

1. 大模型推理的心跳:从一句话到下一个 Token

这两年大模型火得一塌糊涂,但说实话,大多数人对大模型的认识都停留在“输入一句话,输出一段文字”这种黑盒层面。哪怕是在做 AI 开发的工程师,真正把 Transformer 前向传播完整啃下来的人,也不算多。很多人一上来就追着 FlashAttention、量化、KV Cache 这些推理优化手段跑,结果源码一看就懵:QKV 是从哪来的?Mask 到底怎么加的?为什么推理时只算一个 Token?

这篇博文我想把“大模型推理的前向传播”这件事从头到尾拆开讲清楚。内容主线锁定 Transformer 架构和注意力机制,我从最底层的矩阵运算讲起,一直推到完整的自回归解码流程,中间会把多头注意力、位置编码、残差连接、LayerNorm、FFN 这些模块全部过一遍。

这个内容适合谁看?我觉得有三类人收获最大。第一类是准备面试大模型岗的算法工程师或者应届生,面试官特别喜欢从“你讲讲 QKV 的计算流程”这种问题入手,然后一路追问到 KV Cache 的原理;第二类是已经开始跑推理服务,但遇到性能瓶颈不知道怎么定位问题的部署工程师;第三类就是纯粹想搞清楚“大模型内部到底是怎么工作的”的爱好者。

我尽量用直白的话把每个计算环节讲透,所有步骤都会落到具体的矩阵形状上,而不是飘在概念层面。搞懂前向传播之后,你再去看各种推理加速方案,会发现很多优化技巧其实就是在和前向传播的某个瓶颈死磕。

2. 前向传播的总览:一个 Token 的冒险旅程

2.1 输入长什么样:Token 序列到 Embedding 矩阵

要理解 Transformer 的前向传播,第一步得搞清楚输入到底是什么形状。假设我们现在用的是 GPT 类的自回归模型,输入一句话,比如“人工智能正在改变世界”,这句话进来之后会经历两次转换。

第一次转换是 Tokenization,也就是分词。中文分词现在主流走的是 BPE 或者 Unigram 方案,模型有一个词表,假设词表大小是 50257(GPT-2 的标准配置),这句话会被切成若干个 Token。每个 Token 对应词表里的一个整数 ID,比如“人工”对应 1024,“智能”对应 2048,等等。假设这句话最终被切成了 6 个 Token,那我们就拿到了一个形状为 (6,) 的整数向量。

第二次转换是 Embedding,也就是查表。模型的参数里有一个嵌入矩阵 W_e,形状是 (vocab_size, hidden_size),其中 hidden_size 我们一般简写为 d_model。GPT-2 的 d_model 是 768,GPT-3 是 12288,7B 级别的模型通常用 4096。我们把 (6,) 的整数向量通过查表映射,得到一个形状为 (6, 4096) 的矩阵,这个就是模型内部真正处理的输入 X。

这一步有几个细节容易踩坑。第一,Embedding 层是有参数的,而且参数量不小(词表大小乘以隐藏维度),这部分参数在大模型里通常占据总参数的很大一块,所以有些优化方案会做 Embedding 层权重共享(比如 GPT-2 就共享了输入和输出的 Embedding 权重)。第二,输入序列的长度是动态变化的,训练时一般用固定长度(比如 2048 或者 4096),但推理时是一步一步增长的,这也是为什么推理前向传播和训练前向传播在实现上有区别。

2.2 Transformer Block 的内部结构:一个完整的处理流水线

拿到 (batch_size, seq_len, d_model) 的输入矩阵之后,数据会依次穿过很多个相同的 Transformer Block。一个 Block 内部包含两个大的子层:第一个是多头自注意力模块(Multi-Head Self-Attention,MHSA),第二个是前馈神经网络模块(Feed-Forward Network,FFN)。

每个子层外面都套着残差连接和 LayerNorm,这个设计是整个 Transformer 能训练得动、训得深的核心。标准 GPT 架构用的是 Post-LN,也就是“残差相加之后再归一化”;但大模型时代很多模型(比如 GPT-3)其实用的是 Pre-LN,也就是“先归一化再进子层”。这两种排列方式在实践中差别很大,Pre-LN 训练更稳定,对学习率不那么敏感,所以现在的开源大模型基本都走 Pre-LN 路线。

一个 7B 规模的模型大约有 32 个这样的 Block,每个 Block 的参数包括:注意力模块里的 W_q、W_k、W_v、W_o 四个矩阵,加上 FFN 里的两个线性层(一般是先升维到 4 倍 d_model,再降回来),再加两套 LayerNorm 的 gamma 和 beta 参数。前向传播就是数据在这个流水线上按顺序流动一遍。

很多初学者会混淆“Transformer”和“GPT”这两个概念。严格来说,Transformer 是编码器-解码器架构,但 GPT 是只保留了解码器部分的变体,而且把原来的“编码器-解码器注意力”拿掉了,只保留自注意力层。这种设计使得 GPT 天然适合做自回归生成——每一步只预测下一个 Token。大模型推理的前向传播,本质上就是在跑一个纯解码器的自回归循环。

3. 注意力机制深度拆解:QKV 到底是什么

3.1 从一个直觉问题开始:怎么让模型知道“谁该关注谁”

我先用大白话把注意力的直觉讲清楚,因为很多人在矩阵公式里绕晕了,其实底层的想法非常简单。

假设你读一句话:“小明把球传给小李,因为他跑到了空位。”这里“他”指的是谁?人类能通过上下文推断出“他”大概率是“小李”。注意力机制要解决的就是这个问题——让模型在处理当前位置的时候,自动找出输入序列里哪些位置的信息更重要,然后把它们的向量按权重融合起来。

具体到数学上,Self-Attention 要做的事情是:给定一个序列的向量表示,计算两两位置之间的相关度(权重),然后用这个权重把其他位置的向量加权求和,得到每个位置的新向量。这个“相关度”就是注意力分数,加权求和的结果就是 Attention 输出。

3.2 Q、K、V 的计算过程:矩阵形状视角

Self-Attention 的输入是上一步得到的 X,形状为 (batch_size, seq_len, d_model)。每个 Block 的注意力层内部有四个可学习的参数矩阵:W_q、W_k、W_v、W_o,它们分别用于生成 Query(查询)、Key(键)、Value(值)和输出投影。

计算过程分三步。第一步,用 X 分别乘以 W_q、W_k、W_v,得到 Q、K、V 三个矩阵:

Q = X @ W_q # (batch, seq, d_model) @ (d_model, d_model) -> (batch, seq, d_model) K = X @ W_k # 同理 V = X @ W_v # 同理

第二步,计算注意力分数矩阵。注意力分数等于 Q 和 K 的转置做点积,再除以缩放因子根号 d_k,最后经过 softmax 归一化:

S = Q @ K^T / sqrt(d_k) # (batch, seq, seq) A = softmax(S, dim=-1) # 每一行和为 1

第三步,用注意力权重矩阵 A 对 V 做加权求和,得到 Attention 输出,再过一层输出投影:

O = A @ V # (batch, seq, d_model) Output = O @ W_o # (batch, seq, d_model)

这里有三个非常重要的细节。第一个是缩放因子 d_k。为什么要除根号 d_k?因为当维度比较大的时候,Q 和 K 的点积结果会很大,导致 softmax 的输入进入梯度饱和区,反向传播时梯度会变得非常小,训练不动。除以根号 d_k 就是把点积的方差拉回 1 附近,这个设计看起来只是一个小改动,但没了它 Transformer 根本训不动。

第二个是 softmax 是在最后一个维度上做的,也就是对每一个 query 位置,在所有 key 位置上做归一化。这个顺序不能搞错,否则注意力权重的含义就变了。

第三个是注意力权重矩阵的形状是 (batch, seq, seq)。这个矩阵就是所谓的“注意力图”,它显示了每个位置对哪些位置关注度更高。对长序列来说,这个矩阵的空间复杂度是 O(n²),这也是后面 FlashAttention 优化的核心对象。

3.3 多头注意力:不是一个注意力,是 h 个注意力并联

上面讲的其实是一个注意力头,但真正的 Transformer 用的是多头注意力(Multi-Head Attention)。所谓多头,就是把 d_model 维度的空间切成 h 份,每份 d_k = d_model / h 维,然后每一份独立执行自注意力计算,最后把 h 个头的输出拼接起来。

具体来说,原来的 W_q 形状是 (d_model, d_model),现在拆成 h 个 (d_model, d_k) 的矩阵。实际实现时通常还是用一个大的矩阵算完,再 reshape 分头,性能更好。以 GPT-2 为例,d_model = 768,h = 12,d_k = 64。7B 模型的典型配置是 d_model = 4096,h = 32,d_k = 128。

多头注意力的价值在于:不同的头可以关注不同的关系模式。有的头关注语法上的临近词,有的头关注长距离的指代关系,有的头关注位置信息。这些不同的关注模式拼接在一起,模型就能同时捕捉多种语义关系。

从纯计算角度说,多头注意力并不改变总计算量——把一个大矩阵乘法拆成 h 个小矩阵乘法,总的 FLOPs 基本一样。但它显著提升了模型的表达能力,也让注意力图变得可解释。你现在去看一些大模型的可视化工具,能明显看到不同的头关注的是完全不同的区域。

3.4 Mask 机制:训练和推理场景下的不同处理

注意力机制里有个特别关键的细节是 Mask。在 GPT 类模型的推理前向传播里,Mask 的处理方式直接决定了效率和实现复杂度。

先说说为什么需要 Mask。自回归模型的核心假设是:预测当前位置的时候,只能看到当前位置及其之前的信息,不能看到未来的信息。所以在计算注意力分数的时候,矩阵的右上三角部分必须被遮住,否则当前位置会“偷看”到后面位置的 Token,这就是因果掩码(Causal Mask)。

训练的时候实现方式是这样的:给注意力分数矩阵 S 的右上三角位置填上一个非常大的负数(比如 -1e9),这样 softmax 之后这些位置的权重就趋近于 0。注意,不是直接置 0,因为 softmax 的输入如果直接是 0,经过指数运算后还是会贡献权重;填大负数才能把指数结果压到无限接近 0。

推理的时候情况就完全不同了,而且这是大模型推理前向传播和普通 Transformer 前向传播最大的区别之一。推理时我们通常只生成一个 Token,输入序列已经完整地在内存里了,所以不需要重新计算整个注意力矩阵。只需要把新 Token 的 Query 拿出来,跟之前所有 Token 的 Key 做点积,再跟所有 Value 做加权和。这就是 KV Cache 的由来——之前的 K 和 V 可以直接缓存住,不用重新算。

因果掩码在推理时几乎不产生额外成本,因为新 Token 本来就和所有历史 Token 计算注意力,不存在“未来”位置。但是训练时,因果掩码是必须的,而且如果你用的是 FlashAttention 这类融合算子,Mask 的传递方式也有讲究,后面我会展开。

4. 从注意力输出到 Transformer Block 输出:完整的模块衔接

4.1 残差连接与 LayerNorm 的位置之争

注意力计算完之后,输出还不是直接传到下一层,而是走一条标准的“残差 + 归一化”路径。这一步看起来简单,但对整个模型的训练稳定性影响极大。

残差连接就是 Output = x + Attention_Output,也就是把输入直接跨越子层加到输出上。这样做的好处是梯度可以有一条直达的高速公路,反向传播的时候不会因为层数太深而消失。这就像城市里修了一条高架桥,不用在地面红绿灯路口(每一层的非线性变换)里折腾,直接能回去。

LayerNorm 是对每个 Token 的整个 d_model 维向量做归一化,把均值拉到 0、方差拉到 1,然后再用可学习的 gamma 和 beta 做线性变换。它的计算公式是:

LayerNorm(x) = gamma * (x - mean(x)) / sqrt(var(x) + eps) + beta

这里的 eps 是一个很小的常数,比如 1e-5,主要防止除零。gamma 和 beta 的初始值一般是 1 和 0,这保证了初始状态下 LayerNorm 几乎是恒等映射,不会破坏预训练好的分布。

关于 Post-LN 和 Pre-LN,我多说两句。原始 Transformer 论文用的是 Post-LN,就是先让数据经过子层,再残差相加,最后归一化。但这种结构在深层次模型上容易出现训练不稳定问题,梯度爆炸比较频繁。GPT-3 等大模型普遍改用 Pre-LN,就是先归一化再进子层。你可以简单理解:Pre-LN 让每层的输入分布更稳定,所以即使层数很深、学习率很大,也不容易崩。

4.2 FFN 层:大模型里最被低估的参数量大头

Transformer Block 的第二个子层是前馈神经网络。这是整个模型里参数量和计算量的大头,但很多人对它关注不够,注意力都放在注意力机制上了。

FFN 的结构极其简单:两个线性层夹一个非线性激活函数。标准实现是:

FFN(x) = W2 * GELU(W1 * x + b1) + b2

第一层 W1 把 d_model 维上升到 4 * d_model 维,第二层 W2 再降回来。GPT-2 的 d_model 是 768,所以 FFN 的中间维度是 3072。7B 模型一般是 4096 升到 11008(LLaMA 的结构)或者 14336(新一点的模型),每家改动不太一样。

FFN 的计算量怎么估算?每 Token 每个 Block 的 FLOPs 大约是 2 * d_model * (4 * d_model) * 2(两个线性层各乘 2),也就是 16 * d_model²。你可以对比一下注意力模块的 FLOPs,当序列长度比较短的时候,注意力是 O(n²) 而 FFN 是 O(n),所以短序列下 FFN 反而更耗算力;但序列一旦长了,注意力就会变成主导。这也是为什么很多推理优化工具会先优化注意力算子——长文本场景下注意力才是瓶颈。

激活函数方面,早期 Transformer 用 ReLU,后来 GPT-2 改成 GELU,LLaMA 系列用 SwiGLU。SwiGLU 的计算复杂度略高(多了一个门控分支),但效果确实更好,现在已经是开源大模型的主流选择。

4.3 位置编码:没有 RNN 的模型怎么感知顺序

Transformer 没有循环结构,理论上它对输入是“同时看到所有 Token”的,这带来了并行效率的巨大优势,但也带来了一个问题:模型无法感知 Token 的顺序。如果我们把“我打你”和“你打我”这两句话的所有 Token 顺序打乱,模型是完全区分不出来的,因为注意力计算对排列是等变的。

为了解决这个问题,Transformer 需要给每个 Token 加上位置信息。大模型时代主流的方案有两种。

第一种是绝对位置编码(Absolute Positional Encoding),原始 Transformer 论文用的是三角函数公式:

PE(pos, 2i) = sin(pos / 10000^(2i/d_model)) PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))

这种方法的好处是能泛化到训练时没见过的长度,因为三角函数是连续函数,任意位置都能算出来。但它也有明显的缺陷:模型无法直接感受相对位置关系,而且不同位置的位置向量差异会随着维度出现周期性。

第二种是旋转位置编码(Rotary Position Embedding,RoPE),这是 LLaMA 系列和很多现代大模型的选择。RoPE 的思路是在 Q 和 K 上做旋转操作,使得点积运算天然包含相对位置信息。具体来说,对 Q 和 K 的每两个相邻维度,按当前位置的角度做旋转变换。这样计算注意力分数的时候,Q_i 和 K_j 的点积结果里会自然携带 i-j 的相对位置信息。RoPE 的优势是训练时可以外推更长的序列,这也是为什么很多模型声称“支持 128K 上下文”的基础。

第三种是 ALiBi(Attention with Linear Biases),它在注意力分数上直接加一个线性偏置项,偏置的大小和 token 之间的相对距离成正比。ALiBi 的好处是完全不需要额外的位置编码参数,外推性好,早期一些长上下文模型用过,但现在被 RoPE 取代得比较多。

4.4 一个 Block 的完整前向传播伪代码

把上面的所有组件串起来,一个 Transformer Block 的前向传播逻辑可以写成下面这段伪代码。这里的输入 x 是形状为 (batch, seq_len, d_model) 的张量:

def transformer_block(x, params): # Pre-LN 残差注意力子层 h = layer_norm(x, params['ln1']) q = h @ params['wq'] k = h @ params['wk'] v = h @ params['wv'] attn_out = attention(q, k, v, mask) x = x + attn_out @ params['wo'] # Pre-LN 残差FFN子层 h = layer_norm(x, params['ln2']) ff_out = params['w2'] @ gelu(params['w1'] @ h + params['b1']) + params['b2'] x = x + ff_out return x

就这几行代码,堆叠 32 层,最终再加一个输出层,就是一个能做文本生成的大模型。全部的智能,就藏在这套“先归一化、再变换、再残差”的流水线里。

5. 大模型推理的完整前向传播:从输入到输出一个 Token

5.1 预填充阶段 vs 解码阶段:为什么要分开看

大模型推理的前向传播和训练前向传播最大的区别在于:推理是自回归的,一次只生成一个 Token,生成的 Token 又会拼接到输入序列后面,继续参与下一轮预测。

这个过程分两个阶段。第一个阶段叫预填充(Prefill),就是将用户输入的整个 Prompt 一次性跑一遍前向传播,计算出整个序列的中间状态,并且把每一层的 K 和 V 缓存下来。第二个阶段叫解码(Decode),每生成一个 Token 就执行一次前向传播,这次只计算新 Token 对应的结果,充分利用 KV Cache 跳过历史 Token 的重复计算。

为什么要这样设计?核心原因就是效率。假设用户输入了 100 个 Token,要生成 100 个 Token。如果每生成一个 Token 就把整个序列从第一层重新算一遍,总计算量是 O(200²) 级别的注意力计算;而采用 KV Cache,预填充阶段算一次 O(100²),解码阶段每次只算 O(1),总计算量大幅降低。现在的推理框架全都采用这个策略,但实现细节千差万别,这也是为什么有些框架快、有些框架慢。

5.2 解码阶段的每一步:详细推演一个 Token 的诞生过程

我们来完整推演一下解码阶段的一次前向传播。假设当前序列长度是 L(包括已生成的 Token),我生成一个 Token 需要做什么?

第一步,把当前最新的 Token ID 转成向量。注意,这里只转一个 Token,不是整个序列——因为之前的 Token 对应的 Embedding 结果已经算过并且不需要保留了(只要保留 KV Cache)。

第二步,把这个 Token 的向量作为输入,送入第一个 Transformer Block。在 Block 里,先做 LayerNorm,然后生成新的 Q、K、V。新的 Q、K、V 的形状都是 (batch, 1, d_model),其中 Q 的 seq 维度只有 1。

第三步,把新的 K 和 V 追加到 KV Cache 里。这时候 KV Cache 的形状是 (batch, L+1, d_model)。然后拿新的 Q(形状 (batch, 1, d_model))和更新后的完整 K(形状 (batch, L+1, d_model))做注意力计算。这里 Q 的 seq 维度是 1,K 的 seq 维度是 L+1,所以注意力分数矩阵的形状是 (batch, 1, L+1)——这就是为什么解码阶段不需要计算完整的注意力矩阵,只需要计算一行。

第四步,在 L+1 个注意力权重上做加权求和,得到 (batch, 1, d_model) 的注意力输出。经过输出投影、残差、LayerNorm、FFN,得到这个 Block 的最终输出,然后传给下一个 Block。这个过程在 32 层里依次执行。

第五步,最后一个 Block 的输出送入 LM Head(语言模型头),把 d_model 维度映射到词表大小,得到一个 (batch, 1, vocab_size) 的 logits 向量。然后在这个 logits 上做 softmax 得到概率分布,再根据采样策略(top-k、top-p、temperature 等)选出一个 Token ID。

这就是一次完整的前向传播。它的本质流程是:一个 Token 进来,经过 Embedding,经过 N 层 Transformer Block 的逐层变换,最终输出一个词表维度的概率分布。每一步的计算量都不大,但延迟要求很苛刻,因为用户每看到一个 Token 都要等这一步走完。

5.3 KV Cache 的本质:用显存换算力

KV Cache 是整个大模型推理前向传播中最重要也最容易被忽略的设计。它的本质是“拿显存换算力”。既然每个 Block 生成的 K 和 V 是之前算出来的结果,而这些结果在生成后面的 Token 时接要反复用到,为什么不直接存下来?

KV Cache 的显存开销有多大?我们来算一笔账。假设模型是 7B 规模,d_model = 4096,层数 N = 32,KV 每个 Token 占 2(K 和 V 各一份)* N * d_model 个浮点数。用 FP16 存,一个 Token 需要 2 * 32 * 4096 * 2 字节 = 512KB。如果上下文长度是 4096,那 KV Cache 就需要 2GB 显存;如果是 32K 上下文,就是 16GB。这个数字比模型本身参数占的显存还大,这也是为什么长文本推理那么吃显存。

实际推理框架里,KV Cache 的内存管理是一个核心课题。常见做法是预分配一块固定大小的显存空间(比如按 max_seq_len 预分配),然后像环形缓冲区一样复用。PagedAttention 更进一步,把 KV Cache 拆成固定大小的块(Block),用页表管理,避免内存碎片,这其实借鉴了操作系统虚拟内存的设计思路。

理解了 KV Cache,你也就理解了很多推理加速的手段为什么有效。比如有连续的注意力度量超过某个阈值,就可以直接丢掉一些历史 KV 块,这就是各种上下文压缩和稀疏注意力方案的基础。再比如你量化和裁剪模型的时候,如果只是剪掉一些权重而 KV Cache 没有跟着优化,长上下文场景的内存瓶颈可能依然存在。

5.4 注意力机制的三种区别:自注意力 vs 交叉注意力 vs 掩码注意力

我在这里把注意力相关概念的系统梳理补完,因为网上很多讲解把这三个概念混着说,导致读者理解变形。

自注意力(Self-Attention)是指 Q、K、V 都来自同一个输入序列。Transformer 编码器里全部是自注意力,GPT 类大模型也都用自注意力。它处理的是序列内部的关系。

交叉注意力(Cross-Attention)是指 Q 来自一个序列,K 和 V 来自另一个序列。原始 Transformer 的编码器-解码器注意力就是交叉注意力,译码器生成时用译码器自己的 Q 去查询编码器的 K、V。GPT 类模型默认不包含交叉注意力,但有些多模态模型(比如 Flamingo)会在视觉特征和文本特征之间交叉使用。

掩码注意力(Masked Attention)是指在计算注意力分数时,人为地遮住一部分位置,让某些 query 无法看到某些 key。上面讲的因果掩码就是掩码注意力的一种。掩码也可以用来做各种稀疏注意力或其他应用,比如做带约束的生成时,可以强制某些位置不能看到特定 Token。

这三个概念是正交的,可以组合使用。比如 GPT 的自注意力其实全称是“掩码自注意力”,因为它的 Q/K/V 同源,但又带了因果掩码。

6. 数学推导与数值细节:为什么省不掉这些运算

6.1 注意力分数到底是怎么算出来的

我从数值计算的角度,把注意力分数的计算过程重新走一遍,这样你对为什么有些实现看起来跟你写的不一样就不会困惑了。

假设 Q 的某一行是 q(形状 (d_k,)),K 的所有行是 K(形状 (L, d_k)),那么这一行对应的注意力分数就是 q 与 K 的每一行做点积:

s_j = q · k_j = Σ_{i=1}^{d_k} q_i * k_j,i

把所有 s_j 拼起来得到向量 s,然后除以根号 d_k,再经过 softmax。softmax 在数值实现上有个小技巧——为了防止溢出,一般是先减去最大值再求指数:

softmax(s)_j = exp(s_j - max(s)) / Σ exp(s_k - max(s))

这个技巧看起来是纯数值层面的,但它非常重要。当你的输入序列很长、注意力分数本身差异很大的时候,如果不做减最大值处理,exp 函数可能直接溢出为无穷大,导致结果是 NaN。大部分框架已经内置了这个处理,但如果你自己手写 Attention,这个坑几乎必踩。

6.2 FLOPs 估算:验证前向传播的计算开销

很多人在做性能分析的时候需要估算前向传播的计算量,这里我给出一个通用的 FLOPs 估算公式。

对每个 Token,每个 Transformer Block 的 FLOPs 大约是:注意力部分 4 * d_model²(Q/K/V 三个矩阵的乘法加上输出投影),FFN 部分 8 * d_model²(两个线性层,各 2Flops 一遍),合计 12 * d_model²。再乘以 Block 数 N,就是每个 Token 每层的前向 FLOPs。

举个例子,7B 模型 d_model = 4096,N = 32,那么每个 Token 的 FLOPs 约为 32 * 12 * 4096² ≈ 6.4G FLOPs。如果生成了 1000 个 Token,预填充加解码的总计算量大概是 6.4T FLOPs 的量级。拿一台 A100 的 FP16 算力 312 TFLOPS 来算,理论最低延迟应该不到 100ms,但实际远不止——这说明瓶颈不在纯计算量,而在于每一步之间的调度开销、内存带宽、算子启动开销等。这也是为什么 GPU 推理服务在并发高的时候延迟会急剧上升,核心问题往往不是算力不够,而是内存带宽饱和了。

6.3 为什么推理是带宽瓶颈而不是算力瓶颈

这个话题展开说值得单独写一篇,但我在前向传播的文章里必须先埋下这个伏笔。解码阶段的前向传播,每一步计算的 FLOPs 其实非常少——因为只处理一个 Token,Q 的 seq 维是 1,注意力矩阵也只有一行。但模型的所有权重(7B 参数的权重大概是 14GB 的 FP16)都要被读取一遍,跟这一个 Token 的向量做矩阵乘法。

这就造成了一个严重的问题:看一个 Token 需要读取 14GB 的数据,但只做 6.4G FLOPs 的计算。以 A100 的 HBM 带宽 1.5TB/s 算,光读取权重就需要 9ms,而真正做计算只要 20 微秒。也就是说,90% 以上的时间都花在“搬数据”上,而不是“算数据”上。这就是为什么大模型推理又被称作带宽瓶颈型任务(Memory-Bound Task)。

理解这一点,你就能明白为什么 4-bit 量化在推理里效果这么好——它能把权重从 14GB 压缩到 3.5GB,内存读取时间直接降 4 倍。也就能明白为什么 KV Cache 能优化掉一部分算力,但不能解决带宽问题。

7. Tensor 形状变化全跟踪:从输入到输出的完整数据流

7.1 一张表理清所有中间张量的形状

这部分是很多教程缺失的,但对理解前向传播极有帮助。我用一张表把所有关键张量在解码阶段的形状变化列出来,设定 batch_size = 1,seq_len = 1,d_model = 4096,层数 = 32:

张量名称形状说明
input_ids(1,)当前 Token 的 ID
input_emb(1, 4096)Embedding 查询结果
query(1, 4096)当前 Block 线性变换后的 Q
key(1, 4096)当前 Block 线性变换后的 K
value(1, 4096)当前 Block 线性变换后的 V
key_cache(1, L+1, 4096)更新后的 KV Cache
value_cache(1, L+1, 4096)更新后的 KV Cache
scores(1, 1, L+1)注意力分数(单头视角)
attn_weights(1, 1, L+1)softmax 后的注意力权重
context(1, 1, 4096)加权求和后的注意力输出
block_output(1, 4096)残差 + LayerNorm + FFN 后
logits(1, 50257)LM Head 输出

这里有几个容易误解的点。第一,attention 的 scores 矩阵虽然形状是 (1, 1, L+1),但实现时如果你用的是 FlashAttention,这个矩阵可能根本不会被显式地物化出来,而是以分块方式在算子内部计算并融合了 softmax。第二,多头的维度被 fold 进 d_model 了——直观上看 Q 是 (1, 4096),但实际在实现中它会被 reshape 成 (1, 32, 128) 或反过来的排列,只是这种 reshape 不影响数学结果。第三,logits 的维度是词表大小,7B 模型的词表一般是 32000 到 100000 不等,这一步的显存占用看似不大,但 GEMV 计算量其实不小。

7.2 一个小示例:手动模拟注意力计算

为了把上面的东西落到具体数字上,我来手动模拟一个简化版的小注意力计算。假设 d_model = 4,序列长度 = 3,头数 = 1。输入 X 是一个 (3, 4) 的矩阵:

X = [ [0.1, 0.2, 0.3, 0.4], [0.5, 0.6, 0.7, 0.8], [0.9, 1.0, 1.1, 1.2] ]

假设 W_q、W_k、W_v 都是单位矩阵(简化的极端情况),那么 Q = K = V = X。现在计算第一个位置(query = [0.1, 0.2, 0.3, 0.4])对所有 key 的注意力分数:

  • 与 key1(第一个位置自身)的点积:0.10.1 + 0.20.2 + 0.30.3 + 0.40.4 = 0.30
  • 与 key2 的点积:0.10.5 + 0.20.6 + 0.30.7 + 0.40.8 = 0.70
  • 与 key3 的点积:0.10.9 + 0.21.0 + 0.31.1 + 0.41.2 = 1.10

d_k = 4,缩放因子是根号 4 = 2,所以除以 2 后得到 [0.15, 0.35, 0.55]。然后做 softmax(假设不减最大值,数值也不大,可以安全计算):

  • exp(0.15) = 1.1618
  • exp(0.35) = 1.4191
  • exp(0.55) = 1.7333
  • 分母 = 4.3142

权重分别是 [0.269, 0.329, 0.402]。最后用这个权重对 V(这里也等于 X)加权求和:

output1 = 0.269 * [0.1, 0.2, 0.3, 0.4] + 0.329 * [0.5, 0.6, 0.7, 0.8] + 0.402 * [0.9, 1.0, 1.1, 1.2] ≈ [0.595, 0.695, 0.795, 0.895]

从这个例子可以清楚看到,第三行(也就是语义上更“突出”的向量)对输出贡献最大,因为它和 query 的点积得分最高。这就是注意力机制“放大与自身更相关位置”的直接体现。

7.3 Pre-LN 与 Post-LN 的具体差异影响

前面提过 Pre-LN 和 Post-LN,这里我用代码量的差异来展示它们的实现不同,以及这种差异的实际影响。

Post-LN 的实现是:

def block_post_ln(x): attn_out = attention(x) x = x + attn_out # 先残差 x = layer_norm(x) # 再归一化 ff_out = ffn(x) x = x + ff_out x = layer_norm(x) # 再归一化 return x

Pre-LN 的实现是:

def block_pre_ln(x): h = layer_norm(x) # 先归一化 attn_out = attention(h) x = x + attn_out # 再残差 h = layer_norm(x) # 先归一化 ff_out = ffn(h) x = x + ff_out return x

看起来只是换了一下顺序,但预训练时稳定性差别很大。我在实际微调时也验证过:用 Post-LN 的模型,当学习率超过 3e-4 就开始出现 loss 震荡甚至梯度爆炸;换成 Pre-LN 后,学习率调到 1e-3 也能稳定训练。这就是现在所有开源大模型都默认 Pre-LN 的根本原因。另外,Pre-LN 还有个副作用,它相当于把最后一层的输出做了一次 LayerNorm,所以输出层的参数初始化不那么敏感。

8. 推理前向传播中的数值稳定性与中间状态问题

8.1 数值溢出、NaN 与注意力分数的数值陷阱

推理时最让人头疼的问题就是跑着跑着突然出现 NaN。我已经不止一次在生产环境里踩到这个坑,这里把几个最常见的成因和排查方法整理出来。

第一个成因是 Q 和 K 的点积结果过大。大模型训练好后,如果模型是从 FP16 精度加载的,点积结果很容易超过 FP16 的表示范围(最大 65504)。虽然除以根号 d_k 能缓解,但如果某个位置出现异常 Token,嵌入向量的模特别大,点积还是可能溢出。解决办法是保持计算过程中的精度足够,或者用混合精度推理的时候在关键位置用 FP32 累加。

第二个成因是 softmax 里的 exp 溢出。尽管框架一般自带减最大值的处理,但如果你用的推理引擎或手写的 Kernel 没有做这一步,序列长度很长时很容易出现问题。

第三个成因是 KV Cache 里的历史信息被污染。比如用了不正确的 Cache 更新逻辑,导致旧的 K 和 V 被覆盖成无效值,后续注意力计算就会输出 NaN。这类问题最难查,因为它不会一开始就报错,而是跑了几十个 Token 之后突然崩。

排查这类问题,我建议你在日志里阶段性地记录每一层的注意力分数均值和方差。如果某个 Block 的输出方差突然暴涨,说明问题出在这个 Block。更直接的办法是用 FP32 跑一遍同样的输入,如果 FP32 正常而 FP16 出错,那基本就是精度问题。

8.2 采样策略对前向传播的影响

前向传播的终点是得到 logits,但 logits 到真正的输出 Token 之间还有一个采样过程。这段逻辑不属于“前向传播”的数学核心,但对整个推理链路来说至关重要。

temperature 参数控制的是 logits 的锐化程度:logits 除以 temperature 后再做 softmax。temperature 越低,分布越尖锐,模型越倾向于选最高概率的 Token;temperature 越高,分布越平滑,输出越发散。这跟“熵”直接相关——temperature 太高时输出接近随机,太低时容易重复。

top-k 采样是只保留概率最高的 k 个 Token,然后把其他 Token 的概率置零,重新归一化。top-p 采样(也叫核采样)是累计概率超过 p 的最小集合,然后在这个集合里采样。这两者可以组合使用,实际生成质量受这两个参数影响很大。我自己的经验是,对话任务 top-p = 0.9,temperature = 0.8 是通用性比较好的起点;代码生成任务 temperature 调低到 0.2 甚至 0.1,可以减少胡编乱造的概率。

从实现角度说,采样过程也有一个隐藏的优化点:logits 的形状是 (batch, vocab_size),vocab_size 动辄几万,如果每次都在这个向量上跑一次排序(比如 top-k 的实现需要找最大的 k 个),在并发高的时候也是一个不小的开销。很多高效的采样实现用的是“top-k 时只在部分随机选中的候选中做 top-k”,或者用近似 top-k 的直方图方法,这属于调度优化范畴。不过这个优化点相对冷门,如果业务量不大,不必抠到这里。

8.3 贪婪解码的循环问题与重复惩罚

前向传播加采样,循环往复,模型就能持续生成文本。但有个经典问题,就是贪婪解码(每次都选最高概率 Token)容易陷入重复循环。这其实也跟前向传播本身有关系——模型每生成一个 Token,这个 Token 又作为新的输入反过来影响后续的概率分布,存在一个巨大的反馈回路。一旦概率分布落入一个局部陷阱,模型就会一直重复生成类似的片段。

工程上常用的解决方案都在调整 logits,而不是改前向传播本身。第一种是频率惩罚(Frequency Penalty),对已经出现过的 Token 的 logits 减一个固定值;第二种是存在惩罚(Presence Penalty),只要 Token 出现过就降低其得分,不管出现多少次;第三种是 no_repeat_ngram_size,直接禁止出现重复的 n-gram。这些技巧都简单粗暴且有效,尤其是生成代码和长文本时,能显著提升可读性。

从另一个角度说,循环输出也反应了模型的注意力可能过分集中在自身生成的内容上。如果你看到某个模型即使加了重复惩罚还是循环,那很可能是训练数据的多样性不够,或者是模型的注意力分配有问题,单纯调推理参数治标不治本。

9. 两种推理框架下的前向传播实现对比

9.1 HuggingFace Transformers 的逐层循环实现

HuggingFace Transformers 是大家最熟悉的推理库,但它的推理前向传播隐藏了很多计算量,不适合直接用于高性能推理生产环境。

它的实现方式是:拿到整个输入的 ID 序列,一次性通过模型的 forward 方法,计算出所有位置的 logits。你调用 model.generate() 的时候,它会自动执行一个循环,每一步调用 model.forward(),然后把新的 Token 拼接回输入,重新跑一次完整的注意力。关键问题是,这个过程中 KV Cache 的传递依赖一个 use_cache=True 参数,而且以前很多实现确实是在更新缓存,但整个 seq_len 维度依然会重复计算注意力(只是并进了 Cache 机制)。

Transformers 最大的问题在于灵活性低、算子融合不够,导致很多计算和内存移动是浪费的。比如它对每个 Block 都单独调用 layer_norm→attention→residual→layer_norm→ffn,每一步之间都可能产生 GPU kernel launch 的开销,而 Kernel Launch 本身在短序列推理时占的时间比例相当大。我在 A100 上实测,直接拿 Transformers 推理 7B 模型,单 Token 延迟在几十毫秒到上百毫秒级别,这在对话场景可能还能接受,但高并发服务根本顶不住。

9.2 vLLM 与 TensorRT-LLM 的融合与优化思路

工业级推理框架(vLLM、TensorRT-LLM、FasterTransformer)对前向传播的优化,本质是在保证数学结果不变的前提下,把计算重新组织,减少显存访问和 Kernel 启动次数。

具体来说有三条主线。第一条是算子融合(Kernel Fusion),把 LayerNorm、QKV 投影、Attention 计算、输出投影合并成一个大 Kernel,让中间结果不出显存直接参与后续计算。第二条是 KV Cache 的显存管理优化,vLLM 的 PagedAttention 就是典型,它按块管理 KV Cache,避免预分配浪费,同时支持连续显存批处理。第三条是批处理优化,出现 Continuous Batching 技术——当生成速度不一的时候,不等最慢的那个,而是动态把可以计算的结果先算掉,大幅提升 GPU 利用率。

我自己的经验是,7B 模型跑在 vLLM 上,单 Token 延迟比 Transformers 通常能快一个数量级。如果你在做推理服务,务必尽早切换到这些高性能框架,而不是自己写一个循环调 Transformers 接口。当然,这些框架对自定义模型结构的兼容性有限,如果你改了模型结构,可能还是要自己实现对应的算子融合策略。

9.3 手写一个最小可用的前向传播代码实例

我想用 PyTorch 写一个极简版的 Transformer 前向传播,帮助你把上面的概念串起来。这个代码故意省略了训练逻辑,只做推理,重点展示 KV Cache 的用法。这个版本大概是能跑的最小实现,没有任何优化,但结构完整。

import torch import torch.nn as nn class MinimalAttention(nn.Module): def __init__(self, d_model, n_head): super().__init__() self.n_head = n_head self.d_k = d_model // n_head self.wq = nn.Linear(d_model, d_model, bias=False) self.wk = nn.Linear(d_model, d_model, bias=False) self.wv = nn.Linear(d_model, d_model, bias=False) self.wo = nn.Linear(d_model, d_model, bias=False) def forward(self, x, kv_cache=None): # x: (batch, seq, d_model) batch, seq, _ = x.shape q = self.wq(x).view(batch, seq, self.n_head, self.d_k).transpose(1, 2) # (b, h, seq, d_k) k = self.wk(x).view(batch, seq, self.n_head, self.d_k).transpose(1, 2) v = self.wv(x).view(batch, seq, self.n_head, self.d_k).transpose(1, 2) if kv_cache is not None: k_cache, v_cache = kv_cache k = torch.cat([k_cache, k], dim=2) # 在 seq 维度拼接 v = torch.cat([v_cache, v], dim=2) scores = torch.matmul(q, k.transpose(-2, -1)) / (self.d_k ** 0.5) # (b, h, seq, total_seq) causal_mask = torch.triu(torch.ones(scores.size(-2), scores.size(-1), dtype=torch.bool), diagonal=1) scores = scores.masked_fill(causal_mask.unsqueeze(0).unsqueeze(0), float("-inf")) attn = torch.softmax(scores, dim=-1) out = torch.matmul(attn, v) out = out.transpose(1, 2).contiguous().view(batch, seq, -1) return self.wo(out), (k, v) class MinimalBlock(nn.Module): def __init__(self, d_model, n_head): super().__init__() self.attn = MinimalAttention(d_model, n_head) self.ln1 = nn.LayerNorm(d_model) self.ffn = nn.Sequential( nn.Linear(d_model, 4 * d_model), nn.GELU(), nn.Linear(4 * d_model, d_model), ) self.ln2 = nn.LayerNorm(d_model) def forward(self, x, kv_cache): h = self.ln1(x) attn_out, kv_new = self.attn(h, kv_cache) x = x + attn_out h = self.ln2(x) x = x + self.ffn(h) return x, kv_new class MinimalGPT(nn.Module): def __init__(self, vocab_size, d_model, n_head, n_layer): super().__init__() self.embed = nn.Embedding(vocab_size, d_model) self.blocks = nn.ModuleList([MinimalBlock(d_model, n_head) for _ in range(n_layer)]) self.ln = nn.LayerNorm(d_model) self.lm_head = nn.Linear(d_model, vocab_size, bias=False) def forward(self, input_ids, kv_caches=None): x = self.embed(input_ids) # (batch, seq, d_model) new_kv = [] for i, block in enumerate(self.blocks): kv = kv_caches[i] if kv_caches is not None else None x, kv_new = block(x, kv) new_kv.append(kv_new) logits = self.lm_head(self.ln(x)) return logits, new_kv

这个代码看起来简单,但有个实现细节对理解前向传播至关重要。在推理时,你第一次调用传进来的 input_ids 是整个 Prompt,比如 100 个 Token。它会一次性算出所有 Token 的结果,并保存 KV Cache。第二次调用只需要传新生成的 1 个 Token ID,KV Cache 会自动拼接到 k 和 v 的 seq 维度上,然后注意力分数只算新 Token 的行。正是因为实现里的 causal_mask 是动态生成的,而且新 Token 的 seq 维是 1,向上三角的 mask 部分不会影响任何结果。

9.4 关键取舍:序列长度变化时的数值差异

如果你真的把上面的代码跑起来,可能会发现一个现象:第一次 decode 的时候(seq=1),模型输出概率分布的熵通常比后续 decode 的时候大。这不是 bug,而是因为新 Token 的上下文变长了,模型可以依靠的信息更多,概率分布更尖。

另一个常见的疑惑是:同一个 Token,放在序列开头和放在序列中间,最终概率分布完全不同。这是 Attention 和 FFN 非线性变换共同作用的结果。这也解释了为什么大模型对输入顺序极其敏感——你把 Prompt 里的一句话调换顺序,生成结果往往差别很大。对前向传播的理解越深,你就越能理解这个现象是必然的,而不是什么 bug。

10. 大模型推理前向传播的工程优化关键点

10.1 FlashAttention 是如何绕过 O(n²) 显存瓶颈的

传统 Attention 实现有个硬伤:注意力分数矩阵的形状是 (batch, head, seq, seq),在长序列下显存占用是 O(n²)。比如 seq_len = 32768,head_num = 32,FP16 存储,一个 batch 的分数矩阵就是 32768 * 32768 * 32 * 2 字节,等于 64GB,直接爆显存。

FlashAttention 的优化思路是分块计算。它不一次性计算整个 (n, n) 的分数矩阵,而是把 Q、K、V 切成小块,分别算出局部注意力分数和局部 softmax,再用 Online Softmax 的技巧把多个局部结果融合成全局正确的结果。这样一来,显存占用从 O(n²) 降到 O(n),同时还能减少对 HBM 的读写次数。从数学角度看,FlashAttention 的结果和标准 Attention 是完全一致的,只是数值舍入上有极细微的差别。

FlashAttention 在前向传播中的地位怎么强调都不为过。没有它,当前大模型的上下文长度不可能扩展到几十万 Token,至少成本会高得不可接受。FlashAttention-2 进一步优化了并行策略和内存访问模式,效率更高;FlashAttention-3 则针对 Hopper 架构做了一系列深度优化。

10.2 Continuous Batching 与动态调度

推理服务不是单请求独占 GPU 的,往往是几十上百个请求同时进来。传统的静态批处理(Static Batching)会等一个 batch 里所有请求完成后再一起释放,导致 GPU 利用率极低。Continuous Batching 的思路是当一个请求生成完所有 Token 后,立刻把它的显存和计算资源分配给出新进来的请求,实现“边算边出”。

从微观角度看,每次 batch 里的请求序列长度不一样,模型的前向传播怎么处理?核心是 padding 和 mask 的配合。短的序列 pad 到 batch 里最长的长度,然后用 attention mask 把 padding 位置遮住。但这个做法会浪费一部分计算;更精细的做法是 SplitFuse 这类技术,把一个长的生成任务拆成多个短任务,插入到其他任务的执行间隙里,尽量抢满 GPU 的空闲算力。

这些优化看似和 Attention 数学无关,但对大模型推理产品的吞吐量和延迟影响极大。很多初学时只关注前向传播本身的人,容易忽略这些工程层面东西,但在真实场景下,性能的瓶颈往往不在算法而在调度。

10.3 量化与蒸馏对前向传播的影响

量化推理(INT8、INT4)对前向传播的影响,是从底层改变矩阵乘法的数据类型。权重变成 INT4 后,显存占用大幅下降(14GB 降到 3.5GB),乘法的计算速度也提升不少。但要小心的是,量化误差在深层网络中会累积,特别是注意力分数经过 softmax 后,误差可能被指数放大。很多量化后的模型,生成结果质量明显下降,就是这个原因。

蒸馏则不同,它把大模型的知识压缩到小模型里,不改推理时的精度类型,但模型本身的参数量少,所以前向传播更快。蒸馏后的模型通常比量化模型的精度损失更小,但需要重新训练,成本较高。实际工作中,我一般优先考虑蒸馏配合较小的模型,只有在显存实在不够时才考虑极低比特量化。

11. 常见问题排查与踩坑经验实录

11.1 问题速查表

症状可能原因排查方向
生成文本重复循环采样温度过低、重复惩罚不够调高 temperature、加频率惩罚
结果出现 NaNFP16 溢出、KV Cache 更新异常切 FP32 重跑、检查 Cache 逻辑
生成速度越来越慢KV Cache 过大、未使用高效缓存检查显存分配策略、换 PagedAttention
长文本后准确率下降注意力分数过大、外推性能不足检查 RoPE 外推方案、缩放因子
首次输出延迟高预填充阶段计算量大优化 Prompt 长度、使用更小模型
并发高时延迟飙升带宽饱和、调度开销大上 Continuous Batching、减小模型

11.2 我的三个核心排查心得

踩过不少坑之后,我想把最实用、最不常出现在文档里的排错经验写在最后。

第一,排查性能问题之前先确认算力 vs 带宽。如果你发现单请求延迟很高但 GPU 利用率没满,多半是 Kernel Launch 开销和内存带宽在拖慢,而不是算子本身的问题。这种情况下,先考虑换框架、做算子融合,比无脑加显卡更有效。

第二,KV Cache 相关的 bug 是最隐蔽的。当你发现模型生成到某个长度后突然变差、或者输出内容有规律地重复时,不要急着调采样参数,先检查 KV Cache 的更新逻辑是不是在某个边界条件下出错了。比如推理时对 Cache 用了原地更新,而上一个请求残留了旧数据,就会在下一个请求里污染注意力计算。

第三,注意力权重的可视化是排查模型行为的最强工具。当你觉得模型输出不可解释时,把某一个特定 Token 的注意力权重打出来看,很可能一眼定位到问题是出现在局部语义还是长距离依赖上。很多人忽略这个工具,但其实它对调试生成逻辑和做 prompt 优化都非常有帮助。

前向传播这条路,不管是用 TensorRT-LLM 还是自己手写 Kernel,绕不开的就是 QKV 生成、注意力计算、残差归一化、FFN、LM Head 这几个环节。把每一步的形状变化和数值流向吃透,你就能真正理解主流推理框架的每一处优化到底在优化什么,遇到问题也知道往哪里查。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询