每 4 层插 3 次线性注意力:KDA 与 Gated Attention 的混合账本
【免费下载链接】AliceAI-Foundation-80B-A3B-Base项目地址: https://ai.gitcode.com/hf_mirrors/yandex/AliceAI-Foundation-80B-A3B-Base
Yandex 开源的 AliceAI-Foundation-80B-A3B-Base 是一个 80B 总参数、每个 token 仅激活 3B 的稀疏 MoE 基座模型,但真正让它区别于同代 A3B 模型的,不是 MoE 路由,而是注意力层的排布方式:48 层 Transformer 被切分成 12 个"4 层块",每个块内前 3 层是 KDA 线性注意力,第 4 层是 Gated Attention 全注意力。这意味着模型里 75% 的注意力层是线性复杂度的,只有 25% 保留标准 softmax 注意力,而它依然宣称支持 262144 token 的长上下文。
这套"混合账本"是怎么记的、为什么要按 3:1 配比记账,本文直接从仓库源码逐行拆解。
一、先翻开账本:48 层的真实排布
注意力层的配比不是口头上的,config.json里用 48 个字符串的layer_types数组把每一层的类型写得明明白白:
"layer_types": [ "linear_attention", "linear_attention", "linear_attention", "full_attention", "linear_attention", "linear_attention", "linear_attention", "full_attention", ... ], "num_hidden_layers": 48, "block_attn_res_block_size": 4规律非常整齐:每 4 层一组,前 3 层linear_attention,第 4 层full_attention,共 12 组、48 层。这个排布在 configuration_alice_ai.py 里还有一层程序化默认值可以对照:
if layer_types is None: layer_types = [ "full_attention" if (layer_idx + 1) % 4 == 0 else "linear_attention" for layer_idx in range(num_hidden_layers) ]也就是说,(layer_idx + 1) % 4 == 0的层是全注意力,其余全是线性注意力。而block_attn_res_block_size = 4同时控制着残差混合的粒度——下文会看到,这个"4"不是巧合。
注意一个关键点:这里的"线性注意力"并非 Mamba 式的 SSM 或普通的线性注意力近似,而是名为 KDA(Kernel Delta Attention 一类带 delta 状态更新的线性注意力变体)的实现,配置里还专门给出了它的参数:32 个 query 头、32 个 KV 头、head dim 128、因果卷积核大小 4(linear_conv_kernel_dim)。KDA 由社区文章与源码双重印证:config.json中linear_conv_kernel_dim: 4、linear_key_head_dim: 128、linear_num_key_heads: 32等键位与 modeling_alice_ai.py 中AliceAIKDA类的实现一一对应。
二、KDA 线性注意力:一张只记"流水账"的压缩账本
AliceAIKDA的 forward 过程在 modeling_alice_ai.py 中非常直观。它把隐状态分别投影出 query、key、value,外加三组控制信号:alpha(衰减输入)、beta(写入幅度)、output_gate(输出门控),然后走三个深度可分离的因果卷积:
self.q_conv1d = nn.Conv1d( self.key_dim, self.key_dim, groups=self.key_dim, **conv_kwargs ) self.k_conv1d = nn.Conv1d( self.key_dim, self.key_dim, groups=self.key_dim, **conv_kwargs ) self.v_conv1d = nn.Conv1d( self.value_dim, self.value_dim, groups=self.value_dim, **conv_kwargs )conv_kwargs里的kernel_size = config.linear_conv_kernel_dim(即 4)、padding = 3、无 bias。每个通道独立卷积、因果填充,卷积输出再经过 SiLU 激活(hidden_act = "silu")。这就是 KDA 处理局部依赖的方式:kernel size 4 的因果卷积让每个位置能看到自己前面 3 个 token 的原始信号,再叠加循环状态来携带更远的记忆。卷积之后,真正的"记账"发生在_torch_kda的循环里:
state = state * gate[:, token_idx].exp().unsqueeze(-1) prediction = torch.einsum("bhk,bhkv->bhv", key_i, state) delta = (value_i - prediction) * beta[:, token_idx].unsqueeze(-1) state = state + torch.einsum("bhk,bhv->bhkv", key_i, delta) outputs.append(torch.einsum("bhk,bhkv->bhv", query_i, state))这是标准的 delta 规则更新:先用当前 key 从状态里"预测"出 value,算出残差delta,再用beta(sigmoid 后取值 0~1)控制把多少残差写进状态,最后用 query 从更新后的状态读出输出。状态是一个形状为(batch, v_heads, k_dim, v_dim)的固定大小张量,不随序列长度增长——历史信息被"压缩"进这张固定账本,而不是像 KV cache 那样逐 token 累计。
衰减门控则是记账的"折旧率":
gate = -self.a_log_bias.float().exp().view(1, 1, self.num_k_heads, 1) * \ functional.softplus( alpha.float() + self.dt_bias.float().view(1, 1, self.num_k_heads, self.head_k_dim) )a_log_bias是每头一个的可学习参数(exp 后恒为正,保证衰减),dt_bias是逐维的偏置,alpha由网络按 token 动态预测。配合kda_allow_negative_eigenvalues = false(beta不乘 2,保持 0~1 的收缩写入),整个状态更新天然是有界、可衰减的——旧信息随时间指数折旧,这正是"流水账"该有的样子。query 和 key 还会先做 L2 归一化(use_qk_l2norm_in_kernel=True),让点积有界、训练稳定。
三、Gated Attention:每 4 层一次的"精确审计"
流水账记久了会失真,所以每个块的末尾——第 4 层——安排了一次精确的全局审计。AliceAIAttention是标准 softmax 注意力,但带了三个值得注意的工程细节:
- GQA 压缩:16 个 query 头只有 2 个 KV 头(
num_key_value_heads: 2),8:1 的分组共享把 KV cache 压到 1/8; - QK 归一化:query 和 key 在进入注意力前都过 RMSNorm(zero-centered 变体),替代了部分场景下的温度缩放;
- 部分旋转:
partial_rotary_factor = 0.25,只有 1/4 的维度注入 RoPE(rope_theta = 1e6),其余维度保持绝对位置不敏感,兼顾位置感知与通道自由度。
全注意力层还带输出门控:output = output * torch.sigmoid(output_gate),与 KDA 的o_norm+ sigmoid 门控异曲同工。
更值得注意的是块级残差机制。在AliceAIModel.forward里,每 4 层(layer_idx % block_attn_res_block_size == 0)会结算一次"completed block":
if layer_idx > 0 and layer_idx % self.config.block_attn_res_block_size == 0: completed_blocks.append(partial) partial = None每一层的输出并不是简单的残差加法,而是通过_depth_softmax_mix对当前块内所有已完成层的输出做深度维 softmax 加权混合——各层输出先 RMSNorm,再经一个可学习的标量投影打分,softmax 得到权重后加权求和。层与层之间因此存在可学习的"记账权重",3 次 KDA 流水 + 1 次全注意力审计的贡献度不是写死的,而是训练出来的。
四、为什么是 3:1?混合账本的设计推演
配比 3:1 不是拍脑袋,仓库里至少有三本账能对上:
复杂度账本。全注意力是 O(n²),KDA 是 O(n)。48 层里 36 层走线性路径,长序列下注意力部分的计算量主要集中在那 12 层 Gated Attention 上。如果把配比换成 1:1(每 2 层一次全注意力),长上下文成本几乎翻倍;如果全线性(0 次全注意力),又丢失精确检索能力。3:1 是在"成本"与"精度"之间的一个很克制的取点。
KV cache 账本。12 层全注意力即使有 GQA 8:1 压缩,KV cache 依然随上下文线性增长;而 36 层 KDA 层不存 KV,只存固定大小的 recurrent state 加number_of_conv_states = 3个卷积状态(对应 q/k/v 三路因果卷积,kernel=4,decode 阶段只需保留 kernel-1=3 个历史元素)。在 262144 的上下文目标下,如果 48 层全是全注意力,KV cache 会是一个天文数字;3:1 的配比让"随长度增长的部分"被压到最小。
局部 vs 全局的边界账本。KDA 的因果卷积 kernel size 恰好也是 4,与block_attn_res_block_size = 4对齐:每个 4 层块内,KDA 的局部感受野与块的边界天然吻合——前 3 层负责用卷积+循环状态消化局部与中程依赖,第 4 层全注意力负责跨块的长程检索。kernel=4 与 block=4 在同一个 config 里出现两次,这很难说是巧合。
五、长上下文上的实账:262K 与 128k 评测
这套混合账本的实际收益,先看配置承诺:max_position_embeddings: 262144,即支持 26 万 token 上下文。再看 README.md 里长上下文基准的真实数字(vLLM 推理、t=0 采样下的 5-shot 结果):
- FinQA 128k(金融财报长文分析):74.1,与 DeepSeek-V4-Flash-Base 并列第一,显著高于 Qwen3.5-35B-A3B 的 73.5 和 GLM-4.5-Air 的 35.5;
- LongMemEval 128k(长对话历史检索):64.6,超过 Qwen3.5-35B-A3B 的 55.6 与 GLM-4.5-Air 的 50.6。
在数学与代码侧,MATH-500 达 91.1、LiveCodeBench v5-6 CoT 1-shot pass@1 达 50.5,说明把 75% 的层换成线性注意力并没有以推理能力为代价——前提是那 25% 的全注意力层与 MoE 层把"精确审计"的活干到位。
推理侧还有一个实打实的收益:KDA 在 GPU 上通过flash-linear-attention的chunk_kda/fused_recurrent_kda内核执行(modeling_alice_ai.py 中_kda方法的 CUDA 分支),decode 阶段走fused_recurrent_kda,prefill 走chunk_kda,两者共享同一套initial_state/final_state接口,prefill 算出的 final state 可以直接作为 decode 的初始状态,无需重算历史。因果卷积状态同样通过cache.update_conv_state增量维护,decode 每步只处理 1 个新 token。
六、落地时的几个关键细节
- 双份 mask:
AliceAIModel.forward接受 dict 形式的attention_mask,分别给linear_attention层和full_attention层传不同的 mask(modeling_alice_ai.py 的attention_mask["linear_attention"]/attention_mask["full_attention"]分支)。deploy 时别把两份 mask 混用。 - 依赖约束:Transformers 侧跑 KDA 层需要
flash-linear-attention>=0.5.0,否则 CUDA 路径会直接抛ImportError(源码里写明了这一点);参考版本是transformers==5.16.1。 - MTP 已固化:训练时的 MTP 头(
mtp_num_hidden_layers: 1)已融合进权重,推理时通过_keys_to_ignore_on_load_unexpected = [r"^mtp\."]忽略,配合 vLLM 的--speculative-config '{"method":"mtp","num_speculative_tokens":1}'做投机解码,一次前向生成 2 个 token——这是社区文章中验证过的 1.2–1.8× 加速路径。
把三本账合起来看,AliceAI-Foundation-80B-A3B-Base 的注意力设计逻辑非常自洽:用 36 层线性注意力承担"广覆盖、低成本"的流水记账,用 12 层全注意力承担"高精度、有限次数"的全局审计,再用 kernel=4 的因果卷积与 block=4 的残差混合把两层账本的边界对齐。在 262K 上下文的约束下,这套混合账本既守住了长程检索的精度,又把随序列增长的成本锁死在 25% 的层上——它给出的不是一个"更便宜的近似注意力",而是一个把注意力预算明确分成两本账、各记各的架构决策。
【免费下载链接】AliceAI-Foundation-80B-A3B-Base项目地址: https://ai.gitcode.com/hf_mirrors/yandex/AliceAI-Foundation-80B-A3B-Base
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考