【人工智能】使用Transformer训练小型语言模型

使用注意力机制进行类GPT模型训练

一、相关原理

类GPT,把输入的token作为文章的开头,进行自回归输入,最终输出接下来的文本。相比Transformer(翻译),GPT只需要用到Decoder,相比之下比较好写[修改说明:原文保留,以下为扩充内容]。

1. 自回归生成的原理

什么是"自回归"?用一个具体的例子来说明。假设你正在写一段文字,目前已经写下了token序列:[“我”, “今天”, “去”, “超市”, “买”]——这几个字就像写作文写到一半,现在模型的任务是:接下来该写什么?

模型对着这5个token进行计算,最后输出一个概率分布,告诉我们:“了"这个词的概率最高。于是模型预测出了"了”。现在,完整的句子变成了[“我”, “今天”, “去”, “超市”, “买”, “了”]。然后呢?模型把这一整串——注意,是整个6个token——再喂给自己,继续预测第7个token。这次它可能预测出"一",那就变成7个token;再输入,预测出"支"…如此循环往复。

这就是"自回归"(Autoregressive)名字的由来:用自己上一步的输出作为下一步的输入,一步一步地把文本"生长"出来。每一步都依赖于前面已经生成的所有内容。

这也解释了为什么我们使用ChatGPT或者其他GPT类模型时,文本是一个字一个字地蹦出来的,而不是一次性生成整个段落。因为"一次性生成"在数学上就是不成立的——后面的词依赖前面的词,前面的词还没生成,后面的词无从谈起。每一个token的诞生,都必须等待它前面的所有token都已经就位。

这里有一个很重要的对比:翻译任务。在原始Transformer的翻译场景中,Encoder会先把整句源语言读完,得到一个完整的语义表示,然后Decoder在生成目标语言时,每一步既可以"看"已经生成的输出(用Masked Self-Attention),又可以"看"Encoder提供的源语言信息(用Cross-Attention)——这是一种"两边都看"的双向模式。而GPT不需要翻译,它只有一个方向:从前往后、单向生成。因此GPT只保留了Decoder,并且去掉了Cross-Attention——它没有源语言可以"参考",也不需要。这种单向生成的设计,让GPT天然适合文本续写、故事创作、代码补全等任务。

2. 因果掩码的直观解释

理解了自回归的原理之后,一个实际问题就出现了:训练的时候,我们明明有一整段完整的文本,所有token都是一次性摆在模型面前的。那怎么让模型"假装"它看不到后面的词,只能看到前面的词呢?

答案就是因果掩码(Causal Mask),也叫三角形掩码

假设我们的输入序列是 [“我”, “今天”, “去”, “超市”, “买”, “了”, “一”, “支”, “笔”] 这9个token。在自注意力机制中,每个token都要和所有其他token计算注意力分数。用一个9×9的矩阵来表示这个"谁可以看到谁"的关系:

              我 今天 去 超市 买 了 一 支 笔
        我    ✓   ✗   ✗   ✗   ✗  ✗  ✗  ✗  ✗
        今天  ✓   ✓   ✗   ✗   ✗  ✗  ✗  ✗  ✗
        去    ✓   ✓   ✓   ✗   ✗  ✗  ✗  ✗  ✗
        超市  ✓   ✓   ✓   ✓   ✗  ✗  ✗  ✗  ✗
        买    ✓   ✓   ✓   ✓   ✓  ✗  ✗  ✗  ✗
        了    ✓   ✓   ✓   ✓   ✓  ✓  ✗  ✗  ✗
        一    ✓   ✓   ✓   ✓   ✓  ✓  ✓  ✗  ✗
        支    ✓   ✓   ✓   ✓   ✓  ✓  ✓  ✓  ✗
        笔    ✓   ✓   ✓   ✓   ✓  ✓  ✓  ✓  ✓

看第5行(对应token “买”):它只能"看"到前5个token——[“我”, “今天”, “去”, “超市”, “买”],后面的[“了”, “一”, “支”, “笔”]全部被遮住。这就形成了一个下三角(✓的区域)加一个上三角(✗的区域)的图案,所以叫"三角形掩码"。

在代码实现中,这个遮蔽是通过把上三角位置的注意力分数设为 -∞(实际代码中用 -1e4 近似)来实现的。经过Softmax之后,e^(-∞) = 0,那些位置的注意力权重就变成了0,达到了"看不到"的效果。

