Python技术迷

大神啊!徒手只用 200 行 Python 代码重现 GPT 核心!

别被“重现 GPT”这几个字唬住。

真 GPT 不是 200 行代码写出来的,那东西背后是数据、算力、工程、分布式训练、对齐、安全策略一整套东西。200 行 Python 能写出来的,是 GPT 最硬的那条骨架:输入 token,做因果自注意力,堆 Transformer Block,然后一个字一个字往后生成。

这地方我一般不先讲概念,先看代码。能跑起来,比画十张结构图强。

下面这个版本,我故意写成字符级模型,不搞复杂 tokenizer。它不聪明,但链路是对的。

# mini_gpt.py
import torch
import torch.nn as nn
import torch.nn.functional as F

torch.manual_seed(7)

text = """
接口超时了先别急着改业务代码。
先看 trace,再看 SQL,再看线程池。
缓存击穿不是玄学,慢查询也不是玄学。
程序员最怕的不是 bug,是没日志。
"""

chars = sorted(list(set(text)))
stoi = {c: i for i, c in enumerate(chars)}
itos = {i: c for c, i in stoi.items()}

defencode(s):
return torch.tensor([stoi[c] for c in s], dtype=torch.long)

defdecode(ids):
return''.join(itos[int(i)] for i in ids)

data = encode(text)

batch_size = 8
block_size = 32
n_embd = 64
n_head = 4
n_layer = 3
dropout = 0.1
device = "cuda"if torch.cuda.is_available() else"cpu"

defget_batch():
    ix = torch.randint(0, len(data) - block_size - 1, (batch_size,))
    x = torch.stack([data[i:i + block_size] for i in ix])
    y = torch.stack([data[i + 1:i + block_size + 1] for i in ix])
return x.to(device), y.to(device)

classCausalSelfAttention(nn.Module):
def__init__(self):
        super().__init__()
        self.qkv = nn.Linear(n_embd, n_embd * 3, bias=False)
        self.proj = nn.Linear(n_embd, n_embd)
        self.drop = nn.Dropout(dropout)
        self.register_buffer(
"mask",
            torch.tril(torch.ones(block_size, block_size)).view(1, 1, block_size, block_size)
        )

