Flax 文本分类实战:基于 SST-2 情感分类示例训练 BiLSTM + Attention 分类器
2026/9/17 13:01:44 网站建设 项目流程

Flax 文本分类实战:基于 SST-2 情感分类示例训练 BiLSTM + Attention 分类器

【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax

本文围绕 Flax 官方示例 examples/sst2 展开,完整讲解如何在 JAX + Flax 中训练一个端到端的文本情感分类器:从数据管道、词表构建、模型架构(Embedding + BiLSTM + Bahdanau 注意力 + MLP 分类头)、超参数配置到训练评估闭环。读完本文,你可以直接运行该示例复现 85%+ 的验证集准确率,掌握config_flags命令行覆盖超参数的实战技巧,并能对照源码理解每个配置项在底层的作用。

一、示例概览:在 Colab 或本地训练 SST-2 分类器

SST-2(Stanford Sentiment Treebank 二分类子集)是 GLUE 基准中的经典情感分类任务,每条样本是一条英文电影评论,标签为 0(负面)或 1(正面)。Flax 仓库中的 sst2 示例 用不到 10 个 Python 文件搭建了一个完整的文本分类训练管线:

文件职责
main.py入口脚本,解析--workdir--config命令行参数
configs/default.py默认超参数配置(ml_collections)
train.py训练 / 评估循环、指标计算、优化器构建
models.py模型定义:Embedder、BiLSTM、注意力分类器
input_pipeline.py数据加载、分词、长度分桶(bucketed batching)
vocabulary.py词表构建、加载、保存
build_vocabulary.py从训练集生成vocab.txt的独立脚本
sst2.ipynbColab / Jupyter 交互式运行版本

示例本身可以在 Google Colab 中直接运行并自由修改(仓库内提供了对应的 sst2.ipynb 笔记本),无需本地安装任何依赖即可体验;也可以按本文后续命令在本地环境跑通。

二、环境依赖与数据集

运行本示例需要以下依赖(见 requirements.txt):

  • absl-py:命令行 flag 与日志
  • clu:平台与工作单元管理
  • flax:核心神经网络库
  • ml-collections:超参数配置与config_flags
  • numpyoptax(SGD + 动量 + 权重衰减优化器)
  • tensorflowtensorflow-datasetstensorflow-text:数据管道与分词

数据集方面有一个关键特性:TensorFlow datasetglue/sst2会在首次运行时被自动下载并准备(prepared),无需手动下载。input_pipeline.py中通过tfds.load('glue/sst2', split='train')加载,并在TextDataset.__init__中自动识别数据集的文本字段(tfds.features.Text)与标签字段(tfds.features.ClassLabel),因此该数据管道设计上对"单文本 + 单标签"类数据集具有通用性。

三、快速开始:训练命令与基准结果

示例的入口设计非常简洁:main.py只负责解析参数并调用train.train_and_evaluate(config, workdir)(对应 main.py),核心逻辑全部集中在可被 Colab 导入、可被单元测试覆盖的库文件中。启动训练只需一条命令:

python main.py --workdir=/tmp/sst2 --config=configs/default.py

两个 flag 均为必填(见 main.py):

  • --workdir:模型数据与 TensorBoard 日志的输出目录;
  • --config:指向超参数配置文件的路径,由config_flags.DEFINE_config_file(..., lock_config=True)定义,lock_config=True意味着配置字段在训练期间不可被运行时篡改。

仓库 README 给出了默认配置在 TPU 平台上的参考结果:

名称平台Epochs墙钟时间准确率指标
defaultTPU104.3 分钟85.21%TensorBoard 日志

训练过程中日志按 epoch 输出,格式如下:

INFO:absl:train epoch 010 loss 0.1918 accuracy 92.41 INFO:absl:eval epoch 010 loss 0.4144 accuracy 85.21

其中train accuracy是训练集上的准确率(约 92.41%),eval accuracy是验证集(SST-2 官方 validation split)上的准确率(85.21%)。训练与评估日志分别由train_epochevaluate_model中的logging.info打印(见 train.py 与 train.py)。注意该结果是默认配置在特定平台上的单次运行输出,实际数值会因硬件、随机种子与框架版本而略有波动,属于可复现的参考基线而非保证值。

四、超参数配置:完整参数表与命令行覆盖

所有默认超参数集中在 configs/default.py,对应一个ml_collections.ConfigDict。下表汇总了每个字段的默认值及其在源码中的实际作用:

配置字段默认值作用(对应源码位置)
embedding_size300词嵌入维度,见 models.py 中self.param('embedding', ...)
hidden_size256LSTM 隐藏单元数与注意力/MLP 隐藏层大小,见 models.py
vocab_sizeNone词表大小,训练开始时由数据管道自动回填(见 train.py)
output_size1输出维度,SST-2 为二分类,因此输出单个 logit
vocab_path'vocab.txt'词表文件路径,见 input_pipeline.py
max_input_length60分桶时的最大序列长度上限
dropout_rate0.5嵌入层 / 编码输出 / 上下文向量 / MLP 隐藏层的 dropout 比例
word_dropout_rate0.1词级别 dropout:以 10% 概率将输入词替换为<unk>
unk_idx1词表<unk>的索引,word dropout 替换目标
learning_rate0.1SGD 学习率
momentum0.9SGD 动量
weight_decay3e-6权重衰减系数
batch_size64每个 batch 的样本数
bucket_size8长度分桶粒度(见第五节详解)
num_epochs10训练总轮数
seed0全局随机种子(参数初始化、数据 shuffle)

