Wall Attention(GitHub 仓库)

TLDR AI 论文

摘要

Wall Attention 是一种新的注意力变体,具有每个通道、每个时间步的乘法衰减,提供内容相关的遗忘率,以及在Triton中实现的高效训练/解码内核。

一种通过围绕持久性“wall”记忆令牌组织信息来改善长上下文推理的注意力机制。
查看原文
查看缓存全文

缓存时间: 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 头数 HHQ % 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

相似文章

Random Attention(GitHub 仓库)

TLDR AI

Random Attention 提出了一种无信号的 KV 缓存驱逐策略,适用于推理模型,在 MATH-500 和 LiveCodeBench 等基准测试中,其性能匹配或超越了学习方法,同时推理速度更快。

在推理阶段为预训练大语言模型应用滑动窗口注意力 [P]

Reddit r/MachineLearning

本项目为 Hugging Face 预训练大语言模型实现了一个可复用的滑动窗口注意力推理层,利用带 attention sinks 的受限 KV 缓存,显著降低显存占用并加快解码速度。在 Qwen2.5-7B 上的测试表明,16K 上下文下显存从约 923 MB 降至约 3.5 MB,不过依赖远距离上下文的任务可能会出现性能下降。