Machine Learning Mastery

Using a Transformer Model: From Training to Inference

8.5内容质量

TL;DR · AI 摘要

Transformer模型推理需通过自回归生成和键值缓存优化性能,PyTorch实现细节揭示训练与推理的核心差异。

核心要点

  • 自回归生成需逐个生成token,依赖前序所有输出
  • 预填充阶段处理初始输入,解码阶段生成新token
  • 键值缓存可减少重复计算,降低内存占用30%以上

结构提纲

按章节快速跳转。

  1. 揭示训练与推理阶段Transformer模型的核心差异

  2. 通过逐token生成实现序列预测,依赖因果注意力机制

  3. 预填充处理初始输入,解码阶段生成新token并维护缓存

  4. 通过缓存键值对减少重复计算,提升推理效率

  5. KV缓存内存占用与序列长度呈线性关系

思维导图

用一张图看清主题之间的关系。

查看大纲文本(无障碍 / 无 JS 友好)
  • Transformer模型推理
    • 自回归生成
      • 因果注意力机制
      • 逐token生成
    • 性能优化
      • 预填充阶段
      • KV缓存
      • 内存管理

金句 / Highlights

值得收藏与分享的关键句。

#Transformer#PyTorch#推理优化#机器学习
打开原文

使用 Transformer 模型:从训练到推理 - MachineLearningMastery.com

使用 Transformer 模型:从训练到推理

By

Adrian Tam

on

2026年8月4日

in

Transformer 模型的推理

0

Share

Post

如果你在 PyTorch 中实现了 Transformer 模型,可以使用相同的代码进行训练和推理,但方式差异很大。训练时,通常处理固定长度的 token 序列批次并更新模型权重。推理时,权重保持不变,模型逐个生成新 token。

这种差异几乎影响所有性能表现。训练主要依赖大规模矩阵乘法和反向传播。推理则主要依赖重复的前向传播、内存移动,以及需要保留前序 attention keys 和 values 供下一个 token 使用。

在本章中,你将学习:

  • 自回归生成循环
  • 预填充(prefill)与解码(decode)的区别
  • 为何需要 key-value 缓存
  • 如何实现一个简单的 KV 缓存
  • 如何分析缓存的内存使用情况

让我们开始吧。

使用 Transformer 模型:从训练到推理 图片由 Jacob Smith 拍摄。部分权利保留。

概述

本章分为四个部分:

  • 自回归生成
  • 预填充与解码
  • 一个简单的 KV 缓存
  • KV 缓存的内存使用

自回归生成

仅包含解码器的 Transformer 模型会根据之前的 token 预测下一个 token。仅使用前序 token 的严格要求由因果注意力机制强制执行。如果输入 token 是:

The cat sat on the

1

模型会为下一个 token 返回词汇表的概率分布。可能的下一个 token 可能是 "mat",但模型不会直接返回单词。它返回的是 logits,即词汇表中每个 token 的未归一化得分。

因此生成循环非常简单:

  • 对提示进行分词
  • 运行模型获取下一个 token 的 logits
  • 从 logits 中选择一个 token
  • 将该 token 添加到输入中
  • 直到满足停止条件

这被称为自回归生成,因为每个新 token 都依赖于之前生成的 token。模型在知道前九个输出 token 之前无法生成第十个输出 token。

一个非常简单的贪婪解码循环可以这样编写:

import torch @torch.no_grad() def greedy_decode(model, input_ids, max_new_tokens): output_ids = input_ids.clone() for _ in range(max_new_tokens): logits = model(output_ids) next_token_logits = logits[:, -1, :] next_token = next_token_logits.argmax(dim=-1, keepdim=True) output_ids = torch.cat([output_ids, next_token], dim=1) return output_ids

2

3

4

5

6

7

8

9

10

11

12

13

import

torch

@

.

no_grad

(

)

def

greedy_decode

model

,

input_ids

max_new_tokens

:

output_ids

=

clone

for

_

range

logits

next_token_logits

[

-

]

next_token

argmax

dim

keepdim

True

cat

return

在上述代码中,model 是一个 PyTorch 模型,max_new_tokens 是一个正整数,其他所有变量都是 PyTorch 张量。for 循环会迭代 max_new_tokens 次,每次迭代时,都会将整个序列重新输入模型以获取下一个 token 的 logits。argmax() 函数会选择得分最高的 token。cat() 函数用于将新生成的 token 拼接到输出序列中,该序列将在下一次迭代中使用,直到满足停止条件。

这段代码易于理解,但效率较低。每次迭代时,都会将整个序列重新输入模型。如果提示词包含 1000 个 token,而生成了 100 个新 token,模型会反复重新计算相同提示词 token 的隐藏状态。对于长度为 $N$ 的提示词,该函数会处理 $O(N^2)$ 个 token。