你可能会有疑问:这样做是不是在"限制"模型的能力?明明训练数据里有完整信息,为什么不给模型看?其实,这恰恰不是一种限制,而是一种对齐训练和推理的技巧。回想一下:在推理(实际生成文本)的时候,模型确实只能看到已经生成的前半部分,未来的词还没产生,当然看不到。如果训练时让模型能看到后面的词,那训练目标和推理环境就不一致了——模型在训练时学会了"抄答案",到了推理时却无答案可抄,表现必然大幅下降。因果掩码保证:训练时"模拟"了推理时的真实处境,两者保持一致。

这里还有一个重要的对比,值得反复理解:

  • Encoder中的自注意力:双向可见。比如BERT在做句子理解时,每个token可以看到它左边和右边的所有token——这对"理解"任务是必需的,因为要理解一个词往往需要上下文。比如"bank"这个词,不看后面的"river"还是"account",你就不知道它是"河岸"还是"银行"。

  • Decoder/GPT中的因果注意力:单向遮蔽。每个token只能看到它自己和它左边的token——这对"生成"任务是必需的,因为生成永远是从左到右、未知到已知的过程。生成时右边还没有内容,你不能参考不存在的未来。

两者并非谁优谁劣,而是服务于不同的目标:理解用双向,生成用单向。GPT选择了后者,因为它的使命不是理解一段话,而是延续一段话

3. GPT与原始Transformer的关系

理解了自回归和因果掩码之后,我们来说清楚一个经常让人困惑的问题:GPT和原始Transformer到底有什么关系?

一句话总结:GPT的架构 = Transformer的Decoder部分,但去掉了Cross-Attention,只保留了Masked Self-Attention + FFN(前馈网络)。

原始Transformer(出自2017年的论文《Attention Is All You Need》)是一个Encoder-Decoder架构,设计用来做机器翻译:

  • Encoder:读取源语言句子(如英文),通过双向自注意力理解整个句子的含义
  • Decoder:一边用因果自注意力生成目标语言(如中文),一边用Cross-Attention参考Encoder的输出

GPT做的事情完全不同:它不做翻译,它是纯文本生成。输入是"文本的前半段",输出是"文本的后半段"。既然是同一个语言内部的续写,就不需要什么"源语言"的概念。所以:

  • Cross-Attention直接不需要了——没有Encoder的输出可以"交叉注意"
  • 只剩下Decoder中的Masked Self-Attention + FFN,一层一层堆叠起来

这就是为什么说GPT"只用到了Decoder"。

那OpenAI为什么要选择这条路线呢?这背后有很深的洞察:

解码器架构天然适合生成任务。生成永远是单向的——从已有内容推导后续内容。Masked Self-Attention正好完美地建模了这个过程。更重要的是,“预测下一个词"这个训练目标极其简单却极其强大——它不需要任何人工标注。你随便找一段文本,把它切成"前N个词"和"第N+1个词”,就构成了一条训练样本。这意味着你可以在互联网级别的海量无标注文本上进行预训练——这正是GPT系列成功的核心原因。

从GPT-1(2018年,1.17亿参数)到GPT-4(2023年,传闻1.76万亿参数),OpenAI走的就是一条Scaling路线

  • 架构上变化其实不大——核心还是Transformer Decoder的堆叠
  • 变化的是规模:模型越来越大(更多层、更宽维度、更多注意力头)、数据越来越多(从几千本书到整个互联网)、算力越来越猛(从几块GPU到几万块GPU)

正如OpenAI在2020年那篇著名的Scaling Laws论文中所证明的:语言模型的性能与模型大小、数据量、算力呈幂律关系——只要持续加大投入,loss就会持续降低。本文中的小模型训练过程,其实就是这个Scaling Law在微观尺度上的一个缩影。

二、具体效果

现在一共训练了55万个Batch,取其中一些输出作为训练效果的体现,具体如下:

·batch-14000; loss-5.2

320c02ffd735775dd27705ee1d1242e6.png

词汇表还没有多少,正在学习基本语句

·batch-25000; loss-3.5126d2f53936bae91a19a65b571eed8f6.png

句子结构基本正确,但是逻辑欠缺,出现一些不明所以的句子;对部分词语理解有误

·batch-48000; loss-3.0

764335f5b1cc210feb53d98ae7d94e85.png

对部分词的理解欠缺,句子上下文衔接不连贯

·batch-69000; loss-2.5

517caa8175379d5ffe91e304f46442a4.png

逻辑转换莫名其妙,出现不明所以的人物,词性未完全理解

·batch-97000; loss-2.2

84ce2186db4b1e657449ed786fa3cc87.png

基本语句流畅,句间逻辑有很大问题,并且重复的词比较多

·batch-117000; loss-2.0

