Performer 的 TensorFlow 版 FAVOR+ 快速注意力:favor_attention 函数与 Keras 层使用指南
发布时间:2026/9/21 18:15:40 作者:尧图编辑部 阅读量:1,286

人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载本文是 google-research 仓库中 Performer 快速注意力FAVOR模块 TensorFlow 实现的使用指南对应文档位于 performer/fast_attention/tensorflow/README.md。文章围绕核心函数favor_attention与 Keras 层Attention/SelfAttention展开结合 fast_attention.py 源码与 fast_attention_test.py 测试用例帮助你理解如何在 Transformer 中通过随机特征映射把标准注意力替换为线性复杂度的快速注意力并掌握 softmax 与广义 ReLU 两种内核的选择与配置方法。一、背景FAVOR 与 PerformerPerformer 的论文Rethinking Attention with PerformersICLR 2021Oral提出了 FAVORFast Attention Via positive Orthogonal Random features机制。其核心思想是利用结构化的随机特征映射random feature maps与注意力矩阵的低秩分解将传统 softmax 注意力的时间复杂度从序列长度的平方级O(L²)降到线性级同时保证对原始 softmax 注意力的无偏且紧致的近似。仓库中的 FAVOR 模块同时提供 JAX 与 TensorFlow 两种实现见 performer/fast_attention/README.mdJAX 变体位于 performer/fast_attention/jax/通过make_fast_softmax_attention与make_fast_generalized_attention两个工厂函数生成与flax.deprecated.nn.attention.dot_product_attention同 API 的attention_fn用于 Flax 构建的 TransformerTensorFlow 变体位于 performer/fast_attention/tensorflow/核心 API 为函数favor_attention与 Keras 层Attention、SelfAttention本文重点介绍。从源码结构看TensorFlow 变体由三个文件组成核心实现 fast_attention.py、基于tf.einsum的 Keras 稠密层工具 util.py 以及测试文件 fast_attention_test.py。二、核心函数favor_attention文档明确指出TensorFlow 变体的主注意力函数是favor_attention它实现归一化后的 FAVOR 注意力。其完整签名见 fast_attention.pydef favor_attention(query, key, value, kernel_transformation, causal, projection_matrixNone):参数含义参数类型说明query/key/valueTensorQ、K、V 张量形状为[B, L, H, D]batch、序列长度、头数、每头维度kernel_transformationcallable将 Q/K 映射为有限维内核特征的变换函数causalbool是否为因果单向注意力即是否施加 maskingprojection_matrixTensor 或 None随机投影矩阵形状[M, D]为None时表示使用恒等映射不投影2.1 内部执行流程favor_attention的执行分为三步特征映射分别对 Q 与 K 调用kernel_transformation得到query_prime与key_prime形状[B, L, H, M]M 为随机特征数转置对齐将query_prime、key_prime、value统一转置为[L, B, H, ...]布局便于按序列维度累积分子分母计算根据causal标志分流——非因果调用noncausal_numerator与noncausal_denominatorfast_attention.py通过tf.einsum(lbhm,lbhd-bhmd, ks, vs)先合并 K、V 再与 Q 点乘因果调用causal_numerator与causal_denominator以前缀和prefix sum方式逐位置累积 K·V。最后将分子除以归一化项av_attention / attention_normalizer得到归一化的快速注意力输出。2.2 快速测试验证fast_attention_test.py 中的test_relu_noncausal_attention_block_output等用例展示了最小调用方式kernel_transformation fast_attention.relu_kernel_transformation attention_block_output fast_attention.favor_attention( query, key, value, kernel_transformation, False) self.assertListEqual(attention_block_output.get_shape().as_list(), [batch_size, length, num_heads, dim])测试同时覆盖了因果与非因果两种模式并断言输出形状保持[B, L, H, D]不变说明该函数可直接替换原始注意力计算而无需改动张量布局。三、内核变换的选择softmax 与广义 ReLU文档给出的关键配置是通过favor_attention的kernel_transformation参数选择注意力内核。3.1 softmax 注意力默认语义使用 softmax 注意力时设置kernel_transformationsoftmax_kernel_transformationsoftmax_kernel_transformation实现了 FAVOR 对 softmax 内核的随机特征构造见 fast_attention.py。其数学要点利用等式e^{q·k^T/√d} e^{q_norm · k_norm^T}先以data_normalizer 1 / d^(1/4)归一化数据通过ratio 1/√M缩放投影后的点积引入diag_data ‖data‖² / 2项完成平方展开e^{-‖q‖²/2} · e^{‖q‖²/2 q·k} · e^{-‖k‖²/2}的展开形式在exp内部减去reduce_max以保证数值稳定性加上numerical_stabilizer默认1e-6确保特征值非负且远离零。注意 query 与 key 的 max-reduction 维度不同query 仅沿最后一个维度特征维做 maxkey 则沿最后维度与注意力维度attention_dims_t一起做 max源码中last_dims_t attention_dims_t这是为了保证因果情形下归一化正确。3.2 广义 ReLU 注意力使用广义 ReLU 注意力时设置kernel_transformationrelu_kernel_transformationrelu_kernel_transformation对应论文中广义注意力的 ReLU 内核见 fast_attention.py。实现上若projection_matrix为None直接返回tf.nn.relu(data) numerical_stabilizer默认0.001即退化为主机上的纯 ReLU 特征否则先以ratio 1/√M缩放投影结果再施加tf.nn.relu并加上 stabilizer保证特征非负。3.3 关于投影矩阵当使用 softmax 内核时通常需要传入由create_projection_matrix生成的随机正交投影矩阵fast_attention.py。该函数构造形状为[m, d]的随机投影矩阵每个投影方向均匀随机选取长度可以取确定值√dscaling1或服从χ(d)分布scaling0此时投影的边际分布为 d 维高斯向量。struct_modeTrue时使用 Givens 旋转乘积构造随机正交矩阵可绕过 Gram-Schmidt 正交化对大规模场景更高效。测试test_softmax_noncausal_attention_block_outputfast_attention_test.py给出了完整的 softmax 使用范式序列长度 10000、随机特征数 1000并与精确 softmax 注意力对比验证最大误差小于 0.5projection_matrix fast_attention.create_projection_matrix( num_random_features, dim) attention_block_output fast_attention.favor_attention( query, key, value, kernel_transformation, False, projection_matrix)四、以 Keras Layer 的方式使用Attention 与 SelfAttention文档指出若要作为tf.keras.layers.Layer使用应使用 FAVOR 的Attention类在设置好 FAVOR 配置后其 API 与tf.keras.layers.Attention()类似。Attention类的构造函数见 fast_attention.pyAttention(hidden_size, num_heads, attention_dropout, kernel_transformationrelu_kernel_transformation, numerical_stabilizer0.001, causalFalse, projection_matrix_typeNone, nb_random_features0)参数默认值说明hidden_size—必填隐藏层输出维度必须能被num_heads整除否则抛出ValueErrornum_heads—必填注意力头数attention_dropout—必填训练时注意力内部的 dropout 率kernel_transformationrelu_kernel_transformation产生内核特征的变换numerical_stabilizer0.001使内核值远离零的稳定项causalFalse是否为因果注意力projection_matrix_typeNoneNone表示使用恒等映射否则应用随机投影矩阵nb_random_features0随机特征数量仅在投影矩阵非 None 时有效在build阶段层内部通过util.DenseEinsum一个以tf.einsum为底层计算、可处理任意维度的 Keras 稠密层见 util.py构造 Q/K/V 三组投影与输出变换并使用 Glorot 初始化。hidden_size // num_heads即为每头维度size_per_head。call方法的输入约定与tf.keras.layers.Attention()风格对齐query_input[batch_size, length_query, hidden_size]source_input[batch_size, length_source, hidden_size]bias[batch_size, 1, length_query, length_source]的注意力偏置training是否处于训练模式cache预测解码阶段可选的 KV 缓存格式为{k: ..., v: ...}用于增量自回归解码decode_loop_step解码循环步数TPU 上自回归推理使用。调用时层先对 Q/K/V 做线性投影并自动切分多头再根据projection_matrix_type决定是否构造随机投影矩阵非 None 时以输入的统计量派生种子调用create_projection_matrix(self.nb_random_features, dim, seedseed)随后交给favor_attention完成快速注意力最后经output_dense_layer输出。SelfAttention是Attention的子类fast_attention.py其call直接把query_input同时作为source_input传入即自注意力形式。测试test_fast_attentionfast_attention_test.py展示了典型用法layer fast_attention.SelfAttention(hidden_size64, num_heads4, dropout0.5) y layer(x, bias, trainingTrue, cachecache) # y.shape (1, length, 64)五、因果单向变体的自定义梯度与内存优化文档特别强调与 JAX 变体一样TensorFlow 的因果单向变体通过tf.custom_gradient提供自定义梯度从而获得显著的内存缩减。在 fast_attention.py 中causal_numeratorL226-L273与causal_denominatorL276-L319均用tf.custom_gradient装饰。前向计算以 Python 循环逐位置维护累积和sumsfor index in range(qs.shape[0]): sums sums tf.einsum(ijk,ijl-ijkl, ks[index], vs[index]) result.append(tf.einsum(ijkl,ijk-ijl, sums, qs[index])[None, Ellipsis])自定义grad函数则从最后一个位置反向遍历range(qs.shape[0] - 1, -1, -1)通过从总和中逐个减去ks[index]·vs[index]的方式递推计算 Q/K/V 三者的梯度。这种前缀和 反向剥离的写法避免了显式构造并存储完整的L×L注意力矩阵正是因果 FAVOR 内存优势的来源。测试test_custom_causal_gradientsfast_attention_test.py在 L64、B128、H4、D64、M128 的规模下用tf.GradientTape验证了自定义梯度可正常求出且形状正确with tf.GradientTape() as tape: tape.watch([qs, ks, vs]) num fast_attention.causal_numerator(qs, ks, vs) den fast_attention.causal_denominator(qs, ks) loss tf.reduce_sum(num * num_coefs) tf.reduce_sum(den * den_coefs) * 0 grads1 tape.gradient(loss, [qs, ks, vs])六、与 JAX 变体的对应关系若读者同时接触两种实现可按下表对照JAX 侧见 performer/fast_attention/jax/fast_attention.py能力TensorFlow 变体JAX 变体softmax 内核softmax_kernel_transformation配合favor_attentionmake_fast_softmax_attention广义内核relu_kernel_transformationReLU 等非线性映射make_fast_generalized_attention支持 cos、sin、tanh、gelu 等可配置非线性因果变体梯度tf.custom_gradientJAX 自定义 VJP均用于显著降低内存集成方式函数调用或 Keras LayerAttention/SelfAttention生成与dot_product_attention同 API 的attention_fn可配合 gin 配置如FlaxModel.attention_fn make_fast_softmax_attention()七、小结与使用建议若只想在自定义张量计算中替换注意力直接调用favor_attention(query, key, value, kernel_transformation, causal, projection_matrix)若需要可训练的端到端模块使用 Keras 层Attention跨注意力或SelfAttention自注意力并注意hidden_size需被num_heads整除近似 softmax 时务必配合create_projection_matrix生成的随机正交投影矩阵并适当增大nb_random_features测试中 softmax 场景使用了 3501000 个随机特征以逼近原始注意力精度因果解码场景优先启用causalTrue借助自定义梯度在长序列上获得内存收益配合cache参数支持增量自回归推理。赞分享人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载相关推荐终极指南如何利用annotated_deep_learning_paper_implementations掌握高效注意力机制变体终极指南如何利用annotated_deep_learning_paper_implementations掌握高效注意力机制变体 annotated_deep人工智能深度学习大模型NLP计算机视觉强化学习LoRADeepTutor 入门指南从零搭一个会消化资料、记得住你的私人导师DeepTutor 入门指南从零搭一个会消化资料、记得住你的私人导师 凌晨两点你卡在一道题上搜了半小时每条结果都只解决了一半问题。你缺的不是资料是人工智能AI 应用AI Agent多智能体RAG教育后端前端【亲测免费】 探索Keras中的注意力机制层Keras Attention Layer探索Keras中的注意力机制层Keras Attention Layer 在这个充满创新的世界里深度学习的工具和库持续发展以满足日益增长的需求。其中 K创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考