defforward(self, x):
        B, T, C = x.shape

        q, k, v = self.qkv(x).chunk(3, dim=-1)

        q = q.view(B, T, n_head, C // n_head).transpose(1, 2)
        k = k.view(B, T, n_head, C // n_head).transpose(1, 2)
        v = v.view(B, T, n_head, C // n_head).transpose(1, 2)

        score = q @ k.transpose(-2, -1)
        score = score * (k.size(-1) ** -0.5)

        score = score.masked_fill(self.mask[:, :, :T, :T] == 0, float("-inf"))
        att = F.softmax(score, dim=-1)
        att = self.drop(att)

        out = att @ v
        out = out.transpose(1, 2).contiguous().view(B, T, C)

return self.proj(out)

classFeedForward(nn.Module):
def__init__(self):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(n_embd, 4 * n_embd),
            nn.GELU(),
            nn.Linear(4 * n_embd, n_embd),
            nn.Dropout(dropout),
        )

defforward(self, x):
return self.net(x)

classBlock(nn.Module):
def__init__(self):
        super().__init__()
        self.ln1 = nn.LayerNorm(n_embd)
        self.attn = CausalSelfAttention()
        self.ln2 = nn.LayerNorm(n_embd)
        self.ffn = FeedForward()

defforward(self, x):
        x = x + self.attn(self.ln1(x))
        x = x + self.ffn(self.ln2(x))
return x

classMiniGPT(nn.Module):
def__init__(self):
        super().__init__()
        vocab_size = len(chars)

        self.token_emb = nn.Embedding(vocab_size, n_embd)
        self.pos_emb = nn.Embedding(block_size, n_embd)
        self.blocks = nn.Sequential(*[Block() for _ in range(n_layer)])
        self.ln = nn.LayerNorm(n_embd)
        self.head = nn.Linear(n_embd, vocab_size)

defforward(self, idx, targets=None):
        B, T = idx.shape

        token = self.token_emb(idx)
        pos = self.pos_emb(torch.arange(T, device=device))
        x = token + pos

        x = self.blocks(x)
        x = self.ln(x)
        logits = self.head(x)

        loss = None
if targets isnotNone:
            loss = F.cross_entropy(
                logits.view(B * T, -1),
                targets.view(B * T)
            )

return logits, loss

    @torch.no_grad()
defgenerate(self, idx, max_new_tokens=80):
for _ in range(max_new_tokens):
            ctx = idx[:, -block_size:]
            logits, _ = self(ctx)
            logits = logits[:, -1, :]
            probs = F.softmax(logits, dim=-1)
            next_id = torch.multinomial(probs, num_samples=1)
            idx = torch.cat([idx, next_id], dim=1)
return idx

model = MiniGPT().to(device)
opt = torch.optim.AdamW(model.parameters(), lr=3e-4)

for step in range(1200):
    xb, yb = get_batch()
    _, loss = model(xb, yb)

    opt.zero_grad(set_to_none=True)
    loss.backward()
    opt.step()

if step % 200 == 0:
        print("step =", step, "loss =", round(loss.item(), 4))

start = torch.tensor([[stoi["接"]]], device=device)
out = model.generate(start, max_new_tokens=100)[0]
print(decode(out))

这段代码最关键的地方,不是训练循环,也不是 Embedding。

我第一眼会看这里:

score = score.masked_fill(self.mask[:, :, :T, :T] == 0, float("-inf"))

没有这行,它就不是 GPT 这种自回归模型了。

GPT 生成下一个 token 时,只能看左边已经出现过的内容,不能偷看右边答案。这个遮罩就是“不能作弊”的地方。训练时模型一次性吃进去一整段文本,但每个位置只能关注自己前面的 token。

再看这两行:

x = x + self.attn(self.ln1(x))
x = x + self.ffn(self.ln2(x))

这就是 Transformer Block 里很要命的残差结构。

没有残差,层数一深,梯度容易散,训练会变得很难受。很多文章喜欢把注意力讲得玄乎,我倒觉得这里更像工程上的老习惯:别让一层网络把原来的信息全洗掉,先留条旁路。

CausalSelfAttention 里面那段 q、k、v,就是注意力的核心。

q, k, v = self.qkv(x).chunk(3, dim=-1)
score = q @ k.transpose(-2, -1)
att = F.softmax(score, dim=-1)
out = att @ v

这几行翻译成人话就是:

当前这个 token 拿着自己的 query,去和前面所有 token 的 key 对一下眼神。谁更相关,就从谁的 value 里多拿点信息。

别小看这个机制,它解决的是“上下文关联”的问题。

比如输入是:

接口超时了先别急着改业务代码。
先看 trace,再看 SQL,再看线程池。

模型看到“超时”后面,更可能接“trace”“SQL”“线程池”这些词,而不是乱接一个“香蕉”。当然,这个小模型训练语料太少,效果肯定很糙。它只是把流程跑通,不是让它真去写生产故障报告。

生成部分也很直接:

logits = logits[:, -1, :]
probs = F.softmax(logits, dim=-1)
next_id = torch.multinomial(probs, num_samples=1)

只取最后一个位置的输出,因为我们要预测下一个 token。

这里不用 argmax,而是采样。argmax 太死板,每次都拿概率最高的那个字,生成内容容易卡成复读机。采样会有一点随机性,虽然也可能胡说八道,但至少像是在“往下写”。

这个迷你 GPT 跑起来以后,你会发现一个很现实的问题:结构不难,难的是数据和规模。

模型太小,它记不住复杂模式。

语料太少,它只会鹦鹉学舌。

训练太短,loss 还没压下来。

上下文太短,前面说过什么很快就忘了。

所以别把这个东西吹成“复刻 ChatGPT”。那就过了。

但如果只是想摸清 GPT 的骨架,这 200 行以内够了。输入变成 token,token 加位置编码,丢进多层因果注意力,最后预测下一个 token。循环这个过程,就成了生成。

大模型吓人的地方在规模,不在第一眼的代码。

代码拆到最后,其实就这几块。 能把这几块自己敲一遍,再去看那些几万行的训练框架,心里就不会那么虚了。