70e56d8f64e2e48939a047442d419ed9.png

句子语法正确,词义理解基本正确,重复的词汇减少,但是句间逻辑依然有比较大的问题

·batch-197000; loss-1.75

9036cdd29a36d7632097ec97f40b1f93.png

还在学习句子之间的逻辑,明显比上面好,但是出现莫名其妙的转折点

·batch-384000; loss-1.55

89952d488929c6be00ae80b0ed4d4762.png

句子之间的逻辑明显好很多,形容词增加,句子成分更加复杂

·batch-425000; loss-1.5

d746f254faf24d890e0f70488e903758.png

生成的故事已经比较有逻辑了,中间可能有些断断续续,突然出现了一些莫名其妙的内容,但是至少已经很连贯了

·batch-514000; loss-1.45

1bfb8025841c971381b891c15e4199b8.png

对于长故事的创造力很强,但是句子之间还是缺乏逻辑,以及连接的时候会有很多问题。GPT还在天马行空的想象吧!

对上面的内容进行一个总结:

1、loss的折线图如下所示,逐渐趋缓并且达到一个瓶颈期

f9973dab3cbddd4caeda4bcde025586e.png

2、学习内容、产出文本解读

在多次实验中,生成参数保持一致,唯一的变量是所使用的模型。根据 Transformer 的结构特点,每个注意力头(Multi-Head Attention)都会捕捉不同层面的上下文关系。从生成结果可以看出,GPT 已经能够较好地理解词语的基本含义以及上下文之间的常见组合。然而,它在深入理解整体语义、以及在自回归生成过程中有效记忆较远的上下文内容方面,仍存在一定的不足。

深入解读:从训练过程中我们看到了什么

把上面的所有训练快照串起来看,我们能发现很多有意思的规律。

一、Loss曲线的幂律关系

从batch-14000(loss 5.2)到batch-514000(loss 1.45),loss一直在下降,但下降的速度越来越慢。这在坐标系中形成了一条典型的"先陡后平"的曲线——这正是**Scaling Law(规模定律)**的微观体现。

Scaling Law告诉我们:loss与训练数据量之间存在幂律关系,即 loss ∝ (数据量)^(-α),其中α是一个正指数。翻译成直观的理解就是:数据量翻倍,loss并不是等比例地降,而是以一个递减的速率在降。从5.2降到2.0相对容易(只需要大约12万batch),但从2.0降到1.45却需要大约40万batch——每往下压零点几个点,都要付出比之前大得多的数据代价。

这也意味着,本文的训练已经明显进入了瓶颈期。如果要继续降低loss,靠现在的模型规模已经不够了——可能需要从以下方向突破:

  • 增大模型容量:增加d_model(嵌入维度)、n_heads(注意力头数)或n_layers(层数),让模型有更多的参数来吸收更复杂的模式
  • 增加训练数据:更多的文本、更丰富的领域覆盖,让模型见过更多样的语言现象
  • 扩展上下文长度:当前max_seq_len为128,很多长距离依赖关系可能因此被截断

二、模型"学会语言"的阶段性路径

更有趣的是,如果仔细阅读每个batch阶段的生成输出和评价,我们可以清晰地看到模型学习语言的阶段性路径——这不是这个模型独有的,而是几乎所有语言模型训练过程中都会经历的通用规律:

  • 第一阶段:学词汇(batch ~14k,loss 5.2)——词汇表还没有多少,模型刚从随机初始化开始,正在学习最基础的东西:哪些token经常出现在一起,基本的词与词的搭配。这个阶段输出基本是乱码或极短的片段

  • 第二阶段:学语法(batch ~25k-48k,loss 3.5-3.0)——句子结构开始正确,说明模型已经掌握了基本的语法规则(主谓宾结构、虚词用法等)。但语义层面还很脆弱,会出现一些"语法正确但语义不通"的句子,对部分词语的理解也时常跑偏

  • 第三阶段:学局部逻辑(batch ~69k-117k,loss 2.5-2.0)——句子级别已经比较流畅,词汇理解基本正确,重复现象减少。但句子之间的逻辑衔接依然成问题——会突然引入莫名其妙的人物或出现不合逻辑的转折。这说明模型还在学习"前后两句话应该有什么样的关系"

  • 第四阶段:学篇章连贯(batch ~197k+,loss 1.75及以下)——长故事有了一定的创造力,形容词增多,句子成分变复杂。模型开始理解更宏观的叙事结构和篇章逻辑,但仍然会出现"天马行空"式的跳跃——想象力和逻辑性之间的平衡还需要更多训练来打磨

