从零手写大模型 · 推理篇 12 - Transformer Stack(Transformer 堆叠)

上集回顾

EP11里我们把Pre-LN的两行公式装进了一个TransformerBlock:

x1 = x + Attention(LayerNorm(x))
x2 = x1 + FFN(LayerNorm(x1))

用"猫吃鱼"(d_model=4)走了一遍完整的shape标注,最后拼出了一个能跑的nn.Module。结尾留了个悬念:真正的GPT不是一个block,是一摞block。这一集就来解决这个问题——把单个工位,变成一整条流水线。

为什么要堆叠

一个TransformerBlock做的事情,说白了就是"每个token看看别的token,再自己消化一下",做一轮这样的操作。

问题是:一轮消化能学到的东西有限。"猫吃鱼"这种三个字的句子看着简单,但语言里真实的依赖关系可能隔得很远——代词指代谁、前面埋的伏笔在后面才兑现、一个词的含义要看整段上下文才能确定。一层Attention能捕捉到的,往往是比较直接的关联;更抽象、更远距离的关系,需要多层叠加,一层一层把理解往深处推。

这也是"深度学习"里"深度"这两个字的来源——不是靠一层算得多准,是靠层数堆出来的。GPT-2 small用了12层,GPT-3更是堆到了96层。层数越多,模型能表达的抽象层次理论上越丰富,当然参数量和计算量也跟着涨。

堆叠这件事,公式上其实很朴素

先看一个事实:TransformerBlock的输入输出shape是完全一样的。

输入x是(batch, seq_len, d_model),经过Attention、经过FFN、经过两次残差相加,输出还是(batch, seq_len, d_model)。形状没变。

这个"形状不变"是能堆叠的前提——正因为一个block吐出来的东西,长得跟它吃进去的东西一模一样,所以可以直接把上一个block的输出,原封不动喂给下一个block当输入,一直串下去:

x0 = 输入(embedding之后)
x1 = Block1(x0)
x2 = Block2(x1)
x3 = Block3(x2)
...
xN = BlockN(x_{N-1})

每一层都在上一层的输出基础上继续"消化",但接口始终不变。第 07 期讲过,残差连接为信息和梯度提供了直接通路,有助于深层网络训练。堆叠时还必须让每个 Block 保持相同的输入输出维度;在这个接口约束下,多个 Block 才能直接串联。层数仍受显存、计算量和训练稳定性的限制。

猫吃鱼,走一遍两层堆叠

继续用老朋友"猫吃鱼",d_model=4,3个token,这次堆两层block(n_layers=2)方便手算,实际GPT-2 small是12层,公式完全一样,只是循环次数不同。

输入(EP03/EP04的Embedding + 位置编码之后):

x0.shape = (1, 3, 4)   # batch=1, seq_len=3(猫/吃/鱼), d_model=4

进第1层:

x1 = Block1(x0)
x1.shape = (1, 3, 4)   # 形状不变,内容已经被Attention+FFN加工过一轮

进第2层:

x2 = Block2(x1)
x2.shape = (1, 3, 4)   # 还是不变

到这里,"猫吃鱼"这三个token已经经过了两轮"互相看看再消化"。第一层里,"鱼"这个token可能已经开始关注"吃";到了第二层,"吃"这个已经带上"鱼"信息的token,又能反过来影响"猫"——层数越多,token之间能建立的间接关联就越复杂。

最后一步,也是这集要补的一块拼图:前十一期我们完成了输入准备,并组装了单个 Transformer Block,但一个完整的GPT,在最后一层block之后,还有两个小尾巴:

x_final = LayerNorm(x2)          # 再做一次LayerNorm,稳定输出分布
logits  = Linear(x_final)        # (1, 3, 4) -> (1, 3, vocab_size)

这个Linear就是"输出头"(output head),作用是把d_model=4维的向量,重新映射回词表大小的分数(GPT-2词表是50257)。这一步做完,logits里每个位置的4个数,就变成了50257个数——对应词表里每一个token的"打分",分数越高代表模型觉得这个位置接下来越可能是这个token。

一个容易漏掉的点:堆叠的是"层",不是"参数"

新手很容易以为12层是同一个block重复用12次,参数共享。不是的——每一层的Attention、FFN都有自己独立的一套权重,12层就是12套完全不同的参数。代码里体现出来就是:

self.blocks = nn.ModuleList([
    TransformerBlock(d_model, n_heads, 4 * d_model)
    for _ in range(n_layers)
])

