周瑜零基础手写agent

手写大模型:从零构建一个"微型Transformer"的底层逻辑

一、引言:大模型不是"魔法",是"矩阵乘法"

当你在对话框里向ChatGPT提问时,它背后那个庞大的神经网络正在毫秒级地执行着数十亿次矩阵乘法运算。很多人将大模型视为"黑盒魔法",但事实上,它的核心数学逻辑完全可以用不到200行代码表达出来

当然,我们不可能在单篇文章里复现GPT-4的万亿参数,但我们可以手写一个 "微型大模型"------一个完整的Transformer解码器(Decoder-only)架构 ,它具备多头注意力、前馈网络、层归一化和因果掩码等所有核心机制,并且能在你的笔记本电脑上真正运行(训练或推理)。当你亲手敲下Attention类的代码、看着模型根据输入"哼哧哼哧"地逐字生成文本时,你会恍然大悟:原来所谓的"大模型",不过是一系列精巧的张量变换加上海量的数据喂养。

本文将带你用Python + PyTorch (作为底层张量计算库,避免陷入纯NumPy反向传播的泥潭),从零手写一个微型GPT,并让它学习简单的规律(比如加法算术或莎士比亚风格语句)。我们只写架构,不调现成的transformers,以此揭开大模型最内核的面纱。


二、核心基石:缩放点积注意力(Scaled Dot-Product Attention)

大模型区别于传统RNN/LSTM的核心,就是注意力机制(Attention) 。它允许模型在生成每个词时,"关注"输入序列中所有位置的相关信息。

我们首先实现最底层的 "缩放点积注意力" 函数,这是整个Transformer的"心脏"。

python 复制代码
import torch
import torch.nn as nn
import torch.nn.functional as F

def scaled_dot_product_attention(query, key, value, mask=None, dropout=None):
    """
    参数:
        query: (batch_size, num_heads, seq_len, d_k)
        key:   (batch_size, num_heads, seq_len, d_k)
        value: (batch_size, num_heads, seq_len, d_v)
        mask:  (batch_size, 1, 1, seq_len) 或 (seq_len, seq_len)
    返回:
        加权后的输出及注意力权重
    """
    d_k = query.size(-1)
    
    # 1. 计算点积注意力分数 (Q · K^T)
    scores = torch.matmul(query, key.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k, dtype=torch.float32))
    
    # 2. 应用掩码(Mask)------确保模型不会"偷看"未来的词(因果掩码)
    if mask is not None:
        # 将掩码位置填充为极小的负数,使得softmax后概率趋近于0
        scores = scores.masked_fill(mask == 0, float('-inf'))
    
    # 3. Softmax归一化(得到注意力权重)
    attention_weights = F.softmax(scores, dim=-1)
    
    # 4. Dropout(训练时使用,防止过拟合)
    if dropout is not None:
        attention_weights = dropout(attention_weights)
    
    # 5. 加权求和
    output = torch.matmul(attention_weights, value)
    
    return output, attention_weights

这段代码是大模型中最核心的数学操作。你可以看到:没有复杂的逻辑,只有矩阵乘法和Softmax


三、多头注意力(Multi-Head Attention):从"单视角"到"多视角"

单头注意力只做一次加权,视野受限。多头注意力querykeyvalue投影到多个子空间,并行执行多次上述注意力计算,最后拼接起来。这让模型能从不同维度理解语义。

ini 复制代码
class MultiHeadAttention(nn.Module):
    def __init__(self, embed_size, num_heads):
        super().__init__()
        assert embed_size % num_heads == 0, "嵌入维度必须能被头数整除"
        
        self.embed_size = embed_size
        self.num_heads = num_heads
        self.d_k = embed_size // num_heads  # 每个头的维度
        
        # 定义线性投影层(分别对应Q、K、V)
        self.W_q = nn.Linear(embed_size, embed_size)
        self.W_k = nn.Linear(embed_size, embed_size)
        self.W_v = nn.Linear(embed_size, embed_size)
        self.W_o = nn.Linear(embed_size, embed_size)  # 最终输出投影
        
        self.dropout = nn.Dropout(0.1)
    
    def forward(self, x, mask=None):
        batch_size, seq_len, _ = x.size()
        
        # 1. 线性投影并拆分为多头 (batch, seq_len, embed) -> (batch, num_heads, seq_len, d_k)
        Q = self.W_q(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2)
        K = self.W_k(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2)
        V = self.W_v(x).view(batch_size, seq_len, self.num_heads, self.d_k).transpose(1, 2)
        
        # 2. 调用核心缩放点积注意力
        attn_output, _ = scaled_dot_product_attention(Q, K, V, mask, self.dropout)
        
        # 3. 合并多头 (batch, num_heads, seq_len, d_k) -> (batch, seq_len, embed_size)
        attn_output = attn_output.transpose(1, 2).contiguous().view(batch_size, seq_len, self.embed_size)
        
        # 4. 最终线性映射
        output = self.W_o(attn_output)
        return output

