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

X AI KOLs Timeline 工具

摘要

本教程介绍如何使用Transformer模型,从训练到推理,重点讲解自回归生成、prefill(预填充)与decode(解码)阶段,以及用于高效推理的键值缓存。

使用Transformer模型:从训练到推理 https://t.co/ER3nUPHqnh
查看原文
查看缓存全文

缓存时间: 2026/08/04 12:08

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

https://t.co/ER3nUPHqnh


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

来源:https://machinelearningmastery.com/using-a-transformer-model-from-training-to-inference/ 如果你已经在PyTorch中实现了一个Transformer模型,你可以使用相同的代码进行训练和推理,但方式截然不同。在训练期间,你通常会处理一批固定长度的token序列并更新模型权重。在推理期间,权重是固定的,模型一次生成一个token。

这种差异几乎改变了关于性能的一切。训练主要由大型矩阵乘法和反向传播主导。推理则由重复的前向传播、内存移动以及对前一个注意力键和值保持可用以供下一个token使用的需求所主导。

在本章中,你将学习:

  • 自回归生成循环
  • prefill和decode之间的区别
  • 为什么键值缓存是必要的
  • 如何实现一个简单的KV缓存
  • 如何推理缓存所使用的内存

让我们开始吧。

使用Transformer模型:从训练到推理 照片由Jacob Smith (https://unsplash.com/photos/a-red-double-decker-bus-driving-down-a-busy-street-LcuBRr7pRCc)拍摄。保留部分权利。

概述

本章分为四个部分;它们是:

  • 自回归生成
  • Prefill和Decode
  • 一个简单的KV缓存
  • KV缓存的内存使用

自回归生成

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

模型返回词汇表上关于下一个token的概率分布。一个可能的下一个token可能是“mat”,但模型不直接返回一个单词。它返回logits,即词汇表中每个token的未归一化分数。

因此,生成循环很简单:

  1. 对提示进行分词。
  2. 运行模型以获得下一个token的logits。
  3. 从logits中选择一个token。
  4. 将该token追加到输入。
  5. 重复直到达到停止规则。

这被称为自回归生成,因为每个新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

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

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

代码的实际时间复杂度更糟。没有缓存时,每次前向传播都会对不断增长的序列中的所有token重新计算注意力。如果序列长度为N,自注意力有O(N^2)次分数计算。对于生成而言,这意味着你重复了大量工作。(确切地说,如果输出序列长度为N=P+G,其中提示长度为P,生成的token数为G,朴素的计算复杂度应为O(P^2G + PG^2 + G^3)。使用缓存,我们可以将其降低到O(P^2 + PG)。)

推理系统通过将生成分为两个阶段来缓解这一问题:prefill和decode。

Prefill和Decode

生成通常以提示开始。提示在生成开始之前是已知的。模型可以在一次前向传播中处理所有提示token。这被称为prefill阶段。

在prefill期间,模型计算所有提示token的隐藏状态,并为下一个token生成logits。它还计算所有注意力层的键和值。这些键和值可以被保存,因为它们将被未来的每个token需要。

在选择了第一个新token之后,生成进入decode阶段。在decode中,模型只接收最新的token。它计算该token的查询、键和值,将新的键和值追加到缓存中,并让新的查询对所有缓存的键和值进行注意力计算。

这改变了一个decode步骤的成本。模型不再为整个序列重新计算注意力,而是只为一个新的查询对所有先前的键计算注意力。对于长度为N的序列,每个token的注意力成本从大约O(N^2)变为O(N)。prefill步骤仍然是O(N^2),但它只对提示执行一次。

这一区别非常重要,以至于服务系统通常分别衡量prefill和decode:

  • Prefill影响首token时间。缓慢的prefill会增加首token的时间。
  • Decode影响输出token的流式传输速度。缓慢的decode会降低输出token流式传输的速率。

短提示长回答强调decode。长提示短回答强调prefill。具有长对话历史的聊天应用则两者都强调。

下面的矩阵展示了注意力分数矩阵QK^\top。假设提示有五个token。在prefill期间,模型计算蓝色的5 \times 5块。在decode期间,一次添加一个新token。每个decode步骤向矩阵中添加一行新的内容,以不同的红色阴影显示。黑色的元素由于因果掩码而在计算中被忽略。

注意力分数矩阵在生成过程中增长。Prefill一次计算提示块(蓝色)。每个decode迭代为新生成的token追加一行(由于Q扩展)和一列(由于K扩展)。

一个简单的KV缓存

KV缓存是模型存储先前token产生的注意力键和值的地方。要了解它是如何工作的,你不需要一个大模型。下面的代码构建了一个带有缓存的小型Transformer风格模型。

这个模型的目的是产生有用的文本。它的目的是展示缓存如何在prefill期间创建并在decode期间扩展。

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):
        # Note: Positional encoding and padding masks are not implemented here
        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)
        # A token may attend to all cached tokens and earlier tokens
        # in the current chunk, but not future tokens.
        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

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

