Wall Attention(GitHub 仓库)
摘要
Wall Attention 是一种新的注意力变体,具有每个通道、每个时间步的乘法衰减,提供内容相关的遗忘率,以及在Triton中实现的高效训练/解码内核。
查看缓存全文
缓存时间: 2026/06/03 15:35
tilde-research/wall-attention-release
来源:https://github.com/tilde-research/wall-attention-release
Wall Attention
Wall Attention 是一种注意力机制的变体,它在 QK 内积中融入了每个通道、每个时间步的乘法衰减。标准的注意力对一对 (i, j) 的评分是通过 \sum_n q_{i,n}\, k_{j,n} 计算,而 Wall Attention 对每个通道 n 通过两个位置之间累积的学习衰减进行加权。这使得每个查询通道拥有独立、与内容相关的遗忘率,将标量门控(FoX)和 RoPE 风格的衰减推广到了整个通道维度。当 g = 0 时,退化为标准的 softmax 注意力。
更多信息请参见博客: https://blog.tilderesearch.com/blog/wall-attn
本仓库打包了实际使用的两个内核,各自独立:
- 训练/预填充(
wall_attn):一个融合的前向+反向 Triton 内核(FlashAttention 风格的流式 softmax),包含针对 q, k, v, g 的解析梯度。 - 解码(
wall_attn_decode):一个单步内核,读取预缩放的 KV 缓存,因此每个 token 生成只需要一次小型的 GEMV 操作,而无需重新计算前缀。
安装
# 使用 uv(推荐)
uv sync
source .venv/bin/activate
# 或使用 pip
pip install -e .
使用
训练/预填充
import torch
from wall_attn import wall_attn
B, T, H, HQ, K, V = 2, 1024, 4, 8, 64, 64 # GQA: HQ 个查询头,H 个 KV 头
q = torch.randn(B, T, HQ, K, device="cuda", dtype=torch.bfloat16, requires_grad=True)
k = torch.randn(B, T, H, K, device="cuda", dtype=torch.bfloat16, requires_grad=True)
v = torch.randn(B, T, H, V, device="cuda", dtype=torch.bfloat16, requires_grad=True)
g = torch.randn(B, T, HQ, K, device="cuda", dtype=torch.bfloat16, requires_grad=True) * 0.02
o = wall_attn(q, k, v, g, scale=K**-0.5) # [B, T, HQ, V]
o.sum().backward()
可选参数:g_scalar([B, T, HQ] FoX 风格加性门控)、sink_bias([HQ] 注意力下沉)、window_size(滑动窗口)和 cu_seqlens(变长打包,要求 B == 1)。
解码(缓存生成)
在预填充时一次性构建预缩放的缓存,然后逐 token 解码:
import torch
from fla.ops.utils.constant import RCP_LN2
from fla.ops.utils.cumsum import chunk_global_cumsum
from wall_attn import build_wall_kv_cache, wall_attn_decode
C = 64 # 缓存块大小(锚点粒度)
P = chunk_global_cumsum(g, scale=RCP_LN2) # [B, T, HQ, K] 前缀
k_tilde, r_cache = build_wall_kv_cache(k, P, chunk_size=C)
o, _ = wall_attn_decode(
q=q[:, -1:], # 当前查询 [B, 1, HQ, K]
v=v, # 缓存的值 [B, T_kv, H, V]
p_curr=P[:, -1:], # 当前行的前缀
k_tilde=k_tilde, # 预缩放的键 [B, T_kv, HQ, K]
r_cache=r_cache, # 每个块的锚点 [B, ceil(T_kv/C), HQ, K]
sink_bias=None,
scale=K**-0.5,
cache_chunk_size=C,
)
build_wall_kv_cache 将衰减融入到键中(k_tilde[j] = k[j] · exp2(R_c − P[j])),使用每个块的锚点 R_c,因此解码内核永远不会重新累积前缀。完整的追加式服务循环请参见 tests/test_decode.py::test_decode_streaming_matches_full_forward。
代码结构
wall_attn/
├── __init__.py # 公共 API
├── training.py # 前向/反向 Triton 内核 + autograd Function + wall_attn()
├── decode.py # 单步解码内核 + build_wall_kv_cache()
└── reference.py # 急切 PyTorch 参考实现(正确性基准)
tests/
├── test_training.py # 一致性 + 解析梯度(通过有限差分验证)
└── test_decode.py # 解码 == 预填充前向,流式,缓存形状
特性
- GQA:查询头数
HQ可以超过 KV 头数H(HQ % H == 0)。 - 逐通道衰减
g,包含精确的解析梯度,以及可选的标量门控g_scalar。 - 注意力下沉(
sink_bias)、滑动窗口(window_size)和变长打包(cu_seqlens)。 - 预缩放的解码缓存,用于低成本的自回归生成,在长上下文下数值稳定(每个块的锚点保持
exp2有界)。 - BF16/FP32 输入;针对 Hopper / Ampere 自动调整块大小。
测试
pytest # 需要 CUDA GPU
每个内核路径都与急切实现的 wall_attn_reference 进行对比检查,并且 g / g_scalar 的梯度通过中心有限差分进行验证。解码内核在逐 token 的基础上与训练前向结果一致,包括流式生成循环。
致谢
Triton 内核基于 flash-linear-attention(https://github.com/fla-org/flash-linear-attention)(MIT)中的并行注意力机制构建。我们感谢 FLA 团队在高效注意力方面所做的出色工作。
许可证
MIT,详见 LICENSE。
相似文章
@tilderesearch: https://x.com/tilderesearch/status/2061771450168889432
Wall Attention 将对角遗忘门泛化到 softmax 注意力,实现了从 4k 到 160k+ 上下文的零样本最先进长度外推,并且在预训练中优于 RoPE 和 FoX。它作为即插即用的替换方案发布,附带开源的 Triton 内核。
我构建了一种新的注意力机制(Wave Field)——可在标准注意力机制内存溢出的128K上下文环境下运行,笔记本CPU上实现80+ tok/s
一位独立研究员引入了Wave Field注意力机制,用FFT波卷积替代了标准O(N²)点积注意力,实现了O(N log N)的训练效率和每token O(1)的推理效率。声称在128K上下文的CPU上达到80+ tok/s,并且在零样本性能上优于GPT-2 124M。
Random Attention(GitHub 仓库)
Random Attention 提出了一种无信号的 KV 缓存驱逐策略,适用于推理模型,在 MATH-500 和 LiveCodeBench 等基准测试中,其性能匹配或超越了学习方法,同时推理速度更快。
@thtrkim: FlashAttention 的手动可视化深入讲解(使用 Excalidraw 绘制)https://winterrykim.github.io/blog/2026/training-lm-…
深入理解 FlashAttention 的可视化讲解,涵盖内存优化和算子融合,以实现语言模型训练中的高效注意力计算。
在推理阶段为预训练大语言模型应用滑动窗口注意力 [P]
本项目为 Hugging Face 预训练大语言模型实现了一个可复用的滑动窗口注意力推理层,利用带 attention sinks 的受限 KV 缓存,显著降低显存占用并加快解码速度。在 Qwen2.5-7B 上的测试表明,16K 上下文下显存从约 923 MB 降至约 3.5 MB,不过依赖远距离上下文的任务可能会出现性能下降。