这个学习路径揭示了一个深刻的道理:语言模型的训练,本质上就是从"字面"到"语义"、从"局部"到"全局"的渐进式抽象过程。先学会词是什么,再学会怎么组成句子,然后学会句子之间怎么连接,最后学会怎么构建一个完整的、连贯的篇章。每一步都建立在前一步的基础之上,不能跳跃。

三、实现代码(仅展示部分重要模块)

import math
import torch
import torch.nn as nn
import torch.nn.functional as F


class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, n_heads, dropout=0.1):
        super().__init__()

        assert d_model % n_heads == 0

        self.d_model = d_model
        self.n_heads = n_heads
        self.d_k = d_model // n_heads

        self.W_q = nn.Linear(d_model, d_model)
        self.W_k = nn.Linear(d_model, d_model)
        self.W_v = nn.Linear(d_model, d_model)

        self.W_o = nn.Linear(d_model, d_model)

        self.dropout = nn.Dropout(dropout)


    def forward(self, x, mask=None):
        B, T, C = x.shape

        Q = (
            self.W_q(x)
            .view(B, T, self.n_heads, self.d_k)
            .transpose(1, 2)
        )

        K = (
            self.W_k(x)
            .view(B, T, self.n_heads, self.d_k)
            .transpose(1, 2)
        )

        V = (
            self.W_v(x)
            .view(B, T, self.n_heads, self.d_k)
            .transpose(1, 2)
        )


        # Attention Score
        scores = (Q @ K.transpose(-2, -1)) / math.sqrt(self.d_k)


        # Causal Mask
        if mask is not None:
            scores = scores.masked_fill(mask == 0, -1e4)


        attn = F.softmax(scores.float(), dim=-1).type_as(scores)

        attn = self.dropout(attn)


        output = (
            (attn @ V)
            .transpose(1, 2)
            .contiguous()
            .view(B, T, C)
        )


        return self.W_o(output)
class Model(nn.Module):

    def __init__(
        self,
        vocab_size,
        d_model=512,
        n_heads=8,
        num_layers=12,
        max_seq_len=128,
        d_ff=2048,
        dropout=0.1
    ):
        super().__init__()


        self.d_model = d_model
        self.max_seq_len = max_seq_len


        # Token Embedding
        self.embedding = nn.Embedding(
            vocab_size,
            d_model
        )


        # Position Embedding
        self.pos_embedding = nn.Embedding(
            max_seq_len,
            d_model
        )


        self.dropout = nn.Dropout(dropout)


        # Transformer Blocks
        self.blocks = nn.ModuleList([
            Block(
                d_model,
                n_heads,
                d_ff,
                dropout
            )
            for _ in range(num_layers)
        ])


        self.ln_f = nn.LayerNorm(d_model)


        # Output Projection
        self.head = nn.Linear(
            d_model,
            vocab_size,
            bias=False
        )


        # Weight Tying
        self.embedding.weight = self.head.weight


        self.apply(self._init_weights)



    def _init_weights(self, module):

        if isinstance(module, nn.Linear):

            torch.nn.init.normal_(
                module.weight,
                mean=0.0,
                std=0.02
            )

            if module.bias is not None:
                torch.nn.init.zeros_(module.bias)


        elif isinstance(module, nn.Embedding):

            torch.nn.init.normal_(
                module.weight,
                mean=0.0,
                std=0.02
            )


        elif isinstance(module, nn.LayerNorm):

            torch.nn.init.zeros_(module.bias)

            torch.nn.init.ones_(
                module.weight
            )



    def forward(self, idx, targets=None):

        B, T = idx.shape

        device = idx.device


        # Token embedding
        tok_emb = self.embedding(idx)


        # Position embedding
        pos = torch.arange(
            0,
            T,
            dtype=torch.long,
            device=device
        ).unsqueeze(0)


        pos_emb = self.pos_embedding(pos)


        # Combine embeddings
        x = self.dropout(
            tok_emb + pos_emb
        )


        # Causal Attention Mask
        mask = torch.tril(
            torch.ones(
                T,
                T,
                device=device
            )
        ).view(
            1,
            1,
            T,
            T
        )


        # Transformer layers
        for block in self.blocks:
            x = block(
                x,
                mask
            )


        x = self.ln_f(x)


        logits = self.head(x)


        loss = None


        if targets is not None:

            loss = F.cross_entropy(
                logits.view(-1, logits.size(-1)),
                targets.view(-1)
            )


            # 防止异常 loss
            loss = torch.clamp(
                loss,
                0,
                15
            )


        return logits, loss