[batch_size, num_heads, sequence_length, head_dim]

在prefill期间,sequence_length是提示长度。在decode期间,模型一次接收一个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()
    # Prefill: process the whole prompt once.
    logits, cache = model(input_ids)
    next_token = logits[:, -1, :].argmax(dim=-1, keepdim=True)
    output_ids = torch.cat([output_ids, next_token], dim=1)
    # Decode: process only the most recent token.
    assert max_new_tokens > 0, "max_new_tokens must be positive"
    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)

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

KV缓存的内存使用

KV缓存节省了计算,但它消耗内存。对于每个token,每一层存储一个键张量和一个值张量。大约的内存使用量为:

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 bytes
= 2,147,483,648 bytes
= 2 GiB

这只是一个请求的KV缓存。它不包括模型权重、临时激活、分词缓冲区或框架开销。如果服务同时处理许多用户,KV缓存内存很快成为限制因素。

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

缓存的布局也很重要。在上面的简单代码中,每个decode步骤使用torch.cat()追加张量。这适用于教学,但效率低下,因为它在反复分配新张量并复制旧数据。真正的服务引擎会预先分配缓存内存或使用分页布局。后面的章节将详细重新讨论这个问题。

高效的KV缓存管理是推理系统之间的主要区别之一。

进一步阅读

下面是一些你可能觉得有用的资源:

  • Attention Is All You Need (https://arxiv.org/abs/1706.03762),作者:Vaswani等人。这是原始的Transformer论文。它介绍了缩放点积注意力、多头注意力以及本章中使用的查询-键-值公式。
  • Attention (machine learning) (https://en.wikipedia.org/wiki/Attention_%28machine_learning%29),在维基百科上。这是关于注意力机制的一个有用的快速参考,包括公式\operatorname{softmax}(QK^\top / \sqrt{d_k})V以及注意力、自注意力和Transformer架构之间的关系。
  • Fast Transformer Decoding: One Write-Head is All You Need (https://arxiv.org/abs/1911.02150),作者:Noam Shazeer。这篇论文介绍了多查询注意力。它与推理直接相关,因为它减少了增量解码期间必须读取的键和值数据量。
  • FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (https://arxiv.org/abs/2205.14135),作者:Dao等人。FlashAttention不仅仅是一种推理算法;原始论文强调了更快的Transformer训练和内存高效的精确保留注意力。它与推理仍然相关,因为提示prefill和长上下文注意力也受益于减少内存流量和避免显式物化完整的注意力矩阵。
  • Orca: A Distributed Serving System for Transformer-Based Generative Models (https://www.usenix.org/conference/osdi22/presentation/yu),作者:Yu等人。这篇论文聚焦于推理服务。它介绍了迭代级别调度和选择性批处理,这些是自回归生成中连续批处理背后的重要思想。
  • Efficient Memory 管理

相似文章