资讯详情

Flax 序列标注实战:基于 Transformer 在 Universal Dependencies 上训练词性标注器(examples/nlp_seq 源码全解)

📅 2026/9/17 10:50:52 | 华诺云谱 👁 阅读
Flax 序列标注实战:基于 Transformer 在 Universal Dependencies 上训练词性标注器(examples/nlp_seq 源码全解)
Flax 序列标注实战基于 Transformer 在 Universal Dependencies 上训练词性标注器examples/nlp_seq 源码全解【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax本指南围绕 Flax 官方示例 examples/nlp_seq 展开讲解如何用 Flax Linen 实现一个基于 Transformer Encoder 的词性标注Part-of-Speech Tagging模型并在 Universal Dependencies 语料上完成训练与评估。读完本文你将掌握 CoNLL 数据解析、词表构建、Transformer 序列标注网络搭建、分布式训练循环与学习率调度等一整套可复用的实操方案并能在自己的硬件上复现文档记录的 68.6% 标注准确率基线。示例定位用序列标注任务演示 Flax 的完整工作流examples/nlp_seq是 Flax 仓库中的一个完整、可独立运行的 NLP 序列标注示例。它的任务定义非常直观给定一个句子为每个词预测其词性标签。README 中给出了如下示例句From|ADP the|DT AP|PROPN comes|VBZ this|DT story|NN :|:即 From 被标注为 ADP介词、the 为 DT限定词、AP 为 PROPN专有名词、comes 为 VBZ动词第三人称单数、story 为 NN普通名词。模型在词级别做分类属于经典的序列标注Sequence Tagging任务。该示例目录结构清晰涵盖了一个深度学习项目从数据到训练的全部环节train.py训练主循环包含全部命令行参数定义、学习率调度、训练/评估 stepmodels.pyTransformer Encoder 模型与超参数配置类input_pipeline.pyCoNLL 格式解析、词表构建与 tf.data 批处理管线main.py基于 ml_collections 配置文件的启动入口configs/default.py默认超参数配置文件input_pipeline_test.py数据管线单元测试requirements.txt依赖清单。数据准备下载 Universal Dependencies 语料库README 明确要求使用Universal DependencyUD数据集这是一个跨语言、手工标注的依存语法树库其 CoNLL 格式天然适合词性标注任务。数据获取方式如下下载 UD v2.0 数据包并解压curl -# -o ud-treebanks-v2.0.tgz UD v2.0 数据包下载地址见仓库 README tar xzf ud-treebanks-v2.0.tgz解压后你会得到按语言组织的多个子目录例如示例命令中使用的UD_Ancient_Greek古希腊语其中包含以.conllu结尾的训练集grc-ud-train.conllu与开发集grc-ud-dev.conllu。input_pipeline.py中的CoNLLAttributes枚举定义了 UD CoNLL 文件每一列的含义列索引从 0 开始枚举名列号含义ID0词序号FORM1词形模型输入LEMMA2词元UPOS3通用词性XPOS4语言特定词性模型目标FEATS5形态特征HEAD6依存头DEPREL7依存关系一个典型的 CoNLL 行形如1 They they PRON PRP CaseNom|NumberPlur 2 nsubj。示例默认使用FORM词形作为输入、XPOS语言特定词性标签作为预测目标这一点在 train.py 中可以看到attributes_input [input_pipeline.CoNLLAttributes.FORM] attributes_target [input_pipeline.CoNLLAttributes.XPOS]支持的硬件配置与基线结果README 给出了经过显式验证的软硬件配置及其训练结果同时说明该模型在其他配置与硬件上也应当可以运行只是未逐一测试硬件Batch size学习率训练时长准确率Nvidia Titan V12GB640.055:58h68.6%README 还记录了 TensorBoard.dev 上的实验日志链接2022-05-01说明该基线结果是可以追溯复现的。作为参考这是训练约 7.5 万步后开发集上的准确率属于该任务在当前模型与数据规模下的合理水平。运行命令与全部超参数说明README 给出的核心运行命令为python train.py --batch_size64 --model_dir./ancient_greek \ --devud-treebanks-v2.0/UD_Ancient_Greek/grc-ud-dev.conllu \ --trainud-treebanks-v2.0/UD_Ancient_Greek/grc-ud-train.conllu其中--train与--dev为必填参数缺失时 train.py 会抛出UsageError其余参数均有默认值。所有参数在 train.py 中以 absl flags 形式定义汇总如下参数类型默认值说明--model_dirstring模型日志与输出目录--experimentstringxpos实验名称用于区分 TensorBoard 日志子目录--batch_sizeint64训练 batch 大小--eval_frequencyint100每多少步做一次评估--num_train_stepsint75000总训练步数--learning_ratefloat0.05基础学习率--weight_decayfloat1e-1AdamW 风格权重衰减系数--max_lengthint256句子最大长度超出截断/桶大小--random_seedint0PRNG 随机种子--trainstring训练数据.conllu路径--devstring开发集.conllu路径除了直接传参你还可以通过 main.py 以配置文件方式启动。它通过ml_collections.config_flags.DEFINE_config_file加载 configs/default.py后者用ml_collections.ConfigDict声明了与上述 flag 一一对应的默认超参数model_dir、experiment、batch_size、num_train_steps、eval_frequency、learning_rate、weight_decay、max_length、random_seed、train、dev再统一转写入 FLAGS 后调用train.main(argv)。这意味着你可以通过--configconfigs/default.py轻松切换不同的实验配置。依赖与运行前提requirements.txt 列出了示例的依赖版本absl-py、flax、numpy与tensorflow数据管线依赖 tf.data。运行前请确认已安装 JAX 及其 GPU/TPU 版本示例源码使用jax.pmap做多设备数据并行batch_size必须能被jax.device_count()整除否则 train.py 会直接报错TensorFlow 仅用于数据读取train.py 通过tf.config.experimental.set_visible_devices([], GPU)主动禁止 TF 占用 GPU 显存把显存全部留给 JAX训练数据与开发数据均为 UD 格式的.conllu文件。模型架构源码解析为序列标注定制的 Transformer Encoder模型定义在 models.py 中核心是一个只含 Encoder 的 Transformer其数据流向为输入词 ID → Embed → Dropout → AddPositionEmbs → N × Encoder1DBlock → LayerNorm → Dense(logits)超参数配置类 TransformerConfig所有模型超参数通过 TransformerConfig基于flax.struct.dataclass集中管理默认值与说明如下字段默认值含义vocab_size必填输入词表大小由len(vocabs[forms])决定output_vocab_size必填输出标签数由len(vocabs[xpos])决定dtypejnp.float32计算精度emb_dim512词嵌入维度num_heads8注意力头数num_layers6Encoder 层数qkv_dim512注意力 Q/K/V 维度mlp_dim2048MLP 隐藏层维度max_len2048位置编码最大长度dropout_rate0.3通用 dropout 概率attention_dropout_rate0.3注意力 dropout 概率kernel_initxavier_uniform()Dense/注意力权重初始化bias_initnormal(stddev1e-6)偏置初始化posemb_initNone为None时使用固定正弦位置编码在 train.py 中vocab_size与output_vocab_size由语料实际构建出的词表长度动态填充max_len则被设置为FLAGS.max_length默认 256。模型前向传播Transformer 模块的__call__实现如下x inputs.astype(int32) x nn.Embed(num_embeddingsconfig.vocab_size, featuresconfig.emb_dim, nameembed)(x) x nn.Dropout(rateconfig.dropout_rate)(x, deterministicnot train) x AddPositionEmbs(config)(x) for _ in range(config.num_layers): x Encoder1DBlock(config)(x, deterministicnot train) x nn.LayerNorm(dtypeconfig.dtype)(x) logits nn.Dense(config.output_vocab_size, ...)(x)关键细节输入为(batch, len)的词 ID 张量输出为(batch, len, output_vocab_size)的 logits对每个位置独立做标签分类位置编码默认使用 sinusoidal_init 生成的固定正弦-余弦编码pe[:, 0::2] sin(...)、pe[:, 1::2] cos(...)AddPositionEmbs将其直接加到嵌入输出上若在配置中传入posemb_init则会改为可学习的位置嵌入通过self.param(pos_embedding, ...)声明参数Encoder 层每个 Encoder1DBlock 先做LayerNorm → MultiHeadDotProductAttention → Dropout → 残差连接再做LayerNorm → MlpBlock → 残差连接注意力部分设置了use_biasFalse与broadcast_dropoutFalseMLP 使用 ELU 激活函数两个 Dense 之间及输出处各带一个 dropout模型通过model.apply({params: params}, inputs..., trainTrue, rngs{dropout: rng})调用train标志统一控制所有 dropout 的开关见 train_step。从源码结构看该模型即标准的 Transformer Encoder 堆叠去掉了 Decoder 与因果掩码完全服务于逐词分类的序列标注任务。训练循环源码解析调度、优化器与多设备并行train.py 完整实现了 Flax 典型的分布式训练范式以下几点值得深入理解。自定义学习率调度器create_learning_rate_scheduler 用一个factors字符串描述调度组合支持constant、linear_warmup、rsqrt_decay、rsqrt_normalized_decay、decay_every、cosine_decay六种因子按*分隔并依次相乘。默认调度为constant * linear_warmup * rsqrt_decay即先用 8000 步线性预热到峰值再按1/sqrt(max(step, warmup_steps))衰减这也是 Transformer 论文中常用的 noam 风格调度。示例训练时仅传入base_learning_ratelearning_rate默认 0.05其余参数保持默认。优化器与训练状态优化器采用 optax.adamw参数b10.9, b20.98, eps1e-9, weight_decay1e-1配合flax.training.train_state.TrainState.create统一管理参数、优化器状态与apply_fn。损失与权重掩码train_step 中损失函数对目标做 one-hot 后计算加权交叉熵compute_weighted_cross_entropy并且对每个位置构造权重weights jnp.where(targets 0, 1, 0)——这是关键细节由于 batch 内句子经 padding 对齐词 ID 为 0 的 padding 位置PAD_ID 0会被掩码掉不参与损失与准确率统计normalizing_factor用weights.sum()做分母从而得到真正的平均损失。评估阶段的 eval_step 采用同样的掩码逻辑。多设备数据并行训练采用 SPMD 数据并行state jax_utils.replicate(state)复制参数jax.pmap包装 train/eval step 并指定axis_namebatch梯度经jax.lax.pmean(grads, batch)做跨设备平均输入 batch 用common_utils.shard分发到各设备。donate_argnums(0,)开启参数缓冲区复用以节省显存。评估阶段对最后不足 batch 大小的数据通过 pad_examples 补零到 batch 大小避免丢弃样本。训练-评估循环主循环main每eval_frequency默认 100步做一次评估将累计的 train metrics 汇总求均值写入 TensorBoard含steps per second吞吐指标然后在开发集上计算 loss 与 accuracy并维护best_dev_score记录最优开发集准确率。日志目录为model_dir/experiment_train与model_dir/experiment_eval可用 TensorBoard 直接查看训练曲线。数据管线源码解析从 CoNLL 到可训练的 batchinput_pipeline.py 是整个示例的数据地基主要包含三个环节。1. 词表构建create_vocabscreate_vocabs 扫描语料用collections.Counter统计词形FORM与词性标签XPOS频率然后构建两个映射表vocabs[forms]与vocabs[xpos]。三个特殊 token 固定占用前 3 个 IDpPADpadding→ ID 0uUNKNOWN未登录词→ ID 1rROOT人工根节点→ ID 2其余词形按频率从 ID 3 开始编号且词形表最多保留max_num_forms100000个高频词低频词在编码时会被映射为UNKNOWN_ID。这一设计保证了 OOV未登录词也能被模型处理。2. CoNLL 句子解析sentences_from_conll_dataUD 语料中句子之间以空行分隔、注释行以#开头。解析器为每个句子在最前面插入一个人工根节点create_sentence_with_root即r词、ID 为 0然后逐行读取 token按选定的属性列映射为 ID 列表。max_sentence_length参数用于截断过长的句子默认 1000而在批处理时会被设置为bucket_size。3. 动态 padding 批处理sentence_dataset_dictsentence_dataset_dict 基于 tf.data 构建流水线用tf.data.Dataset.from_generator从生成器产出{inputs: [...], targets: [...]}.cache()将数据集缓存到内存训练多轮 epoch 时显著提速.repeat(repeat)控制重复次数训练时无限重复评估时repeat1.padded_batch(batch_sizebatch_size, padded_shapes[bucket_size])做静态 padding每个句子都被补齐到bucket_size即max_lengthpadding 位置填 0正好与训练时的权重掩码targets 0配合.prefetch(AUTOTUNE)预取下一批数据。当attributes_target为空列表时数据集只包含inputs键可无缝用于无标注数据的推理场景。测试佐证input_pipeline_test.py 用一段两句话的 CoNLL 样例验证了上述逻辑test_vocab_creation断言词表映射结果testInputBatch断言输入 batch 的 padding 布局形如[2, 3, 4, 5, 6, 0, ...]开头是 ROOT2testInputTargetBatch断言 input/target 两个键同时存在且标签序列正确。这些测试既是数据管线正确性的保证也是理解其行为的最好注解。如何验证、运行与扩展运行单元测试在示例目录下执行python -m input_pipeline_test或使用 absl 测试框架可快速验证数据管线在本地环境是否正常完整训练按上文命令指定--train/--dev与--model_dir即可若机器为单 GPU/CPU需保证--batch_size能被设备数整除查看日志训练与评估指标写入model_dir下的 TensorBoard 事件文件README 亦提供了与基线对应的 TensorBoard.dev 实验链接供对照扩展到其他语言/任务UD v2.0 包含数十种语言的树库把--train/--dev指向其他语言的.conllu即可修改train.py中的attributes_input/attributes_target如改用LEMMA、UPOS、HEAD列即可训练词元还原、通用词性标注甚至依存头预测等变体任务而模型与训练框架无需改动。小结examples/nlp_seq是一个麻雀虽小、五脏俱全的 Flax 序列标注参考实现它把 UD 语料解析、tf.data 批处理、Transformer Encoder 建模、masked 交叉熵训练与多设备数据并行完整串联起来并提供了可复现的硬件基线Nvidia Titan V、batch 64、lr 0.05、约 6 小时、68.6% 准确率。无论你是想学习 Flax 的工程范式还是需要一个开箱即用的词性标注起点都可以直接从该示例的 train.py、models.py 与 input_pipeline.py 三个核心文件入手阅读与改造。【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。