四、前馈网络(Feed-Forward)与层归一化(LayerNorm)

注意力层之后,通常会接一个简单的全连接前馈网络(FFN) ,引入非线性变换。同时,每个子层(Attention和FFN)都套用了残差连接(Residual)层归一化(LayerNorm) ,这是保证模型能够深度堆叠而不退化的关键。

python 复制代码
class FeedForward(nn.Module):
    def __init__(self, embed_size, hidden_dim_multiplier=4):
        super().__init__()
        hidden_dim = embed_size * hidden_dim_multiplier  # 通常扩展4倍
        self.fc1 = nn.Linear(embed_size, hidden_dim)
        self.fc2 = nn.Linear(hidden_dim, embed_size)
        self.gelu = nn.GELU()  # 较ReLU更平滑的激活函数
    
    def forward(self, x):
        return self.fc2(self.gelu(self.fc1(x)))

class TransformerBlock(nn.Module):
    """一个完整的Transformer解码器块"""
    def __init__(self, embed_size, num_heads, dropout=0.1):
        super().__init__()
        self.attention = MultiHeadAttention(embed_size, num_heads)
        self.feed_forward = FeedForward(embed_size)
        self.norm1 = nn.LayerNorm(embed_size)
        self.norm2 = nn.LayerNorm(embed_size)
        self.dropout = nn.Dropout(dropout)
    
    def forward(self, x, mask=None):
        # 多头注意力 + 残差连接 + 层归一化
        attn_output = self.attention(x, mask)
        x = self.norm1(x + self.dropout(attn_output))
        
        # 前馈网络 + 残差连接 + 层归一化
        ff_output = self.feed_forward(x)
        x = self.norm2(x + self.dropout(ff_output))
        return x

五、词嵌入与位置编码(让模型理解"顺序")

大模型不认识文字,只认识数字。我们需要词嵌入(Token Embedding) 将词ID映射为稠密向量。同时,由于Attention本身不具备时序概念,我们需要注入位置编码(Positional Encoding) 来告诉模型词的先后顺序。

这里我们采用可学习的位置编码(与原始Transformer的三角函数固定编码不同,现代模型如GPT多用可学习嵌入)。

ruby 复制代码
class PositionalEncoding(nn.Module):
    def __init__(self, max_seq_len, embed_size):
        super().__init__()
        # 可学习的位置参数
        self.pos_embedding = nn.Parameter(torch.randn(1, max_seq_len, embed_size))
    
    def forward(self, x):
        seq_len = x.size(1)
        return x + self.pos_embedding[:, :seq_len, :]

六、组装我们的"微型GPT"(MiniGPT)

将上述所有组件堆叠起来,再加上一个最终的输出层(将隐状态映射回词表大小),就构成了一个完整的自回归语言模型

ini 复制代码
class MiniGPT(nn.Module):
    def __init__(self, vocab_size, embed_size=256, num_heads=8, num_layers=6, max_seq_len=512, dropout=0.1):
        super().__init__()
        self.token_embedding = nn.Embedding(vocab_size, embed_size)
        self.position_encoding = PositionalEncoding(max_seq_len, embed_size)
        self.dropout = nn.Dropout(dropout)
        
        # 堆叠多个Transformer块
        self.blocks = nn.ModuleList([
            TransformerBlock(embed_size, num_heads, dropout) for _ in range(num_layers)
        ])
        
        # 最终的LayerNorm + 线性输出层
        self.ln_f = nn.LayerNorm(embed_size)
        self.head = nn.Linear(embed_size, vocab_size)
        
        # 初始化参数(对训练稳定性很重要)
        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)
    
    def forward(self, idx, mask=None):
        # idx: (batch, seq_len)  输入序列的Token ID
        x = self.token_embedding(idx)  # (batch, seq_len, embed)
        x = self.position_encoding(x)
        x = self.dropout(x)
        
        # 生成因果掩码(确保第i个位置只能看到前i个位置)
        if mask is None:
            mask = torch.tril(torch.ones(1, 1, idx.size(1), idx.size(1))).to(idx.device)
        
        for block in self.blocks:
            x = block(x, mask)
        
        x = self.ln_f(x)
        logits = self.head(x)  # (batch, seq_len, vocab_size)
        return logits

七、让模型"活"起来:训练与推理的极简示例

接下来,我们用一个极简的"加法算术"任务来验证手写模型是否真的能学习。我们生成"a+b=c"格式的数据(例如"12+34=46"),让模型学习做加法。

ini 复制代码
# ---------- 数据准备(字符级分词) ----------
chars = sorted(list(set("0123456789+= ")))
vocab_size = len(chars)
stoi = {ch: i for i, ch in enumerate(chars)}
itos = {i: ch for i, ch in enumerate(chars)}
encode = lambda s: [stoi[c] for c in s]
decode = lambda l: ''.join([itos[i] for i in l])

