Math-03.线性代数-20.Attention与LoRA

本页从线性代数视角解析 Self-AttentionLoRA——读 Transformer / 蛋白 LM / 酶功能大模型代码时的核心矩阵结构。

段末注释Self-Attention(自注意力)用查询(query)与键(key)的相似度对值(value)加权聚合;Multi-Head Attention(多头注意力)在多个子空间并行做 attention。

系列入口00.系列规划 | 前置:02 矩阵运算05 SVD与LoRA


1. Self-Attention 矩阵形式(D2–D3)

图 1 QKV 线性投影

输入 $X \in \mathbb{R}^{L \times d_{\mathrm{model}}}$($L$ 序列长)。三个线性层(权重矩阵):

$$
Q = X W_Q, \quad K = X W_K, \quad V = X W_V
$$

$W_Q, W_K, W_V \in \mathbb{R}^{d_{\mathrm{model}} \times d_k}$(或 $d_k$ 每头维度)。

缩放点积 Attention

$$
\mathrm{Attention}(Q,K,V) = \mathrm{softmax}\left(\frac{Q K^\top}{\sqrt{d_k}}\right) V
$$

中间量 shape 含义
$QK^\top$ $L \times L$ token 两两相似度
$\mathrm{softmax}(\cdot)$ 行归一化 注意力权重
输出 $L \times d_k$ 加权值向量

$\sqrt{d_k}$ 缩放防止点积过大导致 softmax 饱和(梯度变小)。


2. Multi-Head(D3)

图 2 多头 = 并行低维 attention

$h$ 个头,每头 $d_k = d_{\mathrm{model}}/h$:

$$
\mathrm{head}_i = \mathrm{Attention}(X W_Q^{(i)}, X W_K^{(i)}, X W_V^{(i)})
$$

拼接后过输出投影:

$$
\mathrm{MultiHead}(X) = \mathrm{Concat}(\mathrm{head}_1,\ldots,\mathrm{head}_h) W_O
$$

参数量(一层 attention,忽略 bias):$4 d_{\mathrm{model}}^2$($W_Q,W_K,W_V,W_O$ 各约 $d^2$ 量级)。大模型层数 $\times$ 宽度 $d$ 决定总参数规模。


3. LoRA 低秩适配(D3–D7)

图 3 LoRA 增量

冻结 $W_0$,训练低秩增量(05 SVD/低秩):

$$
W = W_0 + \frac{\alpha}{r} B A, \quad B \in \mathbb{R}^{d_{\mathrm{out}} \times r},\ A \in \mathbb{R}^{r \times d_{\mathrm{in}}}
$$

前向:

$$
\mathbf{y} = W_0 \mathbf{x} + \frac{\alpha}{r} B A \mathbf{x}
$$

项目 全量微调 LoRA 秩 $r$
可训练参数(单层) $d_{\mathrm{out}} d_{\mathrm{in}}$ $r(d_{\mathrm{out}}+d_{\mathrm{in}})$
例:768→768, $r=8$ 589,824 12,288

常对 Attention 的 $W_Q,W_K,W_V,W_O$FFN 注入 LoRA。蛋白/酶 LM 微调见 酶功能大模型酶改造-06 ML 实践


4. 与线性代数概念对照(D6)

概念 Attention / LoRA
矩阵乘 $QK^\top$, $W_0 x$, $BAx$
LoRA 设 $\mathrm{rank}(\Delta W) \le r$
范数 权重 decay 作用于 $A,B$
Softmax 非线性,打破纯线性;梯度经链式回传(Math-06

5. 局限与工程注意(D8)

图 4 局限

问题 说明
$L$ 大时 $QK^\top$ 为 $L^2$ FlashAttention 分块;长序列贵
LoRA $r$ 过小 任务表达不足
只 LoRA 部分层 需消融哪层最有效
推理合并权重 可合并 $W_0 + BA$ 无额外延迟

6. PyTorch 形状示例(D12)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
import torch
import torch.nn as nn
import math

B, L, d_model, h, d_k = 2, 128, 512, 8, 64
x = torch.randn(B, L, d_model)

# 简化单头 attention
Wq = nn.Linear(d_model, d_k, bias=False)
Wk = nn.Linear(d_model, d_k, bias=False)
Wv = nn.Linear(d_model, d_k, bias=False)

Q, K, V = Wq(x), Wk(x), Wv(x) # (B,L,d_k)
attn = torch.softmax(Q @ K.transpose(-2,-1) / math.sqrt(d_k), dim=-1)
out = attn @ V # (B,L,d_k)

# LoRA 秩-r 增量
r = 8
d_in, d_out = 512, 512
A = nn.Parameter(torch.randn(r, d_in) * 0.01)
B = nn.Parameter(torch.zeros(d_out, r))
x2 = torch.randn(B, d_in)
delta = (x2 @ A.T) @ B.T # (B, d_out)

7. 小结

Attention = 三次线性投影 + $(QK^\top)V$ 的加权聚合;LoRA = 低秩矩阵乘叠加到冻结权重。Math-03 系列至此覆盖 ML/DL 线性代数主线;训练动力学见 Math-04 优化Math-05 信息论

系列导航05 SVD/LoRA 数学 | 10 PCA | 00 规划

-------------本文结束感谢您的阅读-------------