4.1 命令行覆盖超参数

示例通过ml_collections的 config_flags 支持在命令行直接覆盖任意配置字段,无需修改配置文件

python main.py \ --workdir=/tmp/sst2 --config=configs/default.py \ --config.learning_rate=0.05 --config.num_epochs=5

语法规则是--config.<字段路径>=<值>--config.learning_rate=0.05将学习率改为 0.05,--config.num_epochs=5将训练轮数改为 5。该机制对所有配置字段通用,例如同样可以覆盖--config.batch_size=128--config.dropout_rate=0.3,是进行超参数实验(如扫描学习率、batch size、dropout 对准确率的影响)的标准方式。

4.2 一个值得注意的细节:vocab_size 的运行时回填

配置中vocab_size=None,但 Embedder 需要真实词表大小来分配嵌入矩阵。这一矛盾在 train.py 中解决:

# Keep track of vocab size in the config so that the embedder knows it. config.vocab_size = len(train_dataset.vocab)

即训练开始时先用训练集词表长度回填配置,再据此构建模型。这保证了配置文件的简洁性,同时避免硬编码词表大小导致的数据集版本不匹配问题。

五、数据管道源码解析:分词、词表与长度分桶

5.1 词表构建与特殊符号

TextDataset(input_pipeline.py)在初始化时加载 vocab.txt,并使用tf.lookup.StaticHashTable将 token 映射为整数 ID(未知词默认映射到unk_idx)。仓库附带的vocab.txt共 13523 行,前 4 行是固定特殊符号:

<pad> <unk> <s> </s> the

对应 vocabulary.py 中定义的四个特殊 token:<pad>(填充)、<unk>(未知词)、<s>(序列起始 BOS)、</s>(序列结束 EOS),其索引依次为 0、1、2、3,这也是配置中unk_idx=1的由来。每条样本在prepare_example中会通过add_bos_eos在 token 序列首尾追加 BOS 与 EOS:

def add_bos_eos(self, sequence): return tf.concat([[self.vocab.bos_idx], sequence, [self.vocab.eos_idx]], 0)

如果你需要针对新的语料重建词表,可运行 build_vocabulary.py:它加载glue/sst2的训练集,用WhitespaceTokenizer分词后统计词频,只保留出现次数 ≥ 3(min_freq=3)的词,再按"词频降序、token 字典序"排序写入vocab.txtmin_freq阈值的目的是过滤稀有词、防止过拟合,可以按需调整。词表构建的核心实现在 vocabulary.py,其save/load均为每行一个 token 的纯文本格式。

5.2 长度分桶(Bucketed Batching)与填充

由于句子长度不一,直接 pad 会浪费大量算力。示例采用tf.data.experimental.bucket_by_sequence_length实现按长度分桶 + 桶内填充(input_pipeline.py):

  • bucket_size=8表示每 8 个长度级别共享一个桶:对max_input_length=60,桶边界为[9, 17, 25, 33, 41, 49, 57, 65](见get_bucket_boundaries的示例);
  • 每个桶使用完整batch_size=64pad_to_bucket_boundary=True使填充量最小化;
  • 填充形状由padded_shapes定义:{'idx': [], 'token_ids': [None], 'label': [], 'length': []},其中None表示动态填充到该 batch 内最长序列;
  • 训练数据每个 epoch 重新 shuffle(reshuffle_each_iteration=True),由config.seed控制,保证实验可复现。

六、模型架构源码解析:Embedding → BiLSTM → Attention → MLP

模型定义在 models.py,整体结构为:

token_ids ──► Embedder ──► SimpleBiLSTM ──► AttentionClassifier ──► logits

6.1 Embedder:词嵌入 + 双重 dropout

Embedder(models.py)由三个子模块构成:

  • 词嵌入矩阵self.param('embedding', nn.initializers.normal(stddev=0.1), (vocab_size, embedding_size)),默认用标准差 0.1 的正态分布初始化;
  • WordDropout:以word_dropout_rate=0.1的概率把输入 token 替换为unk_idx(自定义模块 WordDropout,等价于nn.Dropout但允许指定被丢弃元素的值);
  • 标准 dropout:嵌入后以dropout_rate=0.5施加nn.Dropout

此外frozen=True时会对嵌入输出施加jax.lax.stop_gradient,用于保持预训练词向量的固定取值,默认关闭。

6.2 SimpleBiLSTM:scan 展开的双向 LSTM

SimpleLSTM(models.py)使用 Flax 的nn.transforms.scan沿时间轴(in_axes=1, out_axes=1)扫描展开单个nn.OptimizedLSTMCell,参数跨时间步广播(variable_broadcast='params'),从而以常量显存处理任意长度序列

