Skip to content

Apache MXNet - 分布式训练

本章介绍了使用 Apache MXNet 进行分布式训练,它能够利用多个计算资源(跨一台或多台机器的 CPU/GPU)来加速和扩展模型训练。

MXNet 支持两种主要的编程范式:

命令式模式(Imperative Mode)(Gluon API)

Section titled “命令式模式(Imperative Mode)(Gluon API)”

这种模式主要通过 Gluon API 暴露,允许逐步定义和执行计算,非常类似于标准的 Python 代码或 NumPy。操作通常立即执行。

import mxnet as mx
from mxnet import nd
# 在 CPU 和 GPU(如果可用)上创建张量
tensor_cpu = nd.zeros((100,), ctx=mx.cpu())
# tensor_gpu = nd.zeros((100,), ctx=mx.gpu(0)) # 如果存在 GPU 则取消注释
# print(tensor_cpu)
# if 'tensor_gpu' in locals(): print(tensor_gpu)

虽然是命令式的,但 MXNet 的后端执行引擎仍然可以优化和并行化操作,通常会延迟实际计算直到需要结果(某些操作的惰性求值 Lazy Evaluation)。

在符号式模式下,您首先使用 mx.symbol 对象定义计算图。这个图代表了操作及其依赖关系,但不会立即执行。然后将图编译并与数据绑定以高效执行。

import mxnet as mx
from mxnet import sym
x_sym = sym.Variable("data_x")
y_sym = sym.Variable("data_y")
z_sym = (x_sym + y_sym) / 100
# 要执行,此符号需要使用 Module 或 Executor 与数据绑定
# executor = z_sym.simple_bind(ctx=mx.cpu(), data_x=(10,1), data_y=(10,1))
# 对于 Gluon 用户,符号式编程通常被 HybridBlocks 抽象化。

Gluon 的 HybridBlock 允许模型以命令式定义,然后可选地转换(hybridized)为符号图以获得性能优势,将命令式编程的灵活性与符号式执行的速度融合。

MXNet 支持不同的策略来分配训练神经网络的工作负载:

这是最常见的策略。模型在多个设备(例如 GPU)上复制。每个设备处理训练数据的不同 Mini-Batch(Mini-Batch)。每个设备上计算的梯度随后被聚合(例如,求平均),并且模型参数会同步或异步更新。这可以在具有多个 GPU 的单台机器上完成,也可以跨多台机器完成。

当模型太大无法放入单个设备的内存时使用。模型的不同部分放置在不同的设备上。数据按顺序或以更复杂的流水线(Pipeline)流经这些部分。MXNet 对模型并行的支持通常侧重于单机场景,尽管复杂的设置可以扩展此功能。

理解以下组件是掌握 MXNet 分布式训练架构的关键:

在典型的 MXNet 分布式设置中,进程扮演着特定的角色:

  • **Worker(工作节点):**执行实际的训练计算。每个 Worker 通常处理一部分数据,计算梯度,将其发送到参数服务器,并接收更新后的参数。
  • **Server(服务器,即参数服务器 Parameter Server):**存储和管理模型的参数。它接收来自 Worker 的梯度,聚合它们,更新参数(或将梯度提供给 Worker 进行本地更新),并将更新后的参数发送回 Worker。可以使用多个 Server 来分散参数存储和更新的负载。
  • **Scheduler(调度器):**管理集群的设置和协调。它确保所有节点(Worker 和 Server)都能找到并相互通信。集群中通常只有一个 Scheduler。

KVStore 是分布式训练中通信的主干,特别是对于数据并行。它管理参数的存储和同步(键是参数名称/索引,值是参数的 NDArray)。

  • Worker 向 KVStore push 梯度。
  • Worker 从 KVStore pull 更新后的参数。
  • KVStore 处理梯度的聚合和参数的更新,可以在服务器端进行,也可以提供组件供客户端更新。

要启用分布式训练,需要创建一个分布式 KVStore:

# 示例:创建一个分布式同步 KVStore
# 这段代码将作为由 tools/launch.py 或类似工具启动的脚本的一部分运行
# kv_type = 'dist_sync' # 或 'dist_async' 等
# kv = mx.kv.create(kv_type)
# print(f"创建了类型为: {kv_type} 的 KVStore")
print("要使用分布式 KVStore,例如 mx.kv.create('dist_sync'),需要像 DMLC_ROLE 这样的环境变量。")

