TensorFlow models 仓库 TF-NLP 数据管道指南:离线预处理、TFDS 读取与 DataLoader 工厂机制
本篇基于 official/nlp/docs/data_processing.md 展开,讲解 TF-NLP 中训练数据从原始文本到 tf.data.Dataset 的完整链路:离线预处理脚本如何生成 tf.Example 文件、DataConfig 的 input_path 与 TFDS 字段如何被 InputReader 解析、DataLoader 抽象与 data loader 工厂 如何按配置类注册和分发,以及在 TPU 上用 tf.data service + 动态序列长度实现训练加速的工程要点。读完后你可以独立地为新任务编写一个 DataLoader、用 TFDS 或离线 TFRecord 喂入任意 NLP 任务,并理解 TPU 环境下输入管道的约束。
1. 两条数据供给路径:离线预处理 vs TFDS
TF-NLP 提供训练数据给输入管道有两种方式(见 数据文档):
- 离线预处理:用 Python 脚本(也可换成 Beam/Flume 等)在训练前把文本 tokenize 成包含
tf.Exampleproto 的 TFRecord 文件,训练时通过DataConfig.input_path指定这些文件; - TFDS 在线读取:直接从 TensorFlow Datasets 加载文本数据,并用 TF.Text 在
tf.data管道内完成 tokenize 和预处理。
两种路径都收敛到同一个配置类 DataConfig,它同时约束了两种互斥的数据源字段:
| 字段 | 默认值 | 含义 |
|---|---|---|
input_path |
"" |
离线文件路径,支持单个路径、逗号分隔的多路径/通配符、路径列表,以及“字典”形式做命名数据混合;与 tfds_name 二选一 |
tfds_name |
"" |
TFDS 数据集名(可为 Config 字典,支持多数据集);指定时必填 tfds_split |
tfds_split |
"" |
从 TFDS 加载哪个 split(train/test 等) |
tfds_data_dir |
"" |
TFDS 数据读写目录 |
tfds_as_supervised |
False |
True 时返回 (input, label) 二元组,False 时返回含全部特征的字典 |
tfds_skip_decoding_feature |
"" |
逗号分隔的跳过解码特征列表,常用于跳过图像/视频解码 |
global_batch_size / drop_remainder / shuffle_buffer_size |
0 / True / 100 |
全局批大小、是否丢弃尾部残批、训练时 shuffle 缓冲大小 |
sharding |
True |
输入管道是否启用 sharding |
enable_tf_data_service |
False |
是否启用 tf.data service(配合 tf_data_service_address) |
互斥关系在 InputReader 构造函数 中强制校验:input_path 与 tfds_name 同时指定会直接抛 ValueError。
1.1 离线预处理脚本及其关键参数
仓库提供了两个数据预处理脚本:
- create_pretraining_data.py:为 BERT 预训练生成 masked LM / next sentence 训练样本;
- create_finetuning_data.py:为微调任务(如 GLUE)生成训练样本。
以 create_pretraining_data.py 为例,它定义了一套完整的 absl flags,可直接复用到自己的语料上:
| Flag | 默认值 | 说明 |
|---|---|---|
--input_file |
必填 | 原始文本文件(或逗号分隔的多个文件) |
--output_file |
必填 | 输出的 TF example 文件(或逗号分隔列表) |
--tokenization |
WordPiece |
WordPiece(BERT 默认)或 SentencePiece(ALBERT 默认) |
--vocab_file |
无 | WordPiece 词表文件 |
--sp_model_file |
"" |
SentencePiece 模型文件路径 |
--do_lower_case |
True |
是否转小写,uncased 模型应为 True |
--do_whole_word_mask |
False |
整词掩码 |
--max_ngram_size |
None |
掩码最长连续 n-gram(需配合 --do_whole_word_mask=True) |
--gzip_compress |
False |
输出压缩 TFRecord |
--use_v2_feature_names |
False |
使用与模型一致的 v2 特征名 |
--max_seq_length |
128 |
最大序列长度 |
--max_predictions_per_seq |
20 |
每条序列的 masked LM 预测上限 |
--random_seed |
12345 |
数据生成随机种子 |
--dupe_factor |
10 |
输入数据复制份数(每份用不同掩码),扩充有效样本 |
--masked_lm_prob |
0.15 |
masked LM 掩码概率 |
--short_seq_prob |
0.1 |
生成短于最大长度序列的概率 |
运行生成的文件内部是 tf.Example proto,之后只需把这些文件路径写进实验配置的 input_path。input_path 支持 glob 通配符,这在 match_files 中实现:含 */? 的模式会调用 tf.io.gfile.glob 展开(也兼容 GCS 路径),匹配不到文件会直接报错。
1.2 TFDS 路径:InputReader 的统一读取
为方便收敛,仓库构建了通用的 InputReader 库来标准化输入读取,内置 TFDS 支持。只要在 DataConfig 中指定 tfds_name、tfds_data_dir、tfds_split,tf.data 管道就会从 TFDS 对应数据集读取。
从源码看,InputReader.read 的执行链是:
- 数据源读取(
_read_data_source):- 有
tfds_name时走_read_tfds:对mldataset.*前缀的名字用tfds.load直接加载,否则取tfds.builder检查 split 分片数;当分片数少于输入管道数时,会先在主机内存读完整个数据集再手动shard,保证每个 worker 都拿到数据; - 有
input_path时按文件数选择两种策略:文件少于 worker 数时用_read_files_then_shard(全量下发再按数据分片,避免数据浪费),否则用_shard_files_then_read(先在文件级 shuffle+shard 再interleave并发读取)。
- 有
- 解码/解析/批处理(
_decode_and_parse_dataset):依次执行shuffle(仅训练且未 cache 时)→decoder_fn→ 可选combine_fn/sample_fn→parser_fn→ 可选filter_fn→cache(如启用)→transform_and_batch_fn或默认按input_context.get_per_replica_batch_size做 per-replica 批处理。 - 可选 tf.data service 分发(
_maybe_apply_data_service)与最终prefetch。
InputReader 的构造参数中还有一组可插拔回调:dataset_fn(默认 tf.data.TFRecordDataset)、decoder_fn、combine_fn、sample_fn、parser_fn、filter_fn、transform_and_batch_fn、postprocess_fn,各 Data 加载器就是靠替换这些回调实现不同任务的数据语义。
另外注意 InputReader 中对 tf.data service 的处理:启用 tf.data service 时会把 seed 置为 None、sharding 强制关闭——因为分片由 tf.data service 的 processing_mode 统一处理;tf.data service 的 job name 中还会拼入 global_batch_size 和随机数,防止 TPU worker 被抢占后复用旧状态,或调参时因批大小变化读到维度不匹配的 tensor。
2. DataLoader 抽象与工厂:多数据集、多处理函数的管理
为了统一管理多个数据集和处理函数,TF-NLP 定义了 DataLoader 抽象类配合 data loader factory 使用。每个 DataLoader 在 load 方法内构建完整的 tf.data 输入管道:
@abc.abstractmethod
def load(
self,
input_context: Optional[tf.distribute.InputContext] = None
) -> tf.data.Dataset:
input_context 由 distribution strategy 传入,包含计算副本和输入管道信息(多机输入场景);注意 load 返回的是 per-host 数据集,分布式数据集由外层 trainer 负责包装。
2.1 工厂注册机制
data_loader_factory.py 维护一个全局注册表 _REGISTERED_DATA_LOADER_CLS,以 DataConfig 子类为键。典型用法:
@dataclasses.dataclass
class MyDataConfig(DataConfig):
# 添加字段
pass
@register_data_loader_cls(MyDataConfig)
class MyDataLoader: # 继承 def __init__(self, data_config)
pass
my_config = MyDataConfig()
my_loader = get_data_loader(my_config) # 返回 MyDataLoader(my_config)
注册本质是委托给 registry.register,查找通过 get_data_loader 按 data_config.__class__ 完成;测试用例见 data_loader_factory_test.py。仓库中已注册的 DataLoader 覆盖主要任务,例如:
- pretrain_dataloader.py:
BertPretrainDataConfig/XLNetPretrainDataConfig; - pretrain_dynamic_dataloader.py:TPU 动态序列预训练;
- pretrain_text_dataloader.py:原始文本预训练(TF.Text tokenize);
- sentence_prediction_dataloader.py:GLUE 句对任务(含
SentencePredictionTextDataConfig原始文本版); - tagging_dataloader.py、question_answering_dataloader.py、dual_encoder_dataloader.py、wmt_dataloader.py 等。
2.2 与任务(Task)的衔接:build_inputs
load 方法在各 NLP 任务的 build_input(s) 方法中被调用,trainer 再将其包装成分布式数据集。文档给出的模式为:
def build_inputs(self, params, input_context=None):
"""Returns tf.data.Dataset for pretraining."""
data_loader = YourDataLoader(params)
return data_loader.load(input_context)
其中 params 默认是实验配置中 task 字段的 train_data 或 validation_data,类型即 DataConfig。仓库中 masked_lm.py、sentence_prediction.py、tagging.py、question_answering.py、dual_encoder.py、electra_task.py、translation.py 等均遵循这一约定。
TPU 约束(文档重点提示):在 TPU 训练时,整个 load 方法会运行在 TPU worker 上,因此函数内不能访问外部资源,例如 task 实例属性——所有状态必须从 params(DataConfig)或 DataLoader 自身的 __init__ 阶段传入。
2.3 处理原始文本特征(TF.Text)
要处理原始文本特征,需要使用基于 TF.Text 做 tokenize 的 DataLoader。参考实现是 sentence_prediction_dataloader.py,它展示了如何从 TFDS 读取原始文本特征完成 BERT GLUE 微调。
3. 在 TPU 上用 tf.data service + 动态序列长度加速训练
TF 2.x 的编程模型配合 TPUStrategy/XLA,允许在 TPU 上启用某些动态 shape。根据 文档,依赖数据分布,在 BERT 预训练等典型文本场景下,相对 padded 静态 shape 输入可看到 50% 到 90% 的训练加速。
启用动态序列长度的做法:
- 启动实验时加
--enable_tf_data_service参数:动态序列需要 tf.data service 对全序列做全局 bucketizing(分桶),因此必须启用 tf.data service 承担这一角色;DataConfig对应字段为enable_tf_data_service与tf_data_service_address(URI 形如grpc://host:port,可被 binary 的FLAGS.tf_data_service覆盖,见 config_definitions.py)。 - 使用实现了 bucketizing 的 DataLoader:参考 pretrain_dynamic_dataloader.py,它针对 tokenize 后的数据集做 BERT 预训练,并对不同长度序列做分桶批处理,从而让 XLA 动态 shape 路径生效。
与 InputReader 的联动关系是:一旦 enable_tf_data_service 为真,本地 sharding 自动关闭、随机 seed 交由 tf.data service 侧的每个 data service worker 使用不同种子(见 input_reader.py 相关逻辑),保证全序列维度的分桶结果全局一致,这是动态 shape 能正确运行的前提。
4. 小结
TF-NLP 的数据管道设计可归纳为三层:
- 配置层:
DataConfig以 dataclass 统一描述数据来源(input_path或tfds_name三件套)、批处理与 shuffle 参数、tf.data service 开关,且两种数据源互斥; - 读取层:InputReader 把文件 glob 展开、TFDS 加载、多 worker sharding、解码/解析/批处理以及 tf.data service 分发标准化,并通过回调实现可插拔;
- 扩展层:
DataLoader.load抽象 + 工厂注册 让每个任务只关心自己的 DataConfig 子类,build_inputs将train_data/validation_data字段自动映射到对应加载器。
在这套机制上,离线脚本产出 TFRecord 适合大规模预训练语料(配合 --dupe_factor、--max_ngram_size 等控制掩码策略),TFDS + TF.Text 适合直接消费公开文本数据集;而在 TPU 大规模训练中,--enable_tf_data_service 加动态序列 DataLoader 是文档给出的官方加速路径。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0624
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00