# 生成训练数据:1000条两位数加法
import random
train_data = []
for _ in range(1000):
    a = random.randint(10, 99)
    b = random.randint(10, 99)
    c = a + b
    text = f"{a}+{b}={c}"
    train_data.append(encode(text))

# 固定序列长度(补齐到最大长度)
max_len = max(len(d) for d in train_data)
padded_data = torch.tensor([d + [0]*(max_len - len(d)) for d in train_data])  # 0为padding

# ---------- 训练循环(CPU可跑) ----------
model = MiniGPT(vocab_size=vocab_size, embed_size=64, num_heads=4, num_layers=4, max_seq_len=max_len)
criterion = nn.CrossEntropyLoss(ignore_index=0)  # 忽略padding
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)

for epoch in range(50):
    logits = model(padded_data)  # (batch, seq, vocab)
    # 将输出展平为 (batch*seq, vocab),标签展平为 (batch*seq)
    loss = criterion(logits.view(-1, vocab_size), padded_data.view(-1))
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()
    if epoch % 10 == 0:
        print(f"Epoch {epoch}, Loss: {loss.item():.4f}")

# ---------- 推理:模型会做加法吗? ----------
model.eval()
test_input = encode("35+20=")
test_tensor = torch.tensor([test_input + [0]*(max_len - len(test_input))])
with torch.no_grad():
    logits = model(test_tensor)
    # 贪心解码获取下一个token
    next_token = logits.argmax(dim=-1)[0, len(test_input)-1].item()
    print(f"输入: 35+20=, 模型预测下一个字符: {itos[next_token]}")  # 期望输出: '5' (因为35+20=55)

当你运行这段代码,经过几十个epoch的迭代后,你会发现模型真的学会了进位和加法------尽管它只有不到100万参数。这就是大模型的雏形


八、为什么"手写"如此重要?

  1. 破除"神秘感" :当你亲手写完scaled_dot_product_attention后,你会理解所谓的"智能"不过是矩阵乘法和非线性激活函数的组合。
  2. 调试能力跃升 :如果生产环境的大模型表现异常,手写过架构的工程师能更快定位到是LayerNorm位置不对,还是mask没传对。
  3. 架构创新的基础:所有最新的MoE(混合专家)、Mamba(状态空间模型)都是在这些基础组件上做替换。不懂底层,就无法理解前沿论文。

九、从"微型"到"巨型":我们还缺什么?

当然,手写的MiniGPT离真正的生产级大模型(如GPT-4、Claude)还有天壤之别。主要差距在于:

维度 手写MiniGPT 生产级大模型
参数量 ~100万 数千亿(~10^12)
训练数据 1000条合成文本 数万亿Tokens(互联网语料)
并行训练 单卡CPU 数万张GPU(张量并行/流水线并行)
优化技术 基础SGD 混合精度训练、ZeRO优化、Flash Attention
分词器 简单字符级 BPE(字节对编码)或SentencePiece

架构的骨骼是完全一致的 。你手写的TransformerBlock,放在GPT-4中只是将其中的embed_size=256改为embed_size=12288num_layers=6改为num_layers=96,并加上数千倍的数据量罢了。


十、结语:写下第一个矩阵乘法的那一刻,你就超过了99%的"调包侠"

"手写大模型" 不是为了在生产环境中替代成熟的框架,而是为了重建开发者的认知路径。当你亲手调通scaled_dot_product_attention的反向传播、看着loss曲线平滑下降时,大模型在你眼中将不再是遥不可及的神祇,而是一台可以拆解、可以优化、可以批判的精密机器。

在此之后,再去看transformers库源码、再去读《Attention Is All You Need》论文,你会发现字字珠玑------因为每一行公式,你都在代码里亲手"踩过坑"。

动手建议 :将本文的代码复制到本地,把embed_size增大到128,把num_layers增加到8,然后在网上找一份小的莎士比亚文本数据集训练。看着模型从输出乱码到慢慢拼出"To be, or not to be"的那一刻,你会明白工程与艺术交汇的狂喜。

相关推荐
闪学it26 分钟前
大模型与Agent智能体开发实战
agent
尘中远32 分钟前
7大开源Agent源码对比解读——架构对比
ai·开源·agent·codex·harness
潘锦1 小时前
Session First 和 Agent First,本身并没有高低之分
agent
张忠琳2 小时前
【deepseek-harness】Cordis 时空可组合性编程范式 — 三段式精读笔记(二)
ai·agent·deepseek·harness·cordis·dsh
小沈同学呀3 小时前
【Agent开发第七期】记忆系统:让 Agent 跨会话也记得你
人工智能·agent·长期记忆·记忆系统
哥不是小萝莉3 小时前
AI 编码 Agent 从原理到可运行代码
ai·agent
新知图书4 小时前
16.3 基于MCP的多Agent旅行规划助手项目结构
人工智能·agent·ai agent·智能体
似水流年QC5 小时前
什么是 Skill?深入理解 AI Agent Skill 的工作原理与应用实践
人工智能·agent·skill