首页
/ Sequence Projection Models 实战指南:用 seq_flow_lite 在 TensorFlow 中训练端侧文本分类与语言检测模型

Sequence Projection Models 实战指南:用 seq_flow_lite 在 TensorFlow 中训练端侧文本分类与语言检测模型

2026-09-07 23:31:09作者:齐添朝

本文以 TensorFlow Models 仓库中 research/seq_flow_lite/README.md 为骨架,系统讲解 Sequence Projection Models(序列投影模型)的核心思想、PRADO / SGNN 等模型的训练、评估与 TFLite 部署全流程,并结合仓库源码与真实配置文件(configs/civil_comments_prado.txt 等)逐项解读参数含义。读者读完后将掌握该目录下从 JSON 配置编写、bazel 训练命令、源码级网络结构到端侧 TFLite 模型导出的完整实操路径。

一、核心思想:把文本序列“投影”为定长特征

Sequence Projection Models 是 seq_flow_lite 提供的一族模型,其设计目标与常规 NLP 模型有本质区别。正如 README 所述:

We provide a family of models that projects sequence to fixed sized features. The idea behind is to build embedding-free models that minimize the model size. Instead of using embedding table to lookup embeddings, sequence projection models computes them on the fly.

即:构建无 Embedding 表(embedding-free)、体积最小的端侧模型。传统模型需要维护一张巨大的 embedding table 去“查表”获得词向量,而序列投影模型则通过哈希加卷积等手段在推理时动态计算特征,从而把词表规模对模型体积的影响彻底消除。这让模型可以轻松在手机等端侧设备上运行,不需要携带动辄数百万参数的词向量矩阵。

整个实现充分考虑了端侧部署的可量化性(Quantization)与整数算术推理(Integer-Arithmetic-Only),这在后文源码与配置中都会反复出现。

二、仓库覆盖的模型家族与论文依据

该目录是以下三篇论文的官方实现载体(论文标题与出处见 README 开头,仓库不附带其在线链接,需自行在 ACL/arXiv 检索原文):

论文 主题 在本仓库的代表实现
PRADO: Projection Attention Networks for Document Classification On-Device 面向端侧文档分类的投影注意力网络 models/prado.py
Self-Governing Neural Networks for On-Device Short Text Classification 面向端侧短文本分类的自治理神经网络 models/sgnn/
Tiny Neural Models for Seq2Seq 微小 Seq2Seq 模型 模型目录中的 PQRNN / ByteQRNN / Charformer / Transformer 编码解码相关文件

models/ 目录结构看,实现远不止 README 提到的三种:

  • prado.py:PRADO 编码器;
  • sgnn/:SGNN(含自定义 TFLite 算子的 C++ 实现 sgnn_projection.cc 与测试);
  • pqrnn.py / byteqrnn.py / charformer.py:投影类 QRNN、字节级 QRNN 与 Charformer 变体;
  • transformer_encoder.py / transformer_uniform_attn_decoder.py:微小的 Transformer 式 Seq2Seq 编解码组件。

README 中给出的完整引用列表还包括两篇基础文献——Batch Normalization(Ioffe & Szegedy, ICML 2015)与整数算术量化训练(Jacob et al., CVPR 2018),前者是工程上的归一化手段,后者则是整套 TFLite 量化方案的依据,体现该项目“小而可量化”的技术取向。

三、仓库结构速览:从数据到算子的分层布局

在写训练命令前,先理解代码的组织方式,便于定位每一层:

research/seq_flow_lite/
├── BUILD / WORKSPACE          # Bazel 构建与依赖定义
├── trainer.py / trainer_v2.py # PRADO 训练入口(TF1 estimator 风格 / TF2)
├── input_fn_reader.py         # TFDS 数据集读取、标签拼接与文本预处理
├── export_to_tflite.py        # 从 checkpoint 导出 TFLite 模型
├── metric_functions.py        # 评估指标
├── configs/                   # JSON 格式 RunnerConfig(PRADO 任务配置)
├── models/                    # 各模型定义(prado/sgnn/pqrnn/charformer 等)
├── layers/                    # 可量化基础层(卷积/全连接/投影/归一化等)
├── tf_ops/  tflite_ops/       # TensorFlow 自定义算子与 TFLite 对应算子
├── utils/  demo/

其中 layers/ 是“乐高积木”,包括 projection_layers(投影)、conv_layersdense_layersnormalization_layersquantization_layerstransformer_layers 等,全部面向量化训练与 TFLite 导出设计;tf_opstflite_ops 则提供训练图与 TFLite 推理图都能识别的自定义算子(如 SGNN 的投影算子 sgnn_projection.cc)。

