首页
/ TensorFlow models 仓库 TF-NLP 数据管道指南:离线预处理、TFDS 读取与 DataLoader 工厂机制

TensorFlow models 仓库 TF-NLP 数据管道指南:离线预处理、TFDS 读取与 DataLoader 工厂机制

2026-09-06 12:42:37作者:韦蓉瑛

本篇基于 official/nlp/docs/data_processing.md 展开,讲解 TF-NLP 中训练数据从原始文本到 tf.data.Dataset 的完整链路:离线预处理脚本如何生成 tf.Example 文件、DataConfiginput_path 与 TFDS 字段如何被 InputReader 解析、DataLoader 抽象与 data loader 工厂 如何按配置类注册和分发,以及在 TPU 上用 tf.data service + 动态序列长度实现训练加速的工程要点。读完后你可以独立地为新任务编写一个 DataLoader、用 TFDS 或离线 TFRecord 喂入任意 NLP 任务,并理解 TPU 环境下输入管道的约束。

1. 两条数据供给路径:离线预处理 vs TFDS

TF-NLP 提供训练数据给输入管道有两种方式(见 数据文档):

  1. 离线预处理:用 Python 脚本(也可换成 Beam/Flume 等)在训练前把文本 tokenize 成包含 tf.Example proto 的 TFRecord 文件,训练时通过 DataConfig.input_path 指定这些文件;
  2. 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_pathtfds_name 同时指定会直接抛 ValueError

1.1 离线预处理脚本及其关键参数

仓库提供了两个数据预处理脚本:

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_pathinput_path 支持 glob 通配符,这在 match_files 中实现:含 */? 的模式会调用 tf.io.gfile.glob 展开(也兼容 GCS 路径),匹配不到文件会直接报错。

1.2 TFDS 路径:InputReader 的统一读取

为方便收敛,仓库构建了通用的 InputReader 库来标准化输入读取,内置 TFDS 支持。只要在 DataConfig 中指定 tfds_nametfds_data_dirtfds_splittf.data 管道就会从 TFDS 对应数据集读取。

从源码看,InputReader.read 的执行链是:

  1. 数据源读取_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 并发读取)。
  2. 解码/解析/批处理_decode_and_parse_dataset):依次执行 shuffle(仅训练且未 cache 时)→ decoder_fn → 可选 combine_fn/sample_fnparser_fn → 可选 filter_fncache(如启用)→ transform_and_batch_fn 或默认按 input_context.get_per_replica_batch_size 做 per-replica 批处理。
  3. 可选 tf.data service 分发_maybe_apply_data_service)与最终 prefetch

InputReader 的构造参数中还有一组可插拔回调:dataset_fn(默认 tf.data.TFRecordDataset)、decoder_fncombine_fnsample_fnparser_fnfilter_fntransform_and_batch_fnpostprocess_fn,各 Data 加载器就是靠替换这些回调实现不同任务的数据语义。

另外注意 InputReader 中对 tf.data service 的处理:启用 tf.data service 时会把 seed 置为 Nonesharding 强制关闭——因为分片由 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_loaderdata_config.__class__ 完成;测试用例见 data_loader_factory_test.py。仓库中已注册的 DataLoader 覆盖主要任务,例如:

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_datavalidation_data,类型即 DataConfig。仓库中 masked_lm.pysentence_prediction.pytagging.pyquestion_answering.pydual_encoder.pyelectra_task.pytranslation.py 等均遵循这一约定。

TPU 约束(文档重点提示):在 TPU 训练时,整个 load 方法会运行在 TPU worker 上,因此函数内不能访问外部资源,例如 task 实例属性——所有状态必须从 paramsDataConfig)或 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% 的训练加速

启用动态序列长度的做法:

  1. 启动实验时加 --enable_tf_data_service 参数:动态序列需要 tf.data service 对全序列做全局 bucketizing(分桶),因此必须启用 tf.data service 承担这一角色;DataConfig 对应字段为 enable_tf_data_servicetf_data_service_address(URI 形如 grpc://host:port,可被 binary 的 FLAGS.tf_data_service 覆盖,见 config_definitions.py)。
  2. 使用实现了 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_pathtfds_name 三件套)、批处理与 shuffle 参数、tf.data service 开关,且两种数据源互斥;
  • 读取层InputReader 把文件 glob 展开、TFDS 加载、多 worker sharding、解码/解析/批处理以及 tf.data service 分发标准化,并通过回调实现可插拔;
  • 扩展层DataLoader.load 抽象 + 工厂注册 让每个任务只关心自己的 DataConfig 子类,build_inputstrain_data/validation_data 字段自动映射到对应加载器。

在这套机制上,离线脚本产出 TFRecord 适合大规模预训练语料(配合 --dupe_factor--max_ngram_size 等控制掩码策略),TFDS + TF.Text 适合直接消费公开文本数据集;而在 TPU 大规模训练中,--enable_tf_data_service 加动态序列 DataLoader 是文档给出的官方加速路径。

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