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.ipynb | Colab / Jupyter 交互式运行版本 |
示例本身可以在 Google Colab 中直接运行并自由修改(仓库内提供了对应的 sst2.ipynb 笔记本),无需本地安装任何依赖即可体验;也可以按本文后续命令在本地环境跑通。
二、环境依赖与数据集
运行本示例需要以下依赖(见 requirements.txt):
absl-py:命令行 flag 与日志clu:平台与工作单元管理flax:核心神经网络库ml-collections:超参数配置与config_flagsnumpy、optax(SGD + 动量 + 权重衰减优化器)tensorflow、tensorflow-datasets、tensorflow-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 | 墙钟时间 | 准确率 | 指标 |
|---|---|---|---|---|---|
| default | TPU | 10 | 4.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_epoch与evaluate_model中的logging.info打印(见 train.py 与 train.py)。注意该结果是默认配置在特定平台上的单次运行输出,实际数值会因硬件、随机种子与框架版本而略有波动,属于可复现的参考基线而非保证值。
四、超参数配置:完整参数表与命令行覆盖
所有默认超参数集中在 configs/default.py,对应一个ml_collections.ConfigDict。下表汇总了每个字段的默认值及其在源码中的实际作用:
| 配置字段 | 默认值 | 作用(对应源码位置) |
|---|---|---|
embedding_size | 300 | 词嵌入维度,见 models.py 中self.param('embedding', ...) |
hidden_size | 256 | LSTM 隐藏单元数与注意力/MLP 隐藏层大小,见 models.py |
vocab_size | None | 词表大小,训练开始时由数据管道自动回填(见 train.py) |
output_size | 1 | 输出维度,SST-2 为二分类,因此输出单个 logit |
vocab_path | 'vocab.txt' | 词表文件路径,见 input_pipeline.py |
max_input_length | 60 | 分桶时的最大序列长度上限 |
dropout_rate | 0.5 | 嵌入层 / 编码输出 / 上下文向量 / MLP 隐藏层的 dropout 比例 |
word_dropout_rate | 0.1 | 词级别 dropout:以 10% 概率将输入词替换为<unk> |
unk_idx | 1 | 词表<unk>的索引,word dropout 替换目标 |
learning_rate | 0.1 | SGD 学习率 |
momentum | 0.9 | SGD 动量 |
weight_decay | 3e-6 | 权重衰减系数 |
batch_size | 64 | 每个 batch 的样本数 |
bucket_size | 8 | 长度分桶粒度(见第五节详解) |
num_epochs | 10 | 训练总轮数 |
seed | 0 | 全局随机种子(参数初始化、数据 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.txt。min_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=64,pad_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 ──► logits6.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)是分类头,包含两个关键组件:
- KeysOnlyMlpAttention(models.py):即 Bahdanau 注意力(论文:Bahdanau et al., 2015, ICLR),仅基于 key 计算注意力分数——将编码输出经过两层无偏置 Dense +
tanh映射为单个标量,用sequence_mask屏蔽 padding 位置(填充位置分数置为-inf,经 softmax 后权重为 0),最后softmax归一化得到每个位置的注意力权重。它还通过self.sow('intermediates', 'attention', scores)将注意力权重注册为可提取的中间量,方便后续可视化与调试。 - 上下文加权求和:
context = attention^T @ encoded_inputs,即对所有时间步的编码向量按注意力权重加权求和,得到整句的摘要向量; - 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:
- 数据准备:分别构造训练集(
split='train')与验证集(split='validation'),训练集使用分桶批处理,验证集使用普通 padded batch; - JIT 编译:
train_step_fn = jax.jit(train_step)、eval_step_fn = jax.jit(eval_step),将核心步函数编译为 XLA 加速版本; - 状态初始化:
create_train_state用optax.chain(optax.sgd(learning_rate, momentum), optax.add_decayed_weights(weight_decay))构建优化器——即带动量的 SGD 叠加权重衰减(注意:add_decayed_weights的实现是往梯度中加权重项,等效于 L2 正则,与 PyTorch 中 AdamW 的解耦式权重衰减语义不同); - 训练步(train.py):
jax.value_and_grad计算损失与梯度,每步用jax.random.fold_in从 epoch 级 RNG 派生新的 dropout RNG,保证每步随机性独立;损失函数采用 sigmoid 交叉熵(sigmoid_cross_entropy_with_logits,通过@jax.vmap逐样本向量化,实现上对数值稳定性做了 logsumexp 处理); - 评估步(train.py):以
deterministic=True前向计算,指标为二分类准确率(logits >= 0与标签比较); - 指标汇总:每个 epoch 结束时
normalize_batch_metrics将各 batch 的损失与正确数累加后除以总样本数,得到 epoch 级平均指标; - TensorBoard 记录:使用 Flax 自带的
flax.metrics.tensorboard.SummaryWriter写入workdir,记录train_loss、train_accuracy、eval_loss、eval_accuracy四个标量,并调用summary_writer.hparams(dict(config))记录完整超参数,便于复现实验。
测试方面,仓库提供了三层验证:models_test.py用init_with_output断言 Embedder / SimpleLSTM / SimpleBiLSTM / TextClassifier 的输出形状(其中 BiLSTM 输出维度验证了前后向拼接后的2 * hidden_size);train_test.py验证jax.jit(train_step)后参数确实发生了更新(逐叶子比较新旧参数);input_pipeline_test.py覆盖数据管道行为。这些测试可作为理解模块接口契约的补充阅读材料。
八、上手路线总结
从零跑通本示例并完成一次超参数实验的最小路径:
- 安装依赖(
pip install -r examples/sst2/requirements.txt); - 首次运行前确保可联网下载
glue/sst2数据集(或提前用tfds手动准备); - 执行
python examples/sst2/main.py --workdir=/tmp/sst2 --config=examples/sst2/configs/default.py完成基线训练; - 通过
--config.learning_rate=...、--config.batch_size=...等 flag 覆盖超参数,观察验证集准确率变化; - 在
--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),仅供参考