
调用一行
model.generate()只需一秒,手写一整个Transformer的前向传播却要一天——但正是这一天,让你真正理解了“注意力”为何能改变世界。
大模型时代,开发者被层层封装包围:HuggingFace的 AutoModel、PyTorch的 nn.Transformer、TensorFlow的 Keras……我们像熟练的API调酒师,却越来越像“盲人摸象”——模型能生成诗,但我们说不清它的每一层到底在做什么。
手写大模型,不是为了造轮子,而是为了拆轮子。 用纯Python和NumPy,从零搭建一个微型Transformer,你会看清:注意力机制不过是一系列矩阵乘法,前馈网络不过是两层线性变换,层归一化不过是减去均值除以方差。当你亲手敲下这些数学公式的代码时,黑盒开始透光。
这篇文章,我只展示最核心的五个构件,代码加起来不到80行。但每行都能让你在“框架舒适区”之外,触碰到深度学习的原始质感。
Attention Is All You Need——这句口号,在手写代码里只剩一行数学公式:
Attention(Q,K,V)=softmax(QKTdk)VAttention(Q,K,V)=softmax(dkQKT)V
用NumPy实现它,直白到近乎残忍:
import numpy as np
def scaled_dot_product_attention(Q, K, V, mask=None):
"""
Q, K, V: shape (batch, seq_len, d_k)
"""
d_k = Q.shape[-1]
scores = np.matmul(Q, K.transpose(0, 2, 1)) / np.sqrt(d_k) # (batch, seq_len, seq_len)
if mask is not None:
scores = scores + (mask * -1e9) # 将padding位置置为极小数
attn_weights = softmax(scores, axis=-1) # (batch, seq_len, seq_len)
output = np.matmul(attn_weights, V) # (batch, seq_len, d_k)
return output, attn_weights
def softmax(x, axis=-1):
e_x = np.exp(x - np.max(x, axis=axis, keepdims=True))
return e_x / np.sum(e_x, axis=axis, keepdims=True)这10行代码,是所有大模型的“原子核”。Q和K的点积衡量“相关性”,除以√d_k防止梯度消失,softmax归一化成概率,最后加权求和V。你看,没有魔法,只有矩阵乘法。
多头注意力(Multi-Head Attention)的本质,就是把Q、K、V分别投影到多个子空间,并行计算注意力,最后拼接。我们手写它,不过是为上述核心函数加一层循环:
class MultiHeadAttention:
def __init__(self, d_model, num_heads):
assert d_model % num_heads == 0
self.d_model = d_model
self.num_heads = num_heads
self.d_k = d_model // num_heads
# 初始化投影矩阵(实际训练中会学习)
self.W_Q = np.random.randn(d_model, d_model) * 0.01
self.W_K = np.random.randn(d_model, d_model) * 0.01
self.W_V = np.random.randn(d_model, d_model) * 0.01
self.W_O = np.random.randn(d_model, d_model) * 0.01
def forward(self, X, mask=None):
batch, seq_len, _ = X.shape
# 线性投影并拆分为多头
Q = X @ self.W_Q # (batch, seq, d_model)
K = X @ self.W_K
V = X @ self.W_V
# 重塑为 (batch, num_heads, seq, d_k)
Q = Q.reshape(batch, seq_len, self.num_heads, self.d_k).transpose(0, 2, 1, 3)
K = K.reshape(batch, seq_len, self.num_heads, self.d_k).transpose(0, 2, 1, 3)
V = V.reshape(batch, seq_len, self.num_heads, self.d_k).transpose(0, 2, 1, 3)
# 并行计算注意力(每个头共享同一个函数)
attn_out, _ = scaled_dot_product_attention(Q, K, V, mask)
# 合并多头
attn_out = attn_out.transpose(0, 2, 1, 3).reshape(batch, seq_len, self.d_model)
output = attn_out @ self.W_O
return output核心改动只有10多行,却让模型拥有了“同时关注局部和全局”的能力。手写这段代码时你会顿悟:多头不是魔法,只是把同一个数学操作在多个子空间里重复执行。
Transformer的FFN就是两层全连接 + ReLU,代码比MLP还简单:
class FeedForward:
def __init__(self, d_model, d_ff):
self.W1 = np.random.randn(d_model, d_ff) * 0.01
self.b1 = np.zeros(d_ff)
self.W2 = np.random.randn(d_ff, d_model) * 0.01
self.b2 = np.zeros(d_model)
def forward(self, X):
hidden = np.maximum(0, X @ self.W1 + self.b1) # ReLU
output = hidden @ self.W2 + self.b2
return output区区6行,定义了模型80%的参数量。它极其简单,却至关重要——注意力负责“发现关系”,FFN负责“转换表征”,两者缺一不可。
手写层归一化(LayerNorm)只需3行:
def layer_norm(X, eps=1e-6):
mean = np.mean(X, axis=-1, keepdims=True)
var = np.var(X, axis=-1, keepdims=True)
return (X - mean) / np.sqrt(var + eps)而残差连接则是一个加法:output = layer_norm(X + attn_output)。这些简单操作,让100层网络也能稳定训练。手写它们会让你敬畏:那些看起来高深的技巧,底层不过是加减乘除。
现在,我们把所有构件组合成一个完整的Transformer编码器层:
class EncoderLayer:
def __init__(self, d_model, num_heads, d_ff):
self.self_attn = MultiHeadAttention(d_model, num_heads)
self.ffn = FeedForward(d_model, d_ff)
def forward(self, X, mask=None):
# 多头注意力 + 残差 + LayerNorm
attn_out = self.self_attn.forward(X, mask)
X = layer_norm(X + attn_out)
# 前馈网络 + 残差 + LayerNorm
ffn_out = self.ffn.forward(X)
X = layer_norm(X + ffn_out)
return X这个极简的编码器,只有10多行,却能堆叠成GPT、BERT等大模型的骨架。当你亲眼看见这些代码时,会发现大模型并非高不可攀——它只是一叠精心组织的矩阵乘法。
讲完了代码,我想谈谈手写大模型真正改变我的地方:
1. 维度对齐的敏感度 手写时,你必须确保Q、K、V的维度在每一步都正确。一次reshape失误,整个矩阵乘法就会报错。这种训练,让你在调框架时也能一眼看出维度不匹配的问题。
2. 数值稳定性不是玄学
为什么softmax要减去最大值?为什么LayerNorm要加eps?手写时你亲身体验过np.exp(100)溢出inf的崩溃,从此你会敬畏每一行“防御性代码”。
3. 计算复杂度的直观感受
当你亲手写下Q @ K.T时,你会立刻意识到这是O(n2)O(n2)的复杂度。手写迫使你思考长序列的瓶颈所在,而不是机械地使用attn_dropout。
有人会问:既然生产环境用框架,手写有什么意义?我的答案是:手写是“逆向工程”的过程——你拆解了黑盒,才知道它的边界在哪里,才知道如何更有效地使用框架。
当你用PyTorch写nn.MultiheadAttention时,你知道它内部在做matmul和softmax;当你看到model.generate输出的“幻觉”时,你明白那是注意力权重在某个子空间发生了偏差;当你调优学习率时,你联想到手写时那些初始化矩阵的分布。
手写大模型,不是为了替代框架,而是为了让框架的每一行调用,都变成你心中早已演算过的确定动作。
全文核心代码不足80行。但就是这80行,构成了你理解所有大模型的“最小知识图谱”。下次当你面对一个数十亿参数的模型时,请记住:它的灵魂,不过就是这几行数学公式的叠加。知其然,更知其所以然——这才是手写的全部意义。
原创声明:本文系作者授权腾讯云开发者社区发表,未经许可,不得转载。
如有侵权,请联系 cloudcommunity@tencent.com 删除。