这里用列表推导式创建了n_layers个独立的TransformerBlock实例,每个实例内部的nn.Linear初始化时都会拿到自己的随机权重。用nn.ModuleList而不是普通Python列表来装,是因为nn.ModuleList会让PyTorch自动把里面每个block的参数都注册进整个模型的参数列表——这样调用model.parameters()才能拿到全部12层的参数,普通list做不到这件事。

流水线的比喻

EP11把单个block比作"组装车间"的一个工位。这次很自然地往下延伸:多个工位串起来,就是一整条流水线。

原材料(embedding之后的向量)从流水线入口进去,经过工位1加工,半成品传给工位2,工位2再加工,传给工位3……每个工位做的操作是同一类型的活(Attention+FFN两道工序),但每个工位上的工人(权重)是不同的人,各自负责在前一个工位的基础上再往前推进一点。走到流水线尽头,最后再经过质检+贴标签(LayerNorm+输出头),产出的就不再是半成品向量,而是"词表里每个候选词的打分"——一个可以直接拿去做预测的东西。

完整实现

第 11 期为了演示单句话,Block 接收 [seq_len, d_model]。这一期统一改为 [batch, seq_len, d_model],因此直接把三维张量传给第 06 期的注意力层,并解包输出和注意力权重,不再临时添加 batch 维。运行前需先定义或导入第 06 期的 MultiHeadAttention 和第 09 期的 FeedForward。

另外,第 04 期使用固定正余弦位置编码;这里为后续对齐 GPT-2,改用可学习的位置嵌入。当前仍是教学用的堆叠与 shape 演示,第 06 期注意力尚未加入因果掩码,不能直接作为完整的 GPT-2 自回归实现。

import torch
import torch.nn as nn

class TransformerBlock(nn.Module):
    def __init__(self, d_model, n_heads, d_ff):
        super().__init__()
        self.ln1 = nn.LayerNorm(d_model)
        self.attn = MultiHeadAttention(d_model, n_heads)
        self.ln2 = nn.LayerNorm(d_model)
        self.ffn = FeedForward(d_model, d_ff)

    def forward(self, x):
        # x: [batch, seq_len, d_model]
        attn_out, _ = self.attn(self.ln1(x))
        x = x + attn_out
        return x + self.ffn(self.ln2(x))


class GPT(nn.Module):
    def __init__(self, vocab_size, d_model, n_heads, n_layers, max_seq_len):
        super().__init__()
        d_ff = 4 * d_model
        self.token_emb = nn.Embedding(vocab_size, d_model)   # EP03
        self.pos_emb = nn.Embedding(max_seq_len, d_model)    # EP04:这里改用可学习的位置嵌入
        self.blocks = nn.ModuleList([
            TransformerBlock(d_model, n_heads, d_ff)  # EP11
            for _ in range(n_layers)
        ])
        self.ln_f = nn.LayerNorm(d_model)                    # EP08
        self.head = nn.Linear(d_model, vocab_size, bias=False)

    def forward(self, idx):
        B, T = idx.shape
        pos = torch.arange(T, device=idx.device)
        x = self.token_emb(idx) + self.pos_emb(pos)   # (B, T, d_model)

        for block in self.blocks:
            x = block(x)                                # 一层一层往下传,形状不变

        x = self.ln_f(x)
        logits = self.head(x)                           # (B, T, vocab_size)
        return logits

核心就是那个for block in self.blocks循环——朴素到没什么好讲的,这也正说明了残差连接在架构设计上的价值:正是因为每一层接口一致,堆叠这件事才能简单到只是一个for循环。

跑一下shape,验证没问题:

model = GPT(vocab_size=50257, d_model=768, n_heads=12, n_layers=12, max_seq_len=1024)
x = torch.tensor([[464, 3797, 25365]])   # "The cat eats"的token id,长度3
logits = model(x)
print(logits.shape)   # torch.Size([1, 3, 50257])

⚠️ 老规矩要提醒一句:这里的权重还是随机初始化的,logits里的打分现在完全没有语义意义,softmax出来选哪个token基本等于随机数生成器。想让这套结构真正"理解"语言、给出靠谱预测,还差最后一步——把GPT-2训练好的真实权重灌进去。

下集预告

EP13我们去干一件事:把HuggingFace上GPT-2训练好的参数,一个个对应搬进我们自己写的这个GPT类里。这里边有个坑——HuggingFace存权重的方式(Conv1D)和我们用的nn.Linear,矩阵是转置关系,直接复制会形状不匹配。把这个坑填平,模型才算真正"活"过来。

← 返回文章列表