Transformer架构的数学框架解析:从自注意力机制到电路建模

1次阅读
没有评论

共计 1695 个字符,预计需要花费 5 分钟才能阅读完成。

image.webp

Transformer 模型彻底改变了自然语言处理领域,其自注意力机制能够动态捕捉长距离依赖关系。相比 RNN 的序列计算,Transformer 的并行化架构大幅提升了训练效率。如今从 BERT 到 GPT 系列,Transformer 已成为现代 NLP 任务的基石架构。

Transformer 架构的数学框架解析:从自注意力机制到电路建模

自注意力机制的数学本质

自注意力机制的核心是三个权重矩阵:$W_Q\in\mathbb{R}^{d\times d_k}$(Query)、$W_K\in\mathbb{R}^{d\times d_k}$(Key)和 $W_V\in\mathbb{R}^{d\times d_v}$(Value)。在向量空间中:

  • Query 向量代表当前需要关注的内容
  • Key 向量相当于所有位置的索引标签
  • Value 向量存储实际要提取的信息

计算过程可表示为:

$$
\text{Attention}(Q,K,V)=\text{softmax}(\frac{QK^T}{\sqrt{d_k}})V
$$

几何意义是:通过 QK^T 点积计算向量相似度,softmax 将其转化为概率分布,最后对 V 加权求和。

多头注意力的并行化实现

多头注意力(Multi-Head Attention)将输入拆分为 $h$ 个子空间:

  1. 输入矩阵 $X$ 分别乘以 $h$ 组不同的 $W_Q^i,W_K^i,W_V^i$
  2. 每个头计算独立的注意力结果 $\text{head}_i=\text{Attention}(XW_Q^i,XW_K^i,XW_V^i)$
  3. 所有头的结果拼接后通过 $W_O\in\mathbb{R}^{hd_v\times d}$ 输出

矩阵分块运算示意图:

# PyTorch 实现示例
assert x.dim() == 3  # [batch, seq_len, dim]
q = torch.chunk(x @ W_Q, num_heads, dim=-1)  # 分块
k = torch.chunk(x @ W_K, num_heads, dim=-1)
v = torch.chunk(x @ W_V, num_heads, dim=-1)
# 各头独立计算
outputs = [attention(q[i], k[i], v[i]) for i in range(num_heads)]

位置编码的傅里叶分析

位置编码(Positional Encoding)采用正弦函数组合:

$$
PE(pos,2i)=\sin(pos/10000^{2i/d})
$$
$$
PE(pos,2i+1)=\cos(pos/10000^{2i/d})
$$

这实质上是傅里叶级数的离散采样:

  • 不同 $i$ 对应不同频率分量
  • $2i$ 和 $2i+1$ 构成正交基
  • 位置变化体现为相位移动

Transformer 的电路建模

信息流动有向图

将 Transformer 视为计算电路时:

  • 节点代表张量运算(矩阵乘 /softmax 等)
  • 边表示数据流动方向
  • 残差连接形成环路结构

残差连接的梯度分析

与传统 DNN 相比:

$$
\frac{\partial L}{\partial x}=\frac{\partial L}{\partial F(x)}+\frac{\partial L}{\partial (F(x)+x)}
$$

残差结构确保梯度至少有一条恒定通路,缓解梯度消失。

性能优化实践

稀疏注意力存储

对于稀疏注意力矩阵 $A$,采用 CSR 格式存储:

values = A[A != 0]
indices = (A != 0).nonzero()

FlashAttention 对比

算法 时间复杂度 空间复杂度
原始注意力 O(N^2) O(N^2)
FlashAttention O(N^2/d) O(N)

数学验证挑战

问题 :推导 LayerNorm 如何缓解梯度消失

提示:考虑对输入 $X$ 的归一化:
$$
Y=\frac{X-\mu}{\sigma}\cdot\gamma+\beta
$$

参考答案:
$$
\frac{\partial Y}{\partial X} = \frac{\gamma}{\sigma}\left(I – \frac{1}{n}\mathbf{1}\mathbf{1}^T – \frac{(X-\mu)(X-\mu)^T}{n\sigma^2}\right)
$$

通过保持梯度幅度稳定,LayerNorm 能有效抑制梯度指数级衰减。

正文完
 0
评论(没有评论)