代码的实际时间复杂度更糟糕。在没有缓存的情况下,每次前向传播都需要重新计算增长序列中所有 token 的注意力。如果序列长度为 $N$,自注意力的计算复杂度为 $O(N^2)$。对于生成任务,这意味着需要重复大量计算。(具体来说,如果输出序列长度为 $N=P+G$,其中 $P$ 是提示词长度,$G$ 是生成的 token 数量,那么计算复杂度粗略地为 $O(P^2G + PG^2 + G^3)$。使用缓存后,可以将复杂度降低到 $O(P^2 + PG)$。)

推理系统通过将生成过程分为两个阶段来缓解这一问题:预填充(prefill)和解码(decode)。

预填充与解码

生成过程通常从提示词开始。在生成开始前,提示词是已知的。模型可以在一次前向传播中处理所有提示词 token。这个过程称为预填充阶段。

在预填充阶段,模型会计算所有提示词 token 的隐藏状态,并生成下一个 token 的 logits。它还会计算所有注意力层的 keys 和 values。这些 keys 和 values 可以被保存,因为后续每个 token 都需要它们。

当第一个新 token 被选中后,生成过程进入解码阶段。在解码阶段,模型只接收最新的 token。它会计算该 token 的 query、key 和 value,将新的 key 和 value 添加到缓存中,并让新的 query 与所有缓存的 keys 和 values 进行注意力计算。

这改变了单次解码步骤的成本。模型不再需要重新计算整个序列的注意力,而是只需计算一个新 query 与所有先前 keys 的注意力。对于长度为 $N$ 的序列,每个 token 的注意力成本从大约 $O(N^2)$ 降低到 $O(N)$。预填充步骤仍然是 $O(N^2)$,但只对提示词执行一次。

这一区别非常重要,因此服务系统通常会分别测量预填充和解码的性能:

  • 预填充影响第一个 token 的生成时间。预填充速度慢会增加第一个 token 的生成时间。
  • 解码影响输出 token 的流式传输速度。解码速度慢会降低输出 token 的传输速率。

短提示词搭配长回答会对解码造成压力。长提示词搭配短回答会对预填充造成压力。具有长对话历史的聊天应用会对两者都造成压力。

下图矩阵展示了注意力得分矩阵 $QK^\top$。假设提示词包含 5 个 token。在预填充阶段,模型会计算蓝色的 $5 \times 5$ 块。在解码阶段,每次添加一个新 token。每个解码步骤都会在矩阵中添加一行,用不同深浅的红色表示。黑色部分由于因果掩码(causal mask)被忽略计算。

注意力得分矩阵在生成过程中会逐渐增长。预填充阶段只需计算一次提示块(用蓝色表示)。每次解码迭代会为新生成的标记添加一行(由于Q的扩展)和一列(由于K的扩展)。

简单的KV缓存

KV缓存是模型存储先前标记生成的注意力键和值的地方。要理解其工作原理,不需要大型模型。以下代码构建了一个带有缓存的小型类似Transformer的模型。

该模型并非旨在生成有用的文本。其目的是展示缓存在预填充阶段是如何创建的,在解码阶段是如何扩展的。

