从 Horovod 到 tf.distribute:TensorFlow 分布式训练十年演进,大厂为何还在用它训推荐模型
【免费下载链接】tensorflowAn Open Source Machine Learning Framework for Everyone项目地址: https://gitcode.com/GitHub_Trending/te/tensorflow
2017 年,Uber 开源了 Horovod,用一句"只需几行代码就能把 TensorFlow 训练扩展到多机多卡"征服了大量工程团队;同年,TensorFlow 内部开始酝酿一套更彻底的方案——tf.distribute。十年过去,深度学习框架的版图几经更迭,PyTorch 在学术界攻城略地,JAX 在前沿研究里风生水起,但翻开任何一家互联网大厂的推荐与广告系统架构文档,TensorFlow 依然稳定地占据着训练链路的主位。本文不打算再复述一遍"框架之争",而是沿着分布式训练这条主线,拆解从 Horovod 到tf.distribute的演进逻辑,并回到仓库源码里,看看大厂为什么至今仍在用它训练推荐模型。
一、推荐模型的训练,为什么绕不开分布式
推荐与广告模型是深度学习里最"工业化"的一类负载。以典型的 Deep & Cross、DIN、Two-Tower 为代表的模型,结构并不复杂,真正压垮单机的是三样东西:海量的稀疏特征、巨大的 embedding 表,以及天量的样本吞吐。
一张用户行为序列的 embedding 表动辄数十亿行、数百 GB 乃至 TB 级,远超单卡显存;训练样本以 TFRecord 形式分布在成百上千个文件中,一天的数据就是几百 TB。要让模型在"每天全量重训一轮、小时级增量更新"的节奏下跑起来,数据并行是唯一现实的选择——把样本按批次切给多台机器,每台机器算梯度,再统一同步参数。
问题的关键从来不是"要不要分布式",而是"分布式的粒度放在哪一层"。这恰恰是 Horovod 与tf.distribute分道扬镳的起点。
二、Horovod 的贡献与局限:把"梯度同步"做到极致
Horovod 的核心思想非常聚焦:它不关心模型如何构建,只接管梯度的跨机通信。基于 MPI 的 AllReduce 原语,配合 ring-allreduce 算法,让梯度在节点间以环形拓扑高效聚合。对于以"数据并行 + 同步更新"为主、模型是规整稠密网络的场景(比如经典的 CNN/ResNet 图像任务),这是近乎完美的解:不改模型结构、不引入新架构,hvd.DistributedOptimizer包一层就能跑多机。
但推荐模型恰恰戳中了 Horovod 的两个盲区。
第一,它表达不了参数服务器语义。稀疏特征场景下,工程上普遍采用 PS(Parameter Server)架构:worker 负责算梯度,专门的 ps 节点持有 embedding 参数,worker 与 ps 之间以 pull/push 方式异步交互。Horovod 的 AllReduce 模型假设"所有节点持有完整参数的副本、每步全量同步",与异步 PS、稀疏参数分片天然不匹配。
第二,它把问题留给了上层应用。Horovod 只做梯度同步,数据切分、分片、容错、监控、checkpoint 编排都得团队自己搭。当训练规模从"8 卡"膨胀到"上千 worker",这些"边缘问题"会变成主战场。
Horovod 的价值在于证明了"把分布式做成库、而非让用户重写训练代码"这条路是通的。而 TensorFlow 选择走得更远:与其在框架外面包一层通信库,不如把分布式能力直接内建到运行时里。
三、tf.distribute:一套 API,五种策略
在 tensorflow/python/distribute/README.md 里,官方对这套 API 的定位写得很直白:
tf.distribute.Strategyis a TensorFlow API to distribute training across multiple GPUs, multiple machines or TPUs. Using this API, users can distribute their existing models and training code with minimal code changes.
关键句是 "minimal code changes"。它通过让 TensorFlow 底层组件(变量、层、模型、优化器、指标、summary、checkpoint)变得strategy-aware,把"分布式"从用户代码中抽离出去——这正是 Horovod 想做而没做全的事。
tf.distribute覆盖了五种分布式形态,分别对应不同硬件与一致性模型:
| 策略 | 定位 | 同步/异步 | 适用场景 |
|---|---|---|---|
MirroredStrategy | 单机多卡,变量镜像到每张卡 | 同步 | 单机多 GPU 快速迭代 |
MultiWorkerMirroredStrategy | 多机多卡,collective 通信 | 同步 | 多机稠密模型数据并行 |
TPUStrategy | TPU 集群 | 同步 | TPU 大规模训练 |
ParameterServerStrategy | worker + ps 集群 | 异步为主 | 超大规模稀疏模型 |
OneDeviceStrategy | 单设备兜底 | — | 调试、单卡 |
用户侧的编程模型高度统一:在strategy.scope()下构建模型,用strategy.run()执行训练步,用strategy.reduce()聚合跨副本结果。仓库里 mirrored_strategy.py 的 docstring 给出了完整示例:
my_strategy = tf.distribute.MirroredStrategy() with my_strategy.scope(): @tf.function def distribute_train_epoch(dataset): def replica_fn(input): # process input and return result return result total_result = 0 for x in dataset: per_replica_result = my_strategy.run(replica_fn, args=(x,)) total_result += my_strategy.reduce(tf.distribute.ReduceOp.SUM, per_replica_result, axis=None) return total_result同样是"包一层",Horovod 包的是优化器,tf.distribute包的是整个训练上下文。
四、从"梯度搬运工"到"全栈调度器":通信原语与数据管线
tf.distribute与 Horovod 的第二个本质差异,在于它把通信与数据都收编进了框架内部。
通信:NCCL/RING 的自动降级
多机同步训练依赖 all-reduce,但tf.distribute没有把通信实现焊死在一种算法上。在 collective_util.py 中,CommunicationImplementation枚举了三种实现:AUTO(自动选择)、RING(TensorFlow 自研 ring 算法)、NCCL(NVIDIA 集体通信库,GPU all-reduce 首选)。而MirroredStrategy初始化时并非一刀切:在 mirrored_strategy.py 的_make_collective_ops_with_fallbacks里,逻辑会根据设备组成自动降级——纯 CPU 环境退化为 RING 实现,混合 CPU/GPU 环境退化为ReductionToOneDevice,GPU 环境默认NcclAllReduce。
这套"自动降级"策略意味着:同一份用户代码,从开发机的 CPU 单机,到测试机的多卡,再到生产环境的多机集群,可以无缝迁移,通信细节由框架按拓扑自动适配。
数据:AutoShard 与输入管线
梯度同步只是分布式的一半,另一半是数据分发。推荐模型的数据管线极其挑剔:样本文件多、单文件大、需要按 worker 切分又不能破坏随机性。tf.distribute的处理在 tensorflow/python/data/experimental/ops/distribute.py 的_AutoShardDataset中实现:
FILE:向上遍历数据集图,找到 reader 节点,在其前插入ShardDataset,让每个 worker 只读一部分文件;DATA:在输入管线末端插入ShardDataset,按数据元素切分;AUTO:优先尝试文件切分,找不到 reader 时回退到数据切分;HINT:配合tf.data.experimental.SHARD_HINT占位符,由用户显式指定切分点。
对于 embedding 特征这种"必须在不同 worker 看到相同样本全集"的场景,框架还提供了distribute_datasets_from_function,让用户接管 batching 与 sharding。在 distribute_lib.py 的 docstring 中写得很清楚:它接收一个以tf.distribute.InputContext为参数的dataset_fn,返回的 dataset 按per-replica batch size构建——这为稀疏特征场景下"每个 worker 独立构造自己的训练流"留出了精确的控制面。
五、推荐场景的存量依赖:为什么大厂还在用 TF
前面几节解释了"技术上行得通",这一节回答"为什么推荐模型这个具体场景里,TF 至今是存量主力"。核心答案藏在ParameterServerStrategyV2 的架构里。
从"裸 PS"到"中央协调 + 分片变量"
老一代 TF1 的 PS 方案要求用户手写tf.train.Server、job 名、replica_device_setter,容错和调度全靠自己。tf.distribute.experimental.ParameterServerStrategy(V2)则重构为中央协调架构:集群由 worker、ps 和 coordinator 三类角色构成,coordinator 负责创建资源、分发tf.function、保存 checkpoint。在 parameter_server_strategy_v2.py 的类文档里,对这套设计的描述直指推荐训练的核心诉求:
As a result, failures of some workers do not prevent the cluster from continuing the work, and this allows the cluster to train with instances that can be occasionally unavailable (e.g. preemptible or spot instances).
——部分 worker 故障不影响训练继续,因而可以放心使用抢占式/spot 实例降低成本。这正是推荐训练规模化后最现实的工程问题:上千个 worker 的集群,节点故障是常态而非异常,异步 PS + 中央协调让"死一个节点"从事故降级为日常。
embedding 分片:MinSizePartitioner
推荐模型动辄 TB 级的 embedding 表必须切到多张 ps 上。V2 策略通过variable_partitioner参数支持变量自动分片,构造函数的 docstring 给出了推荐配置:
MinSizePartitioner(min_shard_bytes=256 << 10, max_shards=num_ps)每个分片至少 256KB、每台 ps 至多分到一个分片,从而保证 embedding 表在 ps 间均匀分布。底层由 sharded_variable.py 的ShardedVariable承接:在strategy.scope()下创建的tf.Variable会被包装成"分片容器",对用户透明。
调度器:ClusterCoordinator 的异步分发
V2 的 worker 不再各自为战,而是由ClusterCoordinator统一调度。在 cluster_coordinator.py 中,schedule()是异步非阻塞的:把tf.function排队分发给可用 worker,立即返回RemoteValue;join()阻塞等待全部完成。更关键的是容错语义——docstring 里明确承诺:
scheduleguarantees thatfnwill be executed on a worker at least once; it could be more than once if its corresponding worker fails in the middle of its execution.
"至少执行一次"的语义配合异步训练,使得 PS 集群天然容忍 worker 级别故障,这是大厂在生产推荐模型上最看重的能力之一。
存量生态:训练只是半条命
最后必须承认一个事实层面的原因:大厂推荐系统的训练只是链路的一环,前面是特征工程与样本生成(TFRecord + tf.data),后面是模型导出与线上推理(TensorFlow Serving)。十年前基于 TF1 构建的整套特征体系、样本规范、模型格式和运维工具,构成了巨大的迁移成本。tf.distribute的价值正在于让这套存量资产在 TF2 时代继续运转——同一套 Keras 模型代码,改一行策略类即可从单机切到多机、从同步切到异步 PS,而不必重写特征管线与部署链路。
结语
回看这十年,Horovod 用 AllReduce 回答了"如何高效同步梯度",tf.distribute则用一套策略抽象回答了更完整的问题:"如何让分布式成为框架的内建能力"。当推荐模型的规模把工程复杂度推向极致时,后者的设计取向——内建通信原语、内建数据分片、内建 PS 语义与容错调度——恰好命中了所有要害。
框架之争的舆论场里,TensorFlow 未必是"最流行"的叙事主角,但在推荐与广告这类最吃工程、最吃存量、最吃稳定性的工业化场景里,"还在用它"本身就是对这套分布式设计十年演进最有力的投票。
【免费下载链接】tensorflowAn Open Source Machine Learning Framework for Everyone项目地址: https://gitcode.com/GitHub_Trending/te/tensorflow
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考