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 65 66 67 68 69 70 71 72 73 74 75
| import torch import torch.nn as nn import torch.nn.functional as F
class SparseAttention(nn.Module): """ 稀疏注意力机制 论文方法:只计算top-k重要注意力权重 优势: - 减少计算量 O(n²) → O(n·k) - 保持关键注意力模式 """ def __init__(self, dim: int, num_heads: int = 8, sparsity_ratio: float = 0.5): super().__init__() self.num_heads = num_heads self.head_dim = dim // num_heads self.sparsity_ratio = sparsity_ratio self.qkv = nn.Linear(dim, dim * 3) self.proj = nn.Linear(dim, dim) def forward(self, x: torch.Tensor) -> torch.Tensor: """ 稀疏注意力前向传播 Args: x: 输入序列, shape=(B, N, D) Returns: output: shape=(B, N, D) """ B, N, D = x.shape qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim) q, k, v = qkv.unbind(2) q = q.transpose(1, 2) k = k.transpose(1, 2) v = v.transpose(1, 2) attn = torch.matmul(q, k.transpose(-2, -1)) / (self.head_dim ** 0.5) k_sparse = int(N * self.sparsity_ratio) topk_values, topk_indices = torch.topk(attn, k_sparse, dim=-1) sparse_attn = torch.zeros_like(attn) sparse_attn.scatter_(-1, topk_indices, topk_values) sparse_attn = F.softmax(sparse_attn, dim=-1) output = torch.matmul(sparse_attn, v) output = output.transpose(1, 2).reshape(B, N, D) return self.proj(output)
if __name__ == "__main__": sparse_attn = SparseAttention(dim=768, sparsity_ratio=0.3) x = torch.randn(2, 100, 768) output = sparse_attn(x) print(f"稀疏注意力输出: {output.shape}") print(f"稀疏比例: 70% 权重被剪枝")
|