四、环境与构建要求

  • TensorFlow 2.3Python 3.6README “Requirements” 一节标注的徽章版本)。
  • 训练、评估命令均通过 Bazel 运行,因此需要先配置好 Bazel 工作区(目录自带 WORKSPACE 与各层 BUILD 文件)。
  • 数据处理依赖 tensorflow_datasetstensorflow_text,训练 PRADO 模型前需保证对应数据集(如 civil_comments)可下载。
  • 项目许可证为 Apache 2.0。

五、训练实战:用 RunnerConfig JSON 驱动 PRADO

5.1 在 Civil Comments 数据集上训练

以“辱骂评论检测”为经典场景,README 给出的训练命令为:

bazel run -c opt :trainer -- \
--config_path=$(pwd)/configs/civil_comments_prado.txt \
--runner_mode=train --logtostderr --output_dir=/tmp/prado

这条命令的语义是:在 seq_flow_lite 根目录下(BUILD 已定义 :trainer 这个 py_binary),把 configs/civil_comments_prado.txt 这份 JSON 作为 RunnerConfig 传入,以 train 模式训练,checkpoint 等产物写到 /tmp/prado

注意:.txt 后缀不代表文本格式——它实际是 JSON。从 trainer_v2.pyload_runner_config() 可以看到,入口就是直接读取该文件并 json.loads 解析:

def load_runner_config():
  with tf.io.gfile.GFile(FLAGS.config_path, "r") as f:
    return json.loads(f.read())

5.2 RunnerConfig 的运行时参数

trainer_v2.py 定义了训练器自身的 CLI 标志位,即命令行 -- 之后跟的参数:

参数 取值 默认值 说明
config_path 字符串 None RunnerConfig JSON 文件路径
runner_mode train / train_and_eval / eval train 运行模式
output_dir 字符串 /tmp/testV2 checkpoint 输出目录
master 字符串 None TensorFlow master URL
use_tpu 布尔 False 是否使用 TPU(否则 CPU/GPU)
num_tpu_cores 整数 8 仅在 use_tpu=True 时生效

TF2 训练器的 main()trainer_v2.py)展示了核心调用链:model_fn_builder 依据 config 中的 name 字段用 importlib 动态导入模型模块并构建 Encoder;create_input_fn 构造数据集;随后进入 tf.function 装饰的 train_step:先计算 logits,再经 compute_loss 求损失、GradientTape 求梯度、Adam 优化器更新参数,每 100 步打印一次 loss。

5.3 model_config:PRADO 网络结构参数全解

configs/civil_comments_prado.txt 为例,model_config 直接控制网络形态,逐一解读如下(可与 models/prado.py_get_params 的参数定义相互印证):

参数 示例值 含义
labels 7 个毒性类别名 输出类别列表,决定分类头维度 num_classes = len(labels)
multilabel true 是否为多标签:是则用 sigmoid 交叉熵,否则用 sparse softmax 交叉熵(对应 trainer_v2.pycompute_loss 分支)
quantize true 是否在训练中启用量化感知(QAT),由 base_layers.Parameters(quantize=...) 传递到每一层
max_seq_len 128 训练时文本最大长度,超过则截断
max_seq_len_inference 128 推理时最大长度
split_on_space true 是否按空格切词后统计序列长度
embedding_regularizer_scale 35e-3 投影/嵌入层正则化系数
embedding_size 64 投影特征映射到的维数,对应 values_fc/attention_fc 全连接输出单元数
bigram_channels 64 bigram 注意力池化通道数
trigram_channels 64 trigram 注意力池化通道数
feature_size 512 特征维度(与各 n-gram 通道组合语义相关)
network_regularizer_scale 1e-4 网络主体(卷积/注意力池化)正则化系数
keep_prob 0.5 训练期 dropout 保留概率,<1.0 时才启用 dropout
distortion_probability 0.25 文本随机变形/扰动概率

models/prado.py 源码可见,PRADO Encoder 实际读取的参数不止上述几个,还支持 unigram_channelsfourgram_channelsfivegram_channelsskip1bigram_channelsskip2bigram_channels(默认 0,即不启用);这印证了配置与实现之间的松耦合关系——model_config 中出现 bigram_channels/trigram_channels,是因为 _get_params 把二者写成了 64 而非默认的 0。

5.4 训练级参数(RunnerConfig 顶层字段)

两份样例配置的顶层还包含优化与调度设置,含义如下:

参数 civil_comments go_emotion 说明
name models.prado models.prado 模型模块名(importlib 导入对象)
batch_size 1024 1024 训练 batch
save_checkpoints_steps 100 100 每 N 步保存 checkpoint
train_steps 100000 100000 训练总步数
learning_rate 1e-3 0.0006 初始学习率
learning_rate_decay_steps 42000 340 学习率衰减周期(步)
learning_rate_decay_rate 0.7 0.7 每个周期的衰减系数
iterations_per_loop 100 100 每循环的迭代数(TPU/训练循环粒度)
dataset civil_comments goemotions TFDS 数据集名,同时作为 input_fn_reader 中处理器函数名

最后一项 dataset 非常关键:在 input_fn_reader.py_post_processor 中,会通过 getattr(sys.modules[__name__], runner_config["dataset"]) 动态调用同名处理器。例如 civil_comments 处理器(input_fn_reader.py)会按 labels 顺序把各毒性标签 stack[batch, num_classes] 的张量并取整,goemotionsinput_fn_reader.py)则读取 comment_text 字段。随后文本会经过 misc_utils.random_substr 截断、ProjectionLayer 投影得到 projectionseq_length,并用 ByteSplitter 产出 token_ids 供需要字节输入的模型使用。

5.5 第二个 PRADO 任务:GoEmotions 情绪识别

仓库还附带了一份多标签情绪识别配置 configs/go_emotion_prado.txtlabels 覆盖 28 个细粒度情绪类别,multilabel=true,训练超参有两处显著差异:学习率降为 0.0006、衰减周期骤降为 340 步、distortion_probability 置为 0.0。可作为对比调参的参考样例,替换命令中的 --config_path 即可复用同一训练入口。

六、源码级解析:PRADO 编码器为什么“小”

6.1 三层结构

PRADO 的 Encodermodels/prado.py)核心结构可概括为:

  1. 投影ProjectionLayer 把输入文本投影为定长特征序列(而非大词表 embedding lookup);
  2. 并行 n-gram 注意力池化:把输入通过 values_fcattention_fc 两个全连接分别映射为“值”与“注意力 logits”,再喂给多个不同 n-gram 宽度的 AttentionPoolReduce
  3. 拼接量化 + 分类头:所有 n-gram 分支输出经 ConcatQuantization 拼接,最后由 final_fcBaseQDense)映射到 num_classes

其中 AttentionPoolReducemodels/prado.py)内含两个 PaddedMaskedVarLenConv:一个生成 value、一个生成 attention logits;推理(TFLITE 模式)时调用自定义算子 expected_value_op 完成注意力加权求和,训练时则用 softmax 与加权求和实现等价计算。

6.2 可变长掩码与量化贯穿始终

PaddedMaskedVarLenConvmodels/prado.py)体现了两个工程细节:

  • skip-gram 支持:可通过 skip_bigram 参数构造跳字卷积核(把 mask[1]mask[skip_bigram] 置 0),ngram 需在 1–5 之间;
  • 量化参数quantize_parameter 在量化后的权重上再乘卷积核 mask,保证量化与跳字结构兼容。

配合 mask(padding 处为 0)与 inverse_normalizer(序列长度的倒数)归一化,模型天然支持变长输入。训练期还会以 invalid_logit 填充无效位置,推理期直接置 0——这些细节共同保证了“小体积 + 高精度 + 可量化”。

models/prado.pyassert tensors 可以看到:若所有 n-gram 通道均为 0,网络会直接报错,提示至少配置一种 n-gram 通道。

七、评估:验证 PRADO 与运行 SGNN 语言检测

7.1 评估 PRADO

runner_mode 切换为 eval 即可用测试集评估同一个 checkpoint 目录:

bazel run -c opt :trainer -- \
--config_path=$(pwd)/configs/civil_comments_prado.txt \
--runner_mode=eval --logtostderr --output_dir=/tmp/prado

output_dir 指向训练产物目录即可复用 checkpoint。此时 trainer_v2.py 中的 create_input_fn 会读取 split="test"(见 input_fn_reader.py),并以 drop_remainder 控制的批处理方式送入模型。

7.2 SGNN:语言检测的独立小模型

SGNN 是独立于 :trainer 的另一条训练线。先训练语言检测模型:

bazel run -c opt sgnn:train -- --logtostderr --output_dir=/tmp/sgnn

