【大模型】- 差分注意力 v2

算法

差分注意力 v2

改进注意力机制的精确度

类型: 学习 | 语言: Python | 🏷 前置:《投机解码》(本系列第 14 篇)

学习目标

  • 理解差分注意力的原理
  • 实现差分注意力机制
  • 比较差分注意力与标准注意力
  • 优化差分注意力性能
  • 分析差分注意力的优势

差分注意力概述

差分注意力(Differential Attention)是一种改进的注意力机制,通过计算两个注意力分数的差值来提高精确度。

标准注意力的问题

标准自注意力存在以下问题:

  1. 注意力分散:注意力权重分布过于平滑
  2. 噪声干扰:容易受到不相关 token 的干扰
  3. 长距离依赖:难以捕捉长距离依赖关系

差分注意力的解决方案

差分注意力通过计算两个不同缩放因子的注意力分数的差值,来突出重要的 token 并抑制噪声。

差分注意力 v2 实现

核心算法

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
import torch
import torch.nn as nn
import torch.nn.functional as F
import math

class DifferentialAttentionV2(nn.Module):
def __init__(self, dim, n_heads, num_heads_per_group=1):
super().__init__()
self.dim = dim
self.n_heads = n_heads
self.head_dim = dim // n_heads
self.scale = self.head_dim ** -0.5

# Q, K, V 投影
self.wq = nn.Linear(dim, dim, bias=False)
self.wk = nn.Linear(dim, dim, bias=False)
self.wv = nn.Linear(dim, dim, bias=False)

# 差分缩放因子
self.lambda_q1 = nn.Parameter(torch.randn(n_heads, self.head_dim))
self.lambda_k1 = nn.Parameter(torch.randn(n_heads, self.head_dim))
self.lambda_q2 = nn.Parameter(torch.randn(n_heads, self.head_dim))
self.lambda_k2 = nn.Parameter(torch.randn(n_heads, self.head_dim))

# 输出投影
self.wo = nn.Linear(dim, dim, bias=False)

def forward(self, x, mask=None):
batch_size, seq_len, _ = x.shape

# 计算 Q, K, V
q = self.wq(x).view(batch_size, seq_len, self.n_heads, self.head_dim)
k = self.wk(x).view(batch_size, seq_len, self.n_heads, self.head_dim)
v = self.wv(x).view(batch_size, seq_len, self.n_heads, self.head_dim)

# 转置
q = q.transpose(1, 2) # (batch, n_heads, seq_len, head_dim)
k = k.transpose(1, 2)
v = v.transpose(1, 2)

# 计算两组注意力分数
attn1 = torch.matmul(q * self.lambda_q1, (k * self.lambda_k1).transpose(-2, -1))
attn2 = torch.matmul(q * self.lambda_q2, (k * self.lambda_k2).transpose(-2, -1))

# 缩放
attn1 = attn1 * self.scale
attn2 = attn2 * self.scale

# 应用掩码
if mask is not None:
attn1 = attn1.masked_fill(mask == 0, float("-inf"))
attn2 = attn2.masked_fill(mask == 0, float("-inf"))

# 计算差分注意力
attn = F.softmax(attn1, dim=-1) - F.softmax(attn2, dim=-1)

# 加权求和
output = torch.matmul(attn, v)

# 重塑
output = output.transpose(1, 2).contiguous()
output = output.view(batch_size, seq_len, -1)

return self.wo(output)

改进的缩放因子

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
class AdaptiveDifferentialAttention(nn.Module):
def __init__(self, dim, n_heads):
super().__init__()
self.dim = dim
self.n_heads = n_heads
self.head_dim = dim // n_heads

# 基础投影
self.wq = nn.Linear(dim, dim, bias=False)
self.wk = nn.Linear(dim, dim, bias=False)
self.wv = nn.Linear(dim, dim, bias=False)

# 自适应缩放因子
self.lambda_q1 = nn.Parameter(torch.randn(n_heads, self.head_dim))
self.lambda_k1 = nn.Parameter(torch.randn(n_heads, self.head_dim))
self.lambda_q2 = nn.Parameter(torch.randn(n_heads, self.head_dim))
self.lambda_k2 = nn.Parameter(torch.randn(n_heads, self.head_dim))

