
FNet 实战指南用傅里叶变换替换自注意力在 JAX/Flax 中预训练与微调高效编码器【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-researchFNet 是一种以「无参数傅里叶变换」完全取代 self-attention 子层的高效 Transformer 类编码器架构其官方实现与预训练权重就位于本仓库的f_net目录中。本文以 f_net/README.md 为主体结合 f_net/models.py、f_net/layers.py、f_net/configs/base.py 等源码完整讲解 FNet 的架构原理、环境安装、C4 预训练与 GLUE 微调流程、全部配置参数含义以及如何在代码层面验证傅里叶混合层的正确性。读完本文你将能够独立复现 FNet 论文实验并基于预训练检查点在自己的分类任务上完成微调。一、FNet 核心思想用傅里叶变换做 Token 混合FNet论文 arXiv:2105.03824由 James Lee-Thorp、Joshua Ainslie、Ilya Eckstein 与 Santiago Ontanon 提出其关键设计是将编码器中全部的 self-attention 子层替换为标准、无参数unparameterized的傅里叶变换。Transformer 的自注意力通过可学习的 Q/K/V 投影来混合序列各位置的 token 信息而 FNet 直接对序列维度和隐藏维度施加离散傅里叶变换DFT并仅保留结果的实部从而在几乎不损失建模能力的前提下大幅减少参数量与计算开销。本仓库中的f_net目录包含复现论文结果所需的全部模型与代码模型在 C4 数据集上预训练在 GLUE 基准上微调实现基于 JAX / Flax。1.1 编码器块的模块化结构从源码看FNet 并没有为傅里叶变换单独重写整个编码器而是复用了标准的「后归一化」post-norm编码器块。在 f_net/layers.py 中EncoderBlock按以下顺序串联子模块mixing_sublayertoken 混合子层可以是傅里叶变换、注意力或其他架构残差连接inputs mixing_outputLayerNormmixing_layer_normfeed_forward_sublayer前馈网络残差连接LayerNormoutput_layer_norm。mixing_output self.mixing_sublayer(inputs, padding_mask, deterministicdeterministic) x nn.LayerNorm(epsilonLAYER_NORM_EPSILON, namemixing_layer_norm)( inputs mixing_output) feed_forward_output self.feed_forward_sublayer(x, deterministicdeterministic) return nn.LayerNorm(epsilonLAYER_NORM_EPSILON, nameoutput_layer_norm)( x feed_forward_output)也就是说替换注意力只需要替换其中的mixing_sublayer残差、归一化与前馈结构保持不变这正是本仓库可以同时支持多种架构对照实验的架构基础。1.2 FourierTransform 子层的实现在 f_net/layers.py 中FourierTransform是一个极简模块它接收一个「傅里叶变换函数」对输入的最后两个维度通常即序列维和隐藏维施加二维傅里叶变换然后取实部返回return jax.vmap(self.fourier_transform)(inputs).real注意两个细节padding_mask与deterministic参数被显式忽略del padding_mask/del deterministic因为傅里叶变换是确定性的、不需要掩码——这是与 self-attention 子层的重要区别变换结果取real实部这是 FNet 的关键实现选择只保留实部既保留主要信号又保证输出仍为实数张量便于后续 LayerNorm 与 Dense 处理。1.3 傅里叶变换的两种计算方式EncoderModel.setup()见 f_net/models.py根据config.use_fft决定采用哪种 DFT 实现FFT 路径use_fftTrue默认直接使用jnp.fft.fftn在 GPU/CPU 上对任意序列长度都是最优实现矩阵乘法路径use_fftFalse用scipy.linalg.dft预计算序列维[MAX_SEQ_LENGTH, MAX_SEQ_LENGTH]与隐藏维[HIDDEN_DIM, HIDDEN_DIM]的 DFT 矩阵再通过 f_net/fourier.py 中的two_dim_matmuljnp.einsum(ij,jk,ni-nk, ...)批量完成变换。源码注释给出了硬件相关的选型经验在 TPU 上较短的序列用预计算 DFT 矩阵做矩阵乘法更快而较长序列用 FFT 更快前提是max_seq_length为 2 的幂。因此代码在use_fftTrue且max_seq_length 4096时会校验序列长度是否为 2 的幂否则抛出ValueErrorif (self.config.max_seq_length 4096 and not math.log2(self.config.max_seq_length).is_integer()): raise ValueError(...)1.4 对照架构与混合注意力布局为了让实验具有对照性ModelArchitecture枚举f_net/configs/base.py定义了 5 种 mixing 架构枚举值含义BERT标准 self-attention 混合层F_NET傅里叶变换混合论文主推FF_ONLY仅前馈子层不做任何 token 混合LINEAR序列/隐藏维上的可学习矩阵乘法RANDOM固定随机矩阵乘法不参与训练对应的混合层初始化逻辑在 f_net/models.py 的_init_mixing_sublayer中F_NET构建layers.FourierTransformFF_ONLY使用IdentityTransform恒等变换LINEAR使用LinearTransformRANDOM使用固定随机矩阵的RandomTransformBERT则构建nn.SelfAttention。此外仓库还支持「混合模型」用config.attention_layout与config.num_attention_layers让部分层保留注意力、其余层使用指定 mixing 架构。HybridAttentionLayoutf_net/configs/base.py提供四种布局BOTTOM前num_attention_layers层用注意力MIDDLE中间若干层用注意力MIXED注意力层穿插分布于整个模型TOP默认最后几层用注意力。判定逻辑位于 f_net/models.py 的_is_attention_layer默认num_attention_layers0即纯 FNet 配置下所有层都用傅里叶变换。二、安装与环境准备FNet 代码基于 JAX/Flax官方推荐在全新的 Python 虚拟环境如 virtualenv中安装。下载代码并安装依赖的命令如下本仓库已包含完整源码可直接在本目录内操作pip install -r f_net/requirements.txt # 如果使用 Google Cloud TPU额外安装 pip install cloud-tpu-clientf_net/requirements.txt 中的核心依赖包括absl-py、clu、flax0.3.0、jax0.2.4、ml_collections、numpy、scipy、sentencepiece、tensorflow-cpu2.5.0注意使用 CPU 版 TensorFlow以将全部 GPU 显存留给 JAX、tensorflow-datasets、tensorflow_text等。运行单元测试python3 -m unittest discover -s f_net -p *_test.py重要约束运行单元测试以及后文所有 Python 命令时当前工作目录必须是f_net文件夹的父目录即仓库根目录因为代码内部以from f_net import ...方式导入模块。三、预训练与微调命令行快速上手3.1 准备 SentencePiece 词表FNet 使用 SentencePiece 分词。首先需要下载官方预训练的词表模型c4_bpe_sentencepiece.model位于 Google Cloud Storage 的gresearch/f_net/vocab/公开目录下随后即可通过统一的命令行入口训练python3 -m f_net.main --workdir$YOUR_WORKDIR --vocab_filepath$VOCAB_FILEPATH --config$CONFIG参数含义YOUR_WORKDIR模型检查点checkpoint与训练指标的输出目录VOCAB_FILEPATH上面下载的 SentencePiece 词表模型绝对路径CONFIGf_net/configs/pretraining.py预训练或f_net/configs/classification.py微调模型结构、数据、预训练检查点路径与全部训练超参都在对应配置文件中设定。你也可以用 sentencepiece 自行训练词表但代价是无法直接复用下述官方预训练模型——除非在加载检查点后手动清空预训练模型的 embedding 层权重。3.2 入口代码如何工作f_net/main.py 使用 absl flags 定义三个必填参数config、workdir、vocab_filepath随后根据config.modeTrainingMode枚举分派训练逻辑PRETRAINING→ 调用 f_net/run_pretraining.py 中的train_and_evaluateCLASSIFICATION→ 调用 f_net/run_classifier.py 中的train_and_evaluate。值得注意的是入口在启动时会将 TensorFlow 的 GPU 设备设为不可见tf.config.experimental.set_visible_devices([], GPU)避免 TF 抢占显存导致 JAX 无显存可用。训练/微调参数统一从ml_collections配置文件加载lock_configTrue锁定配置避免运行中途被意外修改。四、配置参数全解源码级4.1 基础配置 f_net/configs/base.py所有训练共享的基础参数括号内为默认值参数默认值说明model_archF_NET模型架构见上文ModelArchitecture枚举save_checkpoints_steps1000每多少步保存一次检查点eval_frequency1000训练期间评估频率如每 1000 步train_batch_size32训练总 batch sizeeval_batch_size8评估总 batch sizelearning_rate1e-4Adam 基础学习率init_checkpoint_dir初始检查点目录/路径通常来自预训练模型do_lower_caseTrue是否对输入文本小写化uncased 模型为 Truetype_vocab_size4段类型数预训练仅需 2NSP设为 4 以支持 GLUE/SuperGLUE 最多 4 个输入段d_emb768每个 token 的 embedding 维度d_model768模型隐藏维度d_ff3072前馈层隐藏维度max_seq_length512分词后最大序列长度超长截断、不足填充num_heads12自注意力头数仅 BERT 架构使用num_layers12编码器块层数dropout_rate0.1全局 dropout 率mixing_dropout_rate0.1混合模块如自注意力子层内的 dropout 率use_fftTrue是否用 FFT 代替矩阵乘法计算 DFT仅 FNet 相关attention_layoutTOP混合模型中注意力层的分布位置num_attention_layers0混合模型中替换为注意力的层数seed0随机数种子trial0重复运行的占位参数4.2 预训练配置 f_net/configs/pretraining.py在基础配置之上覆盖mode PRETRAININGtrain_batch_size 64、eval_batch_size 64learning_rate 1e-4clipped_grad_norm None不裁剪梯度num_train_steps 1e6百万步、num_warmup_steps 1e4源码注释提示模型越大通常需要越多 warmup 步save_checkpoints_steps 2000、eval_frequency 2000、max_num_eval_steps 100掩码语言建模MLM相关max_predictions_per_seq 80、masking_rate 0.15选中的 token 总数至多为max_predictions_per_seq、mask_token_proportion 0.880% 被掩 token 替换为[MASK]、random_token_proportion 0.110% 替换为随机 token剩余 10% 保持不变init_checkpoint_dir 预训练不从已有检查点开始。4.3 微调分类配置 f_net/configs/classification.pymode CLASSIFICATIONdataset_name glue/rte可替换为glue/cola、glue/sst2、glue/mrpc、glue/qqp、glue/stsb、glue/mnli、glue/qnli、glue/rte、glue/wnli之一stsb为回归任务其余为分类任务save_checkpoints_steps 200eval_proportion 0.05训练过程中以固定间隔计算训练指标train_batch_size 64、eval_batch_size 32learning_rate 1e-5、num_train_epochs 3、warmup_proportion 0.1前 10% 训练步做线性 warmupmax_num_eval_steps 1e5小数据集实际步数可能更少init_checkpoint_dir 微调时需在此填入预训练检查点路径见下节。4.4 优化器与微调时的检查点恢复f_net/run_classifier.py 使用 Adam 优化器beta10.9、beta20.999、eps1e-6、weight_decay0.01。微调恢复预训练权重时_restore_pretrained_modelf_net/run_classifier.py有一个关键细节分类输出层名为classification会被重新初始化为全新参数因为分类任务与预训练的 MLM/NSP 任务输出维度不同同时将其对应的优化器状态清零避免旧动量干扰新层的训练。五、预训练检查点与模型规格所有预训练模型均使用max_seq_length512、type_vocab_size4最多支持 4 个输入段大多数任务会自动配置输入段数量。5.1 Base 模型Base 模型配置config.d_emb 768 config.d_model 768 config.d_ff 3072 config.num_heads 12 config.num_layers 12Base 预训练检查点单个文件约 1 GB存放于gresearch/f_net/checkpoints/base/架构检查点文件ModelArchitecture.F_NETf_net_checkpointModelArchitecture.LINEARlinear_checkpointModelArchitecture.BERTbert_checkpointModelArchitecture.FF_ONLYff_only_checkpointModelArchitecture.RANDOMrandom_checkpoint5.2 Large 模型Large 模型配置config.d_emb 1024 config.d_model 1024 config.d_ff 4096 config.num_heads 16 config.num_layers 24Large 预训练检查点单个文件 3–4 GB存放于gresearch/f_net/checkpoints/large/架构检查点文件ModelArchitecture.F_NETf_net_checkpointModelArchitecture.LINEARlinear_checkpointModelArchitecture.BERTbert_checkpoint5.3 使用方式下载对应检查点后将其路径填入微调配置的config.init_checkpoint_dir再运行f_net.main即可从预训练权重开始微调。例如微调 SST-2 时在 f_net/configs/classification.py 中设置config.dataset_name glue/sst2 config.init_checkpoint_dir /path/to/downloaded/f_net_checkpoint由于所有预训练模型共享同一套 Base/Large 维度与max_seq_length任意架构的检查点都可被 FNet 微调脚本加载这也正是论文能做 BERT / LINEAR / FF_ONLY / RANDOM 多组对照的原因。六、从源码与测试理解实现细节6.1 傅里叶变换正确性的单元测试f_net/fourier_test.py 的test_two_dim_matmul直接验证了矩阵乘法路径的正确性它构造max_seq_length3、hidden_dim8的随机输入用linalg.dft生成两个维度的 DFT 矩阵通过fourier.two_dim_matmul计算并与np.fft.fftn(inputs)的结果逐行比对在lax.Precision.DEFAULT / HIGH / HIGHEST三种精度下分别要求误差小于1e-4 / 1e-5 / 1e-5。这为「用矩阵乘法逼近 FFT」的 TPU 优化路径提供了直接证据。6.2 完整的模型族从预训练到分类f_net/models.py 定义了三个模型EncoderModelmodels.py无任务头的纯编码器。输入为input_ids、input_mask仅 BERT 注意力使用、type_ids输出隐藏状态与池化结果。池化取序列首 tokenhidden_states[:, 0]经 Dense 投影后做tanh缩放至 (-1, 1)。embedding 由词嵌入 学习式位置嵌入 段类型嵌入求和再经 LayerNorm、d_emb → d_model的hidden_mapping_in投影与 dropout见 layers.py 的EmbeddingLayer。PreTrainingModelmodels.py叠加 MLM 与 NSP 两个预训练任务。MLM 用layers.gather取出掩码位置masked_lm_positions经 Dense → GELU → LayerNorm 后用OutputProjection复用词嵌入矩阵tied weights投影回词表得到 logitsNSP 则基于池化输出做二分类。损失为加权交叉熵MLM 与 NSP 两部分相加。SequenceClassificationModelmodels.py在EncoderModel池化输出上接一个OutputProjection分类头。特殊之处在于对glue/stsb回归任务使用均方误差损失jnp.mean((logits[..., 0] - labels)**2)其余分类任务使用 log-softmax 交叉熵。6.3 训练流程组织训练脚本f_net/run_pretraining.py、f_net/run_classifier.py与 f_net/train_utils.py 协作完成数据加载sentencepiece分词 TensorFlow Datasets 数据管道、参数初始化、优化器构建、JAX 多设备并行训练与 TensorBoard 指标记录输入管道实现在 f_net/input_pipeline.py并配套 f_net/input_pipeline_test.py、f_net/layers_test.py、f_net/models_test.py、f_net/train_utils_test.py 等测试覆盖。七、引用与注意事项如需在论文或技术报告中引用 FNet可使用仓库 README 提供的 BibTeXarticle{lee2021fnet, title{FNet: Mixing Tokens with Fourier Transforms}, author{Lee-Thorp, James and Ainslie, Joshua and Eckstein, Ilya and Ontanon, Santiago}, journal{arXiv preprint arXiv:2105.03824}, year{2021} }最后需要说明本仓库中的 FNet 代码属于 Google Research 的开源研究项目并非 Google 官方产品This is not an official Google product.。在复现实验时请确认你的硬件环境CPU/GPU/TPU、JAX/Flax 版本与requirements.txt中声明的版本范围兼容并按前文要求将工作目录置于f_net的父目录下执行命令。【免费下载链接】google-researchGoogle Research项目地址: https://gitcode.com/gh_mirrors/go/google-research创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考