每 4 层插 3 次线性注意力:KDA 与 Gated Attention 的混合账本
每 4 层插 3 次线性注意力KDA 与 Gated Attention 的混合账本【免费下载链接】AliceAI-Foundation-80B-A3B-Base项目地址: https://ai.gitcode.com/hf_mirrors/yandex/AliceAI-Foundation-80B-A3B-BaseYandex 开源的 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 或普通的线性注意力近似而是名为 KDAKernel Delta Attention 一类带 delta 状态更新的线性注意力变体的实现配置里还专门给出了它的参数32 个 query 头、32 个 KV 头、head dim 128、因果卷积核大小 4linear_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, groupsself.key_dim, **conv_kwargs ) self.k_conv1d nn.Conv1d( self.key_dim, self.key_dim, groupsself.key_dim, **conv_kwargs ) self.v_conv1d nn.Conv1d( self.value_dim, self.value_dim, groupsself.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再用betasigmoid 后取值 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 falsebeta不乘 2保持 0~1 的收缩写入整个状态更新天然是有界、可衰减的——旧信息随时间指数折旧这正是流水账该有的样子。query 和 key 还会先做 L2 归一化use_qk_l2norm_in_kernelTrue让点积有界、训练稳定。三、Gated Attention每 4 层一次的精确审计流水账记久了会失真所以每个块的末尾——第 4 层——安排了一次精确的全局审计。AliceAIAttention是标准 softmax 注意力但带了三个值得注意的工程细节GQA 压缩16 个 query 头只有 2 个 KV 头num_key_value_heads: 28:1 的分组共享把 KV cache 压到 1/8QK 归一化query 和 key 在进入注意力前都过 RMSNormzero-centered 变体替代了部分场景下的温度缩放部分旋转partial_rotary_factor 0.25只有 1/4 的维度注入 RoPErope_theta 1e6其余维度保持绝对位置不敏感兼顾位置感知与通道自由度。全注意力层还带输出门控output output * torch.sigmoid(output_gate)与 KDA 的o_norm sigmoid 门控异曲同工。更值得注意的是块级残差机制。在AliceAIModel.forward里每 4 层layer_idx % block_attn_res_block_size 0会结算一次completed blockif 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 三路因果卷积kernel4decode 阶段只需保留 kernel-13 个历史元素。在 262144 的上下文目标下如果 48 层全是全注意力KV cache 会是一个天文数字3:1 的配比让随长度增长的部分被压到最小。局部 vs 全局的边界账本。KDA 的因果卷积 kernel size 恰好也是 4与block_attn_res_block_size 4对齐每个 4 层块内KDA 的局部感受野与块的边界天然吻合——前 3 层负责用卷积循环状态消化局部与中程依赖第 4 层全注意力负责跨块的长程检索。kernel4 与 block4 在同一个 config 里出现两次这很难说是巧合。五、长上下文上的实账262K 与 128k 评测这套混合账本的实际收益先看配置承诺max_position_embeddings: 262144即支持 26 万 token 上下文。再看 README.md 里长上下文基准的真实数字vLLM 推理、t0 采样下的 5-shot 结果FinQA 128k金融财报长文分析74.1与 DeepSeek-V4-Flash-Base 并列第一显著高于 Qwen3.5-35B-A3B 的 73.5 和 GLM-4.5-Air 的 35.5LongMemEval 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 pass1 达 50.5说明把 75% 的层换成线性注意力并没有以推理能力为代价——前提是那 25% 的全注意力层与 MoE 层把精确审计的活干到位。推理侧还有一个实打实的收益KDA 在 GPU 上通过flash-linear-attention的chunk_kda/fused_recurrent_kda内核执行modeling_alice_ai.py 中_kda方法的 CUDA 分支decode 阶段走fused_recurrent_kdaprefill 走chunk_kda两者共享同一套initial_state/final_state接口prefill 算出的 final state 可以直接作为 decode 的初始状态无需重算历史。因果卷积状态同样通过cache.update_conv_state增量维护decode 每步只处理 1 个新 token。六、落地时的几个关键细节双份 maskAliceAIModel.forward接受 dict 形式的attention_mask分别给linear_attention层和full_attention层传不同的 maskmodeling_alice_ai.py 的attention_mask[linear_attention]/attention_mask[full_attention]分支。deploy 时别把两份 mask 混用。依赖约束Transformers 侧跑 KDA 层需要flash-linear-attention0.5.0否则 CUDA 路径会直接抛ImportError源码里写明了这一点参考版本是transformers5.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 层全注意力承担高精度、有限次数的全局审计再用 kernel4 的因果卷积与 block4 的残差混合把两层账本的边界对齐。在 262K 上下文的约束下这套混合账本既守住了长程检索的精度又把随序列增长的成本锁死在 25% 的层上——它给出的不是一个更便宜的近似注意力而是一个把注意力预算明确分成两本账、各记各的架构决策。【免费下载链接】AliceAI-Foundation-80B-A3B-Base项目地址: https://ai.gitcode.com/hf_mirrors/yandex/AliceAI-Foundation-80B-A3B-Base创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考