# 门控机制
self.gate = nn.Sequential(
nn.Linear(dim, dim // 4),
nn.GELU(),
nn.Linear(dim // 4, dim),
nn.Sigmoid()
)

self.wo = nn.Linear(dim, dim, bias=False)

def forward(self, x, mask=None):
batch_size, seq_len, _ = x.shape

# 计算基础注意力
q = self.wq(x).view(batch_size, seq_len, self.n_heads, self.head_dim)
k = self.wk(x).view(batch_size, seq_len, self.n_heads, self.head_dim)
v = self.wv(x).view(batch_size, seq_len, self.n_heads, self.head_dim)

q = q.transpose(1, 2)
k = k.transpose(1, 2)
v = v.transpose(1, 2)

# 计算差分注意力
attn1 = torch.matmul(q * self.lambda_q1, (k * self.lambda_k1).transpose(-2, -1))
attn2 = torch.matmul(q * self.lambda_q2, (k * self.lambda_k2).transpose(-2, -1))

attn = F.softmax(attn1, dim=-1) - F.softmax(attn2, dim=-1)

# 加权求和
output = torch.matmul(attn, v)

# 门控
gate = self.gate(x)
output = output.transpose(1, 2).contiguous()
output = output.view(batch_size, seq_len, -1)
output = output * gate

return self.wo(output)

与标准注意力比较

计算复杂度

1
2
3
4
5
6
7
8
9
10
11
12
13
def compare_complexity(seq_len, dim, n_heads):
"""比较计算复杂度"""
head_dim = dim // n_heads

# 标准注意力
standard_flops = 2 * seq_len * seq_len * dim + 2 * seq_len * dim * dim

# 差分注意力
differential_flops = 2 * (2 * seq_len * seq_len * dim + 2 * seq_len * dim * dim)

print(f"标准注意力 FLOPs: {standard_flops / 1e9:.2f}G")
print(f"差分注意力 FLOPs: {differential_flops / 1e9:.2f}G")
print(f"开销比: {differential_flops / standard_flops:.2f}x")

性能比较

1
2
3
4
5
6
7
8
9
10
11
12
13
14
def compare_performance(model_standard, model_differential, test_data):
"""比较性能"""
# 准确率
acc_standard = evaluate(model_standard, test_data)
acc_differential = evaluate(model_differential, test_data)

# 推理速度
time_standard = measure_inference_time(model_standard, test_data)
time_differential = measure_inference_time(model_differential, test_data)

print(f"标准注意力 - 准确率: {acc_standard:.4f}, 速度: {time_standard:.4f}s")
print(f"差分注意力 - 准确率: {acc_differential:.4f}, 速度: {time_differential:.4f}s")
print(f"准确率提升: {(acc_differential - acc_standard) / acc_standard:.2%}")
print(f"速度开销: {(time_differential - time_standard) / time_standard:.2%}")

优化实现

Flash 差分注意力

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
class FlashDifferentialAttention(nn.Module):
"""使用 Flash Attention 优化差分注意力"""

def __init__(self, dim, n_heads):
super().__init__()
self.dim = dim
self.n_heads = n_heads
self.head_dim = dim // n_heads

self.wq = nn.Linear(dim, dim, bias=False)
self.wk = nn.Linear(dim, dim, bias=False)
self.wv = nn.Linear(dim, dim, bias=False)

self.lambda_q1 = nn.Parameter(torch.randn(n_heads, self.head_dim))
self.lambda_k1 = nn.Parameter(torch.randn(n_heads, self.head_dim))
self.lambda_q2 = nn.Parameter(torch.randn(n_heads, self.head_dim))
self.lambda_k2 = nn.Parameter(torch.randn(n_heads, self.head_dim))

self.wo = nn.Linear(dim, dim, bias=False)

def forward(self, x, mask=None):
batch_size, seq_len, _ = x.shape

q = self.wq(x).view(batch_size, seq_len, self.n_heads, self.head_dim)
k = self.wk(x).view(batch_size, seq_len, self.n_heads, self.head_dim)
v = self.wv(x).view(batch_size, seq_len, self.n_heads, self.head_dim)

# 使用 Flash Attention
attn1 = F.scaled_dot_product_attention(
q * self.lambda_q1,
k * self.lambda_k1,
v,
is_causal=mask is None
)

attn2 = F.scaled_dot_product_attention(
q * self.lambda_q2,
k * self.lambda_k2,
v,
is_causal=mask is None
)

output = attn1 - attn2

output = output.reshape(batch_size, seq_len, -1)
return self.wo(output)

稀疏差分注意力

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
class SparseDifferentialAttention(nn.Module):
"""稀疏差分注意力,只关注部分位置"""

def __init__(self, dim, n_heads, sparsity=0.5):
super().__init__()
self.dim = dim
self.n_heads = n_heads
self.head_dim = dim // n_heads
self.sparsity = sparsity

self.wq = nn.Linear(dim, dim, bias=False)
self.wk = nn.Linear(dim, dim, bias=False)
self.wv = nn.Linear(dim, dim, bias=False)

self.lambda_q1 = nn.Parameter(torch.randn(n_heads, self.head_dim))
self.lambda_k1 = nn.Parameter(torch.randn(n_heads, self.head_dim))
self.lambda_q2 = nn.Parameter(torch.randn(n_heads, self.head_dim))
self.lambda_k2 = nn.Parameter(torch.randn(n_heads, self.head_dim))

self.wo = nn.Linear(dim, dim, bias=False)

def forward(self, x, mask=None):
batch_size, seq_len, _ = x.shape

q = self.wq(x).view(batch_size, seq_len, self.n_heads, self.head_dim)
k = self.wk(x).view(batch_size, seq_len, self.n_heads, self.head_dim)
v = self.wv(x).view(batch_size, seq_len, self.n_heads, self.head_dim)

q = q.transpose(1, 2)
k = k.transpose(1, 2)
v = v.transpose(1, 2)

# 计算注意力分数
attn_scores = torch.matmul(q, k.transpose(-2, -1))

# 稀疏化:只保留 top-k 分数
k = int(seq_len * self.sparsity)
topk_values, topk_indices = torch.topk(attn_scores, k, dim=-1)

# 创建稀疏掩码
sparse_mask = torch.zeros_like(attn_scores)
sparse_mask.scatter_(-1, topk_indices, 1.0)

# 应用差分注意力
attn1 = F.softmax(attn_scores * sparse_mask, dim=-1)
attn2 = F.softmax(attn_scores * sparse_mask * 0.5, dim=-1) # 不同缩放

attn = attn1 - attn2

output = torch.matmul(attn, v)
output = output.transpose(1, 2).contiguous()
output = output.view(batch_size, seq_len, -1)

return self.wo(output)

应用场景

长序列处理

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
# 差分注意力在长序列上的优势
def long_sequence_benchmark():
"""长序列基准测试"""
seq_lengths = [512, 1024, 2048, 4096, 8192]

for seq_len in seq_lengths:
# 创建测试数据
x = torch.randn(1, seq_len, 768)

# 标准注意力
standard_attn = MultiHeadAttention(768, 8)
start = time.time()
_ = standard_attn(x)
standard_time = time.time() - start

# 差分注意力
diff_attn = DifferentialAttentionV2(768, 8)
start = time.time()
_ = diff_attn(x)
diff_time = time.time() - start

print(f"序列长度 {seq_len}:")
print(f" 标准: {standard_time:.4f}s")
print(f" 差分: {diff_time:.4f}s")

代码生成

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
# 差分注意力在代码生成上的优势
def code_generation_test():
"""代码生成测试"""
# 代码具有复杂的依赖关系
# 差分注意力可以更好地捕捉这些关系

model_standard = load_model("standard_model")
model_differential = load_model("differential_model")

test_cases = [
"def fibonacci(n):",
"class Transformer(nn.Module):",
"SELECT * FROM users WHERE",
]

for test in test_cases:
# 生成代码
code_standard = generate(model_standard, test)
code_differential = generate(model_differential, test)

# 评估代码质量
quality_standard = evaluate_code_quality(code_standard)
quality_differential = evaluate_code_quality(code_differential)

print(f"提示: {test}")
print(f" 标准质量: {quality_standard:.2f}")
print(f" 差分质量: {quality_differential:.2f}")

总结

差分注意力 v2 通过计算两组注意力分数的差值,提高了注意力机制的精确度。它在长序列处理和复杂依赖关系建模上具有优势,但计算开销略高。优化实现(如 Flash 差分注意力)可以减少开销。

下一步

下一课将介绍原生稀疏注意力,进一步优化注意力计算。

📚 本文改编自 AI Engineering from Scratch(MIT License · 作者 Rohit Ghumare),中文内容来自官方中文镜像。原课程共 503 课 · 20 阶段 · 免费开源,教程网站见 aiengineeringfromscratch.com

  • 标题: 【大模型】- 差分注意力 v2
  • 作者:
  • 创建于 : 2026-08-19 09:16:00
  • 更新于 : 2026-08-21 16:20:11
  • 链接: https://sxl-space.tk/2026/08/19/010_LLM/010_LLM-16-DifferentialAttentionV2/
  • 版权声明: 版权所有 © 宋,禁止转载。