ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

flash-linear-attention 中的 Simple GLA:Head-wise 门控线性注意力的数学原理与内核实现

flash-linear-attention 中的 Simple GLA:Head-wise 门控线性注意力的数学原理与内核实现 flash-linear-attention 中的 Simple GLAHead-wise 门控线性注意力的数学原理与内核实现【免费下载链接】flash-linear-attention Efficient implementations for emerging model architectures项目地址: https://gitcode.com/GitHub_Trending/fl/flash-linear-attentionSimple GLA 是 flash-linear-attentionFLA中一类门控粒度更粗的线性注意力算子与标准 GLA 的逐元素门控不同它对每个注意力头共享一个标量遗忘门从而可以复用 RetNet 式的 chunk 训练内核用纯 matmul 完成训练而不引入数值不稳定。本文以仓库内 Simple GLA 的算子定义与测试代码为依据完整讲清它的递推公式、与 GLA/Mamba2/YOCO 门控机制的关系、四种内核入口chunk / fused_recurrent / fused_chunk / parallel的参数约定以及配套SimpleGatedLinearAttention层的构造参数与正确性验证方式。读完后可直接在项目中调用这些内核或该层并理解其数值行为与适用边界。一、Simple GLA 是什么Head-wise 门控的线性注意力仓库中 Simple GLA 的算子背景定义在 Simple GLA 参考文档其核心结论可以概括为三点门控粒度是 head-wise 而非 elementwise。标准 GLA 的门控对 state 矩阵的每个元素独立衰减而 Simple GLA 的每个头只使用一个标量门 $g_t$递推公式为$$S_{t1} g_{t1} \odot S_{t} K_{t1} V_{t1}^{\top}$$其中 $g$ 是标量。这一简化正是它相对 GLA 的全部差异来源。门控机制与 Gated RFA、Mamba2、YOCOGated RetNet同源。由于门是 head-wise 的训练时可以适配 RetNet 的 chunk 内核chunk 间的状态传播退化为按标量衰减的矩阵乘法using matmul w/o numerical instability——即不需要 GLA 中逐元素指数累和那类容易溢出的运算。定位是更快但表达力更弱的基线。参考文档明确说I will use it as a baseline for the GLA它比 GLA 快但表达力更小适合作为对照基线。从朴素参考实现 naive.py 可以看到这个递推的直接编码在naive_recurrent_simple_gla中状态更新就是一个标量门乘子加一个外积累加S S * gate.unsqueeze(-1).unsqueeze(-1) kv见 naive.py与上面的数学定义一一对应。naive_chunk_simple_gla则展示了 chunk 视角chunk 内用带衰减掩码的q k^TL_mask (decay.unsqueeze(-1) - decay.unsqueeze(-2)).tril().exp()做 intra-chunk 注意力chunk 间用S * decay_last做 state 传播——这正是RetNet 内核可被适配的具体含义。二、仓库中的内核入口与参数约定Simple GLA 的公开 API 定义在 fla/ops/simple_gla/init.py共导出四个内核入口文件定位chunk_simple_glachunk.py训练主路径chunk 并行 可微支持变长fused_recurrent_simple_glafused_recurrent.py融合逐 token 递推适合短序列/推理也作为数值参考fused_chunk_simple_glafused_chunk.py融合 chunk 内核同样用于测试对照parallel_simple_glaparallel.py并行形式显式构造衰减注意力矩阵 A用于校验 A 矩阵chunk_simple_gla是训练时最常用的入口其参数约定见 chunk.py 的文档字符串q/k/v形状分别为[B, T, H, K]、[B, T, H, K]、[B, T, H, V]统一采用[B, T, H, ...]布局传入已废弃的head_first会抛DeprecationWarninggforget gates形状[B, T, H]。注意它比 GLA 的 elementwise 门少一个维度——这正是head-wise在 API 上的直接体现。g与g_gamma二者只能提供一个后者是形状[H]的与数据无关的 head-wise log decayscale注意力分数缩放缺省为1 / sqrt(K)initial_state/output_final_state初始/最终状态形状[N, H, K, V]等长输入下N Bstate_v_first把 recurrent state 存成[V, K]布局默认[K, V]便于与特定后端或下游模块对齐旧的transpose_state_layout参数已废弃并映射到它见 fused_recurrent.pycu_seqlens形状[N1]的累积序列长度用于变长训练与 FlashAttention API 一致。此时批次维必须为 1变长输入需要预先拼接扁平化且initial_state的第一维必须等于len(cu_seqlens) - 1否则直接抛ValueError见 chunk.pychunk_size必须为 2 的幂缺省 64。变长用法示例摘自 chunk.py 文档字符串可复制运行import torch import torch.nn.functional as F from einops import rearrange from fla.ops.simple_gla import chunk_simple_gla B, T, H, K, V 4, 2048, 4, 512, 512 q torch.randn(B, T, H, K, devicecuda) k torch.randn(B, T, H, K, devicecuda) v torch.randn(B, T, H, V, devicecuda) g F.logsigmoid(torch.randn(B, T, H, devicecuda)) o, ht chunk_simple_gla(q, k, v, g, initial_stateNone, output_final_stateTrue) # 变长输入B 必须为 1并给出 cu_seqlens q, k, v, g map(lambda x: rearrange(x, b t ... - 1 (b t) ...), (q, k, v, g)) cu_seqlens q.new_tensor([0, 2048, 4096, 6144, 8192], dtypetorch.long) o_var, ht_var chunk_simple_gla( q, k, v, g, initial_stateNone, output_final_stateTrue, cu_seqlenscu_seqlens, )三、chunk 内核内部为什么 head-wise 门不会数值不稳定chunk_simple_gla的自动微分函数ChunkSimpleGLAFunction把前向拆成三步公共内核调用chunk.pychunk_local_cumsum(g, scaleRCP_LN2, ...)对 head-wise 门做chunk 内局部累和。由于门是每头一个标量这个累和是逐头的一维标量序列运算代价远低于 GLA 的[B, T, H, K]逐元素累和chunk_fwd_hfla/ops/common/chunk_h.py计算每个 chunk 内的局部状态轨迹h。注意 Simple GLA 调用它时gkNone, gvNone只传g——即只启用标量衰减通道chunk_fwd_ofla/ops/common/chunk_o.py由h和 intra-chunk 的q k^T得到输出o。代码注释给出关键工程细节g先乘以RCP_LN2预缩放使下游 Triton 内核可以直接用硬件exp2代替exp见 chunk.py。反向传播则复用chunk_bwd_dh/chunk_bwd_dqkwg/chunk_bwd_dv三个公共反向内核并同样以states_in_fp32True计算 state 轨迹以提高梯度精度chunk.py对g的梯度通过一次reverseTrue的chunk_local_cumsum从衰减后梯度还原chunk.py。这一结构印证了参考文档的判断head-wise 门让 chunk 间传播只是标量乘S所有运算都可以落在 TensorCore 友好的 matmul 上这正是adapt the RetNet kernel for training的实现形态。四、SimpleGatedLinearAttention 层把算子组装成可训练层fla/layers/simple_gla.py 提供了完整的 Transformer 层封装SimpleGatedLinearAttention。它的 docstring 说明calls the simplified GLA kernel in which the gating is head-wise instead of elementwise即该层就是第三节所述算子的直接上层消费者。关键构造参数默认值以__init__签名为准simple_gla.py参数默认值说明modechunk使用的内核当前支持chunk与fused_recurrenthidden_size1024输入隐藏维expand_k/expand_v1.0key / value 维扩展比key_dim hidden_size * expand_k必须为整数num_heads4注意力头数key_dim与value_dim必须都能被它整除num_kv_headsnum_headsKV 头数用于 GQA/MQA要求num_heads % num_kv_heads 0feature_mapNone作用于 q/k 的映射函数经ACT2FN注册表查找use_short_convTrue注意签名默认是否为 q/k/v 加短卷积默认 silu 激活conv_size/conv_bias4 /False短卷积核大小与 bias 开关gate_fnswish输出门激活函数gate_logit_normalizer16门 logit 的归一化因子作用于logsigmoid之后fuse_normTrue是否用FusedRMSNormGated把 norm 与输出门融合layer_idxNone层索引供 cache 管理使用前向中的门控计算值得单独看层从 hidden states 投影出 per-head 标量门gk然后执行gk F.logsigmoid(gk) / self.gate_logit_normalizer # fla/layers/simple_gla.py#L236logsigmoid把门约束到负值衰减系数小于 1再除以归一化因子控制遗忘速度——测试矩阵中gate_logit_normalizer覆盖 0.1 / 1 / 10 三档说明它是一个敏感的超参。该层还有一个实用的运行期路由逻辑simple_gla.py训练梯度开启时强制走chunk内核推理且序列长度不超过 64 时改走fused_recurrent——注释解释为launching the triton kernel for just one token will actually be slower即单 token 场景下融合递推内核更快。缓存方面层通过get_layer_cache/update_layer_cache管理 recurrent state 与三个短卷积的 conv state并支持cu_seqlens传入以处理变长序列state_size()给出了单序列状态总大小key_dim * head_v_dim加上各卷积状态可用于规划 KV cache 内存。五、正确性验证测试如何覆盖 Simple GLASimple GLA 的数值正确性由 tests/ops/test_simple_gla.py 系统覆盖其组织方式与仓库的 correctness coverage 技能文档 中给出的覆盖轴序列布局 dense/varlen、前向/反向、状态传递、头维等一致test_fused_recurrent/test_fused_recurrent_varlen以纯 PyTorch 的naive_recurrent_simple_gla为参考逐位比对o、最终状态ht以及dq/dk/dv/dg/dh0全套梯度参数化覆盖scale ∈ {0.1, 1}、gate_logit_normalizer ∈ {0.1, 1, 10}、float32/float16 与不等长cu_seqlenstest_chunk/test_chunk_varlen/test_chunk_with_chunk_size以fused_recurrent_simple_gla为参考验证 chunk 内核chunk_size覆盖 16/32/64 三种取值变长用例使用cu_seqlens [0, 15, 100, 300, 1200, 2000]这类非对齐长度test_fused_chunk/test_fused_chunk_varlen同样的对照方式验证融合 chunk 内核test_parallel验证parallel_simple_gla输出的注意力矩阵A与naive_parallel_simple_gla构造的衰减矩阵一致test_chunk_state_v_first/test_fused_recurrent_state_v_first验证state_v_first布局与默认布局的数值等价初始状态与最终状态梯度需做transpose(-1, -2)对齐同时确认废弃参数transpose_state_layout的告警与互斥行为。一个特别有说服力的用例是test_simple_gla_to_mamba2test_simple_gla.py它把 Mamba2 的 SSD 参数直接映射到 Simple GLA 输入——q C、k B、v x、g A * dt、scale 1.0——然后断言chunk_simple_gla的输出与 Mamba2 参考实现ssd_minimal_discrete乃至官方mamba_chunk_scan_combined融合内核逐位一致float16 容差 1e-2。该测试从实现层面证实了参考文档的核心论断Simple GLA 的门控机制与 Mamba2 是同一套 head-wise 门控线性注意力只是记法不同。这也解释了为什么 head-wise 标量衰减足以覆盖 Mamba2 这类现代架构的 state 动态。六、选型建议与适用边界结合参考文档的结论与上述源码证据可以在以下场景做判断需要更强的记忆调制表达力门控随通道位置变化时应选择 GLAfla/ops/gla/其门是 elementwise 的代价是内核需要逐元素指数运算、实现更复杂追求训练速度或作为基线时Simple GLA 的 chunk 内核路径更短且fused_recurrent变体可承担推理路径与 Mamba2 权重或语义对齐的场景例如转换、对比实验chunk_simple_gla与 Mamba2 SSD 的输入映射关系已由test_simple_gla_to_mamba2固化可直接参考变长训练时注意批次维必须拼成 1、cu_seqlens必须与initial_state条数一致这两条硬约束它们由内核入口的ValueError显式校验。需要说明的边界参考文档所述faster than GLA是基于算子结构matmul 为主、无逐元素指数路径给出的定性结论仓库内 benchmarks/ops/registry.py 注册了chunk_simple_gla基准项具体吞吐差距应结合目标硬件用benchmarks/ops/run.py实测本文不给出未经实测的性能数字。参考路径索引算子背景.agents/skills/fla-correctness-coverage/references/simple-gla.md内核实现fla/ops/simple_gla/chunk.py、fla/ops/simple_gla/fused_recurrent.py、fla/ops/simple_gla/naive.py、fla/ops/simple_gla/parallel.py层封装fla/layers/simple_gla.py测试tests/ops/test_simple_gla.py公共 chunk 内核fla/ops/common/chunk_h.py、fla/ops/common/chunk_o.py【免费下载链接】flash-linear-attention Efficient implementations for emerging model architectures项目地址: https://gitcode.com/GitHub_Trending/fl/flash-linear-attention创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表