在多服务器设置中,参数(键)分布在各个服务器上。KVStore 客户端库透明地将针对特定键的请求路由到正确的服务器。大型参数数组甚至可以分片(sharded)到多个服务器上。

为了有效的数据并行,每个 Worker 应处理整个训练数据集中的一个独特子集。MXNet 的数据迭代器,如 mxnet.io.MNISTIterator 和 mxnet.gluon.data.DataLoader,可以配置用于分布式训练,确保每个 Worker 获得其指定的分片。数据加载工具中的 num_workers 和 rank(或类似)参数有助于正确划分数据。

在单个 Worker 内部(其本身可能拥有多个 GPU),mxnet.gluon.utils.split_and_load 可以进一步分割一个 mini-batch,以便在本地设备上处理。

当与分布式 KVStore 一起使用时,gluon.Trainer 对象管理参数更新。Trainer 中的 update_on_kvstore 参数至关重要:

  • 如果 update_on_kvstore=True(或未设置,对于分布式 KVStore 通常默认为 True),则优化器的更新逻辑在参数服务器上执行。Worker push 梯度,服务器更新参数,Worker pull 更新后的参数。
  • 如果 update_on_kvstore=False,Worker push 梯度,服务器聚合它们并将聚合后的梯度发回。然后 Worker 使用这些聚合后的梯度执行本地参数更新。这有时可以减少通信开销或允许更复杂的本地更新规则。
from mxnet import gluon
# net = gluon.nn.Sequential() # ... 定义你的网络 ...
# net.initialize(ctx=contexts) # contexts 可以是 [mx.gpu(0), mx.gpu(1)]
# kv = mx.kv.create('dist_sync') # 假设这是一个分布式 Worker 节点
# trainer = gluon.Trainer(
# net.collect_params(),
# optimizer='sgd',
# optimizer_params={'learning_rate': 0.01},
# kvstore=kv,
# update_on_kvstore=True # 或 False,取决于期望的策略
# )
print("gluon.Trainer 的 update_on_kvstore 参数控制参数更新发生的位置。")

mx.kv.create() 的字符串参数决定了分布式 KVStore 的类型和行为:

所有 Worker 同步操作。处理完一个 Batch 后,每个 Worker push 其梯度。服务器等待接收所有(或配置数量的)Worker 的梯度,然后聚合它们并更新参数。所有 Worker 然后 pull 相同的、新更新的参数,然后才开始下一个 Batch。

优点:更容易理解,通常能带来更稳定的收敛(Convergence)。缺点:如果某些 Worker 明显比其他 Worker 慢(落后者问题 Straggler Problem),可能会很慢。一个崩溃的 Worker 可能导致进度停止。

Worker 异步操作。当一个 Worker 完成一个 Batch 并 push 梯度时,服务器会立即更新其参数,无需等待其他 Worker。然后该 Worker pull 最新可用的参数,这些参数可能已经使用了更早或更晚完成的其他 Worker 的梯度进行了更新。

优点:由于较快的 Worker 无需等待较慢的 Worker,可以实现更高的吞吐量(Throughput)。对 Worker 故障更具弹性。缺点:梯度可能 ‘过时’(Stale)(基于较旧的参数版本计算),可能导致收敛不稳定或需要仔细调整学习率(Learning Rates)。

dist_device_sync(或类似名称如 dist_sync_device)

Section titled “dist_device_sync(或类似名称如 dist_sync_device)”

类似于 dist_sync,但针对每个 Worker 节点拥有多个 GPU 的场景进行了优化。梯度聚合和参数更新尝试直接在设备内存(GPU)上执行,以减少服务器端(如果适用)或 Worker 侧本地多 GPU 聚合之前的 CPU-GPU 数据传输开销。

dist_device_async(或类似名称如 dist_async_device)

Section titled “dist_device_async(或类似名称如 dist_async_device)”

结合了 dist_async 的异步特性和 dist_device_sync 的设备内存优化尝试。

选择合适的分布式模式取决于具体的硬件设置、网络条件、模型架构以及在训练速度、收敛稳定性和实现复杂性之间期望的权衡。对于启动分布式 MXNet 任务,通常使用 tools/launch.py(来自 MXNet 源代码仓库)或与集群管理器(例如 Kubernetes、YARN)的集成工具来设置必要的环境变量(DMLC_ROLE、DMLC_PS_ROOT_URI 等),以配置 Scheduler、Server 和 Worker 进程。