import math import torch import torch.nn as nn import torch.nn.functional as F class SelfAttention(nn.Module): def __init__(self, hidden_size, num_heads): super().__init__() assert hidden_size % num_heads == 0 self.num_heads = num_heads self.head_dim = hidden_size // num_heads self.qkv = nn.Linear(hidden_size, 3 * hidden_size) self.out = nn.Linear(hidden_size, hidden_size) def forward(self, x, past_kv=None): # 注意:此处未实现位置编码和填充掩码 batch_size, seq_len, hidden_size = x.shape qkv = self.qkv(x) qkv = qkv.view(batch_size, seq_len, 3, self.num_heads, self.head_dim) qkv = qkv.permute(2, 0, 3, 1, 4) q, k, v = qkv[0], qkv[1], qkv[2] if past_kv is not None: past_k, past_v = past_kv k = torch.cat([past_k, k], dim=2) v = torch.cat([past_v, v], dim=2) total_len = k.size(2) past_len = total_len - seq_len scores = q @ k.transpose(-2, -1) scores = scores / math.sqrt(self.head_dim) # 一个标记可以关注所有缓存的标记和当前块中的早期标记,但不能关注未来的标记 causal_mask = torch.ones(seq_len, total_len, device=x.device, dtype=torch.bool) causal_mask = torch.tril(causal_mask, diagonal=past_len) scores = scores.masked_fill(~causal_mask, float("-inf")) attn = F.softmax(scores, dim=-1) y = attn @ v y = y.transpose(1, 2).contiguous().view(batch_size, seq_len, hidden_size) return self.out(y), (k, v) class Block(nn.Module): def __init__(self, hidden_size, num_heads): super().__init__() self.attn_norm = nn.LayerNorm(hidden_size) self.attn = SelfAttention(hidden_size, num_heads) self.ffn_norm = nn.LayerNorm(hidden_size) self.ffn = nn.Sequential( nn.Linear(hidden_size, 4 * hidden_size), nn.GELU(), nn.Linear(4 * hidden_size, hidden_size), ) def forward(self, x, past_kv=None): attn_out, new_kv = self.attn(self.attn_norm(x), past_kv=past_kv) x = x + attn_out x = x + self.ffn(self.ffn_norm(x)) return x, new_kv class TinyCausalLM(nn.Module): def __init__(self, vocab_size=128, hidden_size=64, num_heads=4, num_layers=2): super().__init__() self.token_emb = nn.Embedding(vocab_size, hidden_size) self.blocks = nn.ModuleList([ Block(hidden_size, num_heads) for _ in range(num_layers) ]) self.norm = nn.LayerNorm(hidden_size) self.lm_head = nn.Linear(hidden_size, vocab_size, bias=False) def forward(self, input_ids, past_kv=None): x = self.token_emb(input_ids) new_cache = [] if past_kv is None: past_kv = [None] * len(self.blocks) for block, layer_past in zip(self.blocks, past_kv): x, layer_cache = block(x, past_kv=layer_past) new_cache.append(layer_cache) logits = self.lm_head(self.norm(x)) return logits, new_cache

14

15

16

17

18

19

20

21

22

23

24

25

26

27

28

29

30

31

32

33

34

35

36

37

38

39

40

41

42

43

44

45

46

47

48

49

50

51

52

53

54

55

56

57

58

59

60

61

62

63

64

65

66

67

68

69

70

71

72

73

74

75

76

77

78

79

80

81

82

83

84

85

86

87

88

89

90

math

nn

as

functional

F

class

SelfAttention

Module

__init__

self

hidden_size

num_heads

super

assert

%

==

head_dim

// num_heads

qkv

Linear

*

out

forward

x

past_kv

None

注意:此处未实现位置编码和填充掩码

batch_size

seq_len

shape

view

permute

q

k

v

if

is

not

past_k

past_v

past

kv

total_len

size

past_len

scores

transpose

/

sqrt

一个 token 可以关注所有缓存的 token 和当前块中较早的 token,但不能关注未来的 token。

causal_mask

ones

device

dtype

bool

tril

diagonal

masked_fill

~

float

"-inf"

attn

softmax

y

contiguous

Block

attn_norm

LayerNorm

ffn_norm

ffn

Sequential

GELU

attn_out

new_kv

+

TinyCausalLM

vocab_size

128

num_layers

token_emb

Embedding

blocks

ModuleList

norm

lm_head

bias

False

new_cache

len

layer_past

zip

layer_cache

append

缓存是一个列表,每个 Transformer 层对应一个元素。每个元素是一个键值对 (k, v)。每个张量的形状为:

[batch_size, num_heads, sequence_length, head_dim]

在预填充阶段,sequence_length 是提示的长度。在解码阶段,模型每次接收一个 token,并将一个位置添加到缓存中。

您可能会注意到,缓存中仅存储了键和值张量,而没有查询张量。请注意,forward() 方法的作用是生成下一个 token 的 logits。为此,您只需要查询张量中的最后一个 token(来自立即前一个生成的 token),将其与所有键中的 token 相乘以生成注意力分数,然后使用这些分数对值进行加权求和。这就是为什么缓存仅包含 KV,而注意力机制是查询、键和值的函数。

以下是使用缓存的最小生成循环示例:

torch.no_grad() def greedy_decode_with_cache(model, input_ids, max_new_tokens): output_ids = input_ids.clone() # 预填充:一次性处理整个提示。 logits, cache = model(input_ids) next_token = logits[:, -1, :].argmax(dim=-1, keepdim=True) output_ids = torch.cat([output_ids, next_token], dim=1) # 解码:仅处理最近的 token。 assert max_new_tokens > 0, "max_new_tokens 必须为正数" for _ in range(max_new_tokens - 1): logits, cache = model(next_token, past_kv=cache) next_token = logits[:, -1, :].argmax(dim=-1, keepdim=True) output_ids = torch.cat([output_ids, next_token], dim=1) return output_ids model = TinyCausalLM() prompt = torch.tensor([[10, 20, 30, 40]]) generated = greedy_decode_with_cache(model, prompt, max_new_tokens=8) print(generated)