SGNN 训练入口 models/sgnn/train.py 会从 TFDS 的 wikipedia/20190301.{ar,en,es,fr,ru,zh} 六个语言子集采样文本,构建语言识别数据(代码中硬编码了 LANGIDS 与 6 个语言 id)。该脚本暴露的 CLI 参数很直观:

  • projection_size(默认 600):投影层大小,决定哈希投影后的特征数;
  • ngram_size(默认 3):投影特征的最大 n-gram 长度;
  • fc_layer(默认 256,128):全连接层尺寸,逗号分隔;
  • batch_size(默认 160)、epochs(默认 10)、learning_rate(默认 2e-4)。

训练完成后,脚本内的 save_and_convertmodels/sgnn/train.py)会先保存 SavedModel,再用 TFLiteConverter 转换,并通过 allow_custom_ops=True + SELECT_TF_OPS 保留自定义投影算子,最终产出 /tmp/sgnn/model.tflite

评估(实际是端侧推理体验)SGNN 模型的方法,是直接跑编译好的 TFLite runner:

bazel run -c opt sgnn:run_tflite -- --model=/tmp/sgnn/model.tflite "Hello world"

传入 --model 指向上述 model.tflite,并附上一句待识别文本,程序会返回其语言预测结果。SGNN 之所以能进 TFLite 而不丢算子,依赖 models/sgnn/ 中成对出现的 sgnn_projection.cc(TensorFlow 算子)、sgnn_projection_op_resolver.cc(TFLite op resolver)以及 sgnn_projection_test.cc(C++ 侧测试),是“训练图与推理图算子对齐”的典型工程实现。

八、导出部署:把 PRADO checkpoint 转为 TFLite

SGNN 走的是自带的 keras 导出路径,而 PRADO 的导出由独立的 export_to_tflite.py 完成。它的机制是:读取 output_dir 下的 runner_config.txt(训练时自动保存的配置副本)重新构建 TFLITE 模式图,从最新 checkpoint 恢复权重,然后导出 flatbuffer。其 CLI 标志为:

  • output_dir:模型/checkpoint 目录;
  • output:输出张量类型,取值 logits / sigmoid / softmax,默认 sigmoid(多标签任务通常选 sigmoid,softmax 适合单标签)。

导出产物为 output_dir/tflite.fb。值得注意的细节(export_to_tflite.py):

  • 输入占位符为 tf.placeholder(tf.string, shape=[1], name="Input"),即直接吃原始文本字符串,无需在端侧维护词表——这正是 embedding-free 设计对部署的回报;
  • 非 PQRNN 类模型(含 PRADO)走 ByteSplitter 字节级 token 化并把每个 token id += 3 做偏移;
  • output 参数决定把 logits 再套 sigmoid 或 softmax,便于端侧直接解释结果。

由于 README 训练/评估命令均以 bazel 为目标执行环境,导出前请确保在工作区内完成构建并保证 TF2.3 与自定义算子编译环境一致。

九、调参与自定义任务的一般路径

综合以上内容,将一个新任务接入 seq_flow_lite 的 PRADO 流程通常只需三步:

  1. 确认数据集在 TFDS 中可用,并在 input_fn_reader.py 中添加一个与数据集同名的处理器函数,把数据集的原始 feature 映射为 (text, label)(label 已按 labels 顺序栈成 [batch, num_classes]);
  2. 编写一份 RunnerConfig JSON,参照 configs/civil_comments_prado.txt 配置 labelsmultilabel、n-gram 通道数、batch_size 与学习率调度;
  3. train → eval → export_to_tflite 顺序依次运行对应 bazel 命令。

其中 multilabel 与输出层激活的选择需一致:多标签任务(如毒性评论的 7 个标签同时成立)训练用 sigmoid 交叉熵、导出建议 --output=sigmoid;单标签任务训练用 softmax 交叉熵、导出建议 --output=softmax

十、小结

seq_flow_lite 的 Sequence Projection Models 通过“在线投影替代 embedding 查表”这一设计,把模型体积压缩到可端侧部署的程度,并以量化感知训练、自定义 TFLite 算子、sigmoid/softmax 导出参数等方式打通了从研究训练到移动端推理的完整链路。阅读本文后,你可以直接复用 configs/ 下的两份配置运行 PRADO 训练与评估,也能借助 models/prado.pymodels/sgnn/train.py 理解其网络结构与超参语义,进而把新任务快速接入这一“小而精”的文本建模框架。

该项目基于 Apache License 2.0 开源(见 research/seq_flow_lite/README.md 末尾声明),读者可在此基础上自由扩展自定义层与任务。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.75 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
857
1.35 K
docsdocs
暂无描述
Markdown
899
5.82 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
920
1.85 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.8 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
532
596
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.02 K
521
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.36 K
1.46 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
392