SimpleBiLSTM(models.py)则由前向 LSTM 与后向 LSTM 组成。处理变长序列时,朴素地翻转 padded 序列会把 padding 移到头部,因此示例实现了flip_sequences(models.py):利用jnp.roll+jnp.flip在翻转真实 token 的同时保持 padding 留在序列末尾。后向 LSTM 输出再翻转回原方向,最后将前后向表示沿最后一维拼接,得到维度2 * hidden_size的编码。

6.3 AttentionClassifier:Bahdanau 注意力汇总 + MLP 分类

AttentionClassifier(models.py)是分类头,包含两个关键组件:

  1. KeysOnlyMlpAttention(models.py):即 Bahdanau 注意力(论文:Bahdanau et al., 2015, ICLR),仅基于 key 计算注意力分数——将编码输出经过两层无偏置 Dense +tanh映射为单个标量,用sequence_mask屏蔽 padding 位置(填充位置分数置为-inf,经 softmax 后权重为 0),最后softmax归一化得到每个位置的注意力权重。它还通过self.sow('intermediates', 'attention', scores)将注意力权重注册为可提取的中间量,方便后续可视化与调试。
  2. 上下文加权求和context = attention^T @ encoded_inputs,即对所有时间步的编码向量按注意力权重加权求和,得到整句的摘要向量;
  3. MLP 分类器(models.py):Dense(hidden_size) → tanh → Dropout → Dense(output_size, use_bias=False),最终输出 1 个 logit(二分类),经 sigmoid 交叉熵计算损失。

TextClassifier__call__完整串联了"嵌入 → 编码 → 注意力汇总 → 分类"的流程,同时所有子层共享deterministic开关:训练时为False(开启 dropout),评估时为True(关闭 dropout),该开关通过nn.module.merge_param实现模块级默认值与调用级参数的安全合并。

七、训练与评估循环源码解析

训练逻辑集中在 train.py:

  1. 数据准备:分别构造训练集(split='train')与验证集(split='validation'),训练集使用分桶批处理,验证集使用普通 padded batch;
  2. JIT 编译train_step_fn = jax.jit(train_step)eval_step_fn = jax.jit(eval_step),将核心步函数编译为 XLA 加速版本;
  3. 状态初始化create_train_stateoptax.chain(optax.sgd(learning_rate, momentum), optax.add_decayed_weights(weight_decay))构建优化器——即带动量的 SGD 叠加权重衰减(注意:add_decayed_weights的实现是往梯度中加权重项,等效于 L2 正则,与 PyTorch 中 AdamW 的解耦式权重衰减语义不同);
  4. 训练步(train.py):jax.value_and_grad计算损失与梯度,每步用jax.random.fold_in从 epoch 级 RNG 派生新的 dropout RNG,保证每步随机性独立;损失函数采用 sigmoid 交叉熵(sigmoid_cross_entropy_with_logits,通过@jax.vmap逐样本向量化,实现上对数值稳定性做了 logsumexp 处理);
  5. 评估步(train.py):以deterministic=True前向计算,指标为二分类准确率(logits >= 0与标签比较);
  6. 指标汇总:每个 epoch 结束时normalize_batch_metrics将各 batch 的损失与正确数累加后除以总样本数,得到 epoch 级平均指标;
  7. TensorBoard 记录:使用 Flax 自带的flax.metrics.tensorboard.SummaryWriter写入workdir,记录train_losstrain_accuracyeval_losseval_accuracy四个标量,并调用summary_writer.hparams(dict(config))记录完整超参数,便于复现实验。

测试方面,仓库提供了三层验证:models_test.pyinit_with_output断言 Embedder / SimpleLSTM / SimpleBiLSTM / TextClassifier 的输出形状(其中 BiLSTM 输出维度验证了前后向拼接后的2 * hidden_size);train_test.py验证jax.jit(train_step)后参数确实发生了更新(逐叶子比较新旧参数);input_pipeline_test.py覆盖数据管道行为。这些测试可作为理解模块接口契约的补充阅读材料。

八、上手路线总结

从零跑通本示例并完成一次超参数实验的最小路径:

  1. 安装依赖(pip install -r examples/sst2/requirements.txt);
  2. 首次运行前确保可联网下载glue/sst2数据集(或提前用tfds手动准备);
  3. 执行python examples/sst2/main.py --workdir=/tmp/sst2 --config=examples/sst2/configs/default.py完成基线训练;
  4. 通过--config.learning_rate=...--config.batch_size=...等 flag 覆盖超参数,观察验证集准确率变化;
  5. --workdir目录下用 TensorBoard 查看训练 / 验证曲线,或直接修改 sst2.ipynb 在 Colab 中交互式迭代。

这个示例完整展示了 Flax 的核心工程范式:ml_collections 驱动的可配置训练、Flax Linen 的模块化建模(含scan展开循环网络)、TrainState状态管理、jax.jit编译加速以及 TensorBoard 指标记录,是学习用 Flax 搭建 NLP 训练管线的最佳入门模板之一。

【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询