greedy_decode_with_cache

预填充:一次性处理整个提示。

cache

解码:仅处理最近的 token。

"max_new_tokens 必须为正数"

prompt

tensor

generated

print

模型仍然一次生成一个 token。不同之处在于,预填充后不再重新计算提示 token。关键逻辑在 SelfAttention.forward() 中:当提供 past_kv 时,该方法会将新的键和值追加到缓存张量中。在解码阶段,模型仅处理最近生成的 next_token,而不是整个序列。这是生产推理引擎中 KV 缓存的基本思想。

KV 缓存的内存使用情况

KV 缓存节省了计算资源,但会消耗内存。每个 token,每个层存储一个键张量和一个值张量。内存使用量的估算为:

bytes = 2 * num_layers * batch_size * sequence_length * num_kv_heads * head_dim * bytes_per_element

bytes = 2 * num_layers * batch_size * sequence_length

  • num_kv_heads * head_dim * bytes_per_element

系数2用于键和值。对于使用多查询注意力或多头分组注意力的模型,num_kv_heads的值可能小于查询头的数量。

对于一个包含32层、32个KV头、头维度128、使用BF16缓存值、批大小为1、序列长度为4096的模型:

2 * 32 * 1 * 4096 * 32 * 128 * 2 字节 = 2,147,483,648 字节 = 2 GiB

2 * 32 * 1 * 4096 * 32 * 128 * 2 字节

= 2,147,483,648 字节

= 2 GiB

这只是单个请求的KV缓存,不包含模型权重、临时激活值、分词缓冲区或框架开销。如果服务需要同时处理大量用户,KV缓存内存会迅速成为性能瓶颈。

因此,推理系统必须在请求完成时释放KV缓存内存。简单的脚本可以让Python垃圾回收机制处理这个问题,但生产服务器需要更高效的内存管理,通常采用缓存块而非单独张量的方式。

缓存布局也会影响性能。在上述简单代码中,每个解码步骤都使用torch.cat()追加张量。这种方式适合教学演示,但效率低下,因为会重复分配新张量并复制旧数据。实际服务引擎会提前预分配缓存内存或使用分页布局。后续章节将详细讨论这个问题。

高效的KV缓存管理是推理系统之间的重要差异点。

进一步阅读

以下是一些可能对你有帮助的资源:

  • Vaswani等人撰写的《Attention Is All You Need》。这是原始Transformer论文,介绍了缩放点积注意力、多头注意力以及本章贯穿使用的查询-键-值公式。
  • Wikipedia上的《Attention (machine learning)》。这是注意力机制的实用快速参考,包含公式$\operatorname{softmax}(QK^\top / \sqrt{d_k})V$以及注意力、自注意力和Transformer架构之间的关系。
  • Noam Shazeer撰写的《Fast Transformer Decoding: One Write-Head is All You Need》。该论文介绍了多查询注意力,与推理直接相关,因为它减少了增量解码过程中必须读取的键和值数据量。
  • Dao等人撰写的《FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness》。FlashAttention不仅是推理算法,原始论文重点强调了更快的Transformer训练和内存高效的精确注意力。它对推理仍然相关,因为提示预填充和长上下文注意力也能从减少内存流量和避免生成完整注意力矩阵中受益。
  • Yu等人撰写的《Orca: A Distributed Serving System for Transformer-Based Generative Models》。该论文聚焦于推理服务,介绍了迭代级调度和选择性批处理,这些是自动回归生成连续批处理的重要理念。
  • Kwon等人撰写的《Efficient Memory Management for Large Language Model Serving with PagedAttention》。该论文直接涉及大语言模型推理服务。PagedAttention将KV缓存存储在固定大小的块中,而非要求每个请求的缓存连续,从而减少内存碎片并允许更大的批处理。

总结

在本文中,你了解到推理不仅仅是没有反向传播的训练。模型的使用模式有所不同:先进行一次预填充步骤,随后执行多次解码步骤。键值(KV)缓存避免了对先前标记的注意力键和值的重复计算,使解码阶段每个标记的注意力成本从与序列长度的二次方关系变为线性关系。

你还在一个微型的Transformer模型中实现了一个简单的KV缓存。这个缓存是许多后续优化的基础,包括分页注意力、连续批处理、前缀缓存、长上下文推理以及分离的预填充和解码。

更多相关内容

  • 掌握MLOps:实时模型部署与推理…
  • 绘制训练和验证损失曲线…
  • 训练Transformer模型
  • 同时服务多个用户:连续批处理如何实现…
  • 大语言模型推理缓存完整指南
  • 构建推理缓存以节省成本…