MLX 分布式通信完全指南:四种后端、启动工具与实战配置
2026/9/11 2:32:13 网站建设 项目流程

MLX 分布式通信完全指南:四种后端、启动工具与实战配置

【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx

MLX 是面向 Apple Silicon 的数组框架,本文档以 docs/src/usage/distributed.rst 为主体,系统讲解 MLX 的分布式通信能力:如何用几行代码把训练或推理的计算成本分摊到多台物理机器上,如何用mlx.launchmlx.distributed_config两个辅助工具启动与配置多节点任务,以及 MPI、RING、JACCL、NCCL 四种通信后端各自的适用场景、hostfile 格式与环境变量。读完本文,你将能独立搭建从单机多进程到多机多 GPU(含 Thunderbolt RDMA 全互联 mesh)的分布式 MLX 程序。

分布式通信概述与 API 骨架

MLX 的核心分布式抽象是一个Group(通信组),它把一组参与通信的进程组织起来,并提供统一的集体通信(collective communication)原语。C++ 侧的定义位于 mlx/distributed/distributed.h,Python 侧通过mlx.core.distributed模块暴露。每个 Group 至少提供以下能力:

  • rank():当前进程在组内的编号(从 0 开始);
  • size():组内进程总数;
  • 集体通信原语:all_sumall_maxall_minall_gathersum_scattersendrecv(对应声明见 mlx/distributed/ops.h)。

一个分布式 MLX 程序最短只需三行:

import mlx.core as mx world = mx.distributed.init() x = mx.distributed.all_sum(mx.ones(10)) print(world.rank(), x)

这段代码在所有分布式进程间对mx.ones(10)求和。关键设计:当只用python直接运行脚本时,只启动一个进程,通信组大小为 1,此时mx.distributed下的所有操作都是 noop(空操作)。这一特性让你免去如下所示的手动判断:

import mlx.core as mx x = ... world = mx.distributed.init() # 无需这个判断,直接写 x = mx.distributed.all_sum(x) 即可 if world.size() > 1: x = mx.distributed.all_sum(x)

从源码层面看,单进程时init()会返回一个EmptyGroup(见 mlx/distributed/distributed.cpp 中的detail::EmptyGroup),其size()恒为 1、rank()恒为 0;不过它内部的通信方法(如all_sum)会抛出"Communication not implemented in an empty distributed group."异常——也就是说,只要进程真正加入了分布式组,操作就会真实执行。

支持的通信后端一览

MLX 当前支持四种后端,分别面向不同硬件与网络环境:

后端说明
MPI功能完善、成熟的分布式通信库,适合已有 MPI 生态的集群
RING基于 TCP socket 的环形 all reduce 与 all gather,始终可用(不依赖第三方库),通常比 MPI 更快
JACCL基于 Thunderbolt 上的 RDMA 实现低延迟通信,是张量并行(tensor parallelism)等场景的必需项
NCCLCUDA 环境下的首选后端,支持多 GPU 与多节点

所有已支持操作的完整 API 文档见 MLX 的mlx.core.distributed模块文档。

后端初始化逻辑(init的语义)

mx.distributed.init()通过backend参数选择后端,合法取值为{'any', 'ring', 'jaccl', 'mpi', 'nccl'}

  • any时,MLX 会依次尝试所有可用后端;全部失败则创建单例组(singleton group)。
  • 注意:一旦某个后端成功初始化,后续不带参数(或 backend 为any)调用init()返回同一个后端,不会重复探测。

C++ 侧 mlx/distributed/distributed.cpp 的init()实现印证了这一逻辑:backends是一个静态的unordered_mapinit命中缓存直接返回已注册的 Group;backend="any"时按nccl → ring → mpi → jaccl的顺序尝试(CUDA 环境优先 NCCL),并把成功者注册到"any"键下。以下示例帮助理解:

# Case 1: 强制 MPI,无论 ring 是否可用 world = mx.distributed.init(backend="mpi") world2 = mx.distributed.init() # 后续调用返回 MPI 后端! # Case 2: 初始化任意可用后端 world = mx.distributed.init(backend="any") # 等价于不带参数 world2 = mx.distributed.init() # 同上 # Case 3: 同时初始化两个后端 world_mpi = mx.distributed.init(backend="mpi") world_ring = mx.distributed.init(backend="ring") world_any = mx.distributed.init() # 返回 MPI,因为它先被初始化!

此外,init()还支持strict参数(C++ 签名init(bool strict, const std::string& bk)),strict=true时若any无法初始化任何后端会抛出"Couldn't initialize any backend"异常;非法 backend 名会抛出invalid_argument

分布式程序实战示例

MLX 官方文档提供了两个完整的分布式程序范式:

  • 数据并行(Data Parallelism)
  • 张量并行(Tensor Parallelism)

数据并行中常用的梯度平均可以直接复用mlx.nn.average_gradients(实现见 python/mlx/nn/utils.py),它会把多个小张量的梯度按字节数分组,合并成少数几次大 all reduce 调用,避免大量小通信拖慢性能(详见下文「Tips and Tricks」)。

运行分布式程序:mlx.launch

MLX 提供mlx.launch辅助脚本(入口实现见 python/mlx/_distributed_utils/launch.py),它负责 ssh 到各节点、按 rank 启动进程、注入环境变量并汇聚输出。

沿用开头的例子,在 localhost 上以 4 个进程运行:

$ mlx.launch -n 4 my_script.py 3 array([4, 4, 4, ..., 4, 4, 4], dtype=float32) 2 array([4, 4, 4, ..., 4, 4, 4], dtype=float32) 1 array([4, 4, 4, ..., 4, 4, 4], dtype=float32) 0 array([4, 4, 4, ..., 4, 4, 4], dtype=float32)

-n(即--repeat-hosts)表示每个 host 重复启动的进程数。也可以直接给出远程主机 IP(前提是脚本存在于所有主机上且可 ssh 访问):

$ mlx.launch --hosts ip1,ip2,ip3,ip4 my_script.py 3 array([4, 4, 4, ..., 4, 4, 4], dtype=float32) 2 array([4, 4, 4, ..., 4, 4, 4], dtype=float32) 1 array([4, 4, 4, ..., 4, 4, 4], dtype=float32) 0 array([4, 4, 4, ..., 4, 4, 4], dtype=float32)

mlx.launch的完整行为与参数(hostfile、各后端特有问题)见专门的启动指南文档 docs/src/usage/launching_distributed.rst,下文各后端小节也会展开。从源码看,launch.pymain()会做这些事:

  • --print-python打印当前 Python 可执行文件路径(排查远程主机路径不一致很有用);
  • 解析--hostfile--hosts(默认127.0.0.1),并支持--env注入任意环境变量(如--env MLX_METAL_FAST_SYNCH=1);
  • 自动选择后端:未显式指定时,CUDA 可用则默认nccl,否则默认ring
  • 通过RemoteProcessssh -tt建立交互式会话,把用户 stdin 广播给所有进程、把各进程 stdout/stderr 汇聚到本地输出;某个进程异常退出或mlx.launch被终止时,会终止其余所有进程(对应make_kill_script的 pidfile 机制)。

远程主机的准备清单

要让mlx.launch在远程主机上成功拉起脚本,需要满足:

  • ssh hostname无密码、无 host 确认提示即可登录;
  • Python 二进制在所有主机上位于相同路径(可用mlx.launch --print-python查看本机路径);
  • 待运行的脚本在所有主机上位于相同路径

RING 后端:TCP 环,始终可用

RING 后端不依赖任何第三方库,因此始终可用。它基于 TCP socket,节点间通过网络可达即可。顾名思义,节点按排列:rank 1 只能与 rank 0 和 rank 2 通信,rank 2 只能与 rank 1 和 rank 3 通信,以此类推。因此 RING 后端不支持任意发送者/接收者的send/recv

实现层面,mlx/distributed/ring/ring.cpp 基于非阻塞 socket 与线程池实现通信,内部将单次大于 1 GiB(MAX_IO_BYTES)的传输拆分为多次send(2)/recv(2),以规避send/recv对超过INT_MAX长度返回EINVAL的限制;all sum 默认使用 8 MiB 分块、双缓冲(ALL_SUM_SIZE/ALL_SUM_BUFFERS)流水线化。

用 JSON hostfile 定义环

定义环最方便的方式是写一个 JSON hostfile 配合mlx.launch。对每个节点,指定一个用于 ssh 的主机名(hostname),以及一个或多个该节点监听连接的 IP。下面的 hostfile 定义了一个 4 节点环,hostname1是 rank 0,hostname2是 rank 1,依此类推:

[ {"ssh": "hostname1", "ips": ["123.123.123.1"]}, {"ssh": "hostname2", "ips": ["123.123.123.2"]}, {"ssh": "hostname3", "ips": ["123.123.123.3"]}, {"ssh": "hostname4", "ips": ["123.123.123.4"]} ]

运行mlx.launch --hostfile ring-4.json my_script.py会 ssh 到每个节点并运行脚本,脚本会在每个提供的 IP 上监听连接。具体地,hostname1会主动连接123.123.123.2,并接受来自123.123.123.4的连接,以此类推。

RING 后端的启动参数

mlx.launch中与 RING 强相关的参数(见 launch.py 的launch_ring):

  • --hosts只接受 IP 而非主机名:如果需要 ssh 的主机名与要绑定的 IP 不一致,必须改用 hostfile;
  • --starting-port-p,默认32323):远程主机上绑定的起始端口,rank 0 的第一个 IP 使用该端口,后续每个 IP 或 rank 端口号 +1;
  • --connections-per-ip(默认 1):增加相邻节点间的连接数,等价于mpirun--mca btl_tcp_links 2

launch_ring会把最终的[[ip:port, ...], ...]结构写入MLX_HOSTFILE环境变量并传给每个进程(详细格式见下文「不使用 mlx.launch」一节)。

Thunderbolt 环(高速场景)

虽然 RING 后端在以太网上也能比 MPI 有优势,但其主要用途是借助Thunderbolt 环获得更高带宽。手动配置 Thunderbolt 环比较繁琐,MLX 提供了mlx.distributed_config工具来简化。使用前提:各台电脑能通过以太网或 Wi-Fi 被 ssh 访问;随后用 Thunderbolt 线互联,执行:

mlx.distributed_config --verbose --hosts host1,host2,host3,host4 --backend ring

默认情况下,脚本会尝试发现 Thunderbolt 环,并给出配置每个节点的命令以及可供mlx.launch使用的hostfile.json。如果节点上配置了免密sudo,可用--auto-setup自动配置。

若想手动配置,步骤如下:

  1. 禁用 Thunderbolt bridge 接口;
  2. 对连接 ranki与 ranki + 1的线缆,在节点ii + 1上找到该线缆对应的网络接口;
  3. 为两个节点的对应接口配置一个独立的子网。例如线缆在节点i上对应en2、在节点i + 1上也对应en2,则可以分别为两个节点分配192.168.0.1192.168.0.2。更多细节可参考工具脚本生成的命令。

使用mlx.distributed_config配置网络

工具脚本的完整流程记录在 python/mlx/_distributed_utils/config.py 中,其main()会:

  1. 通过 ssh 检查所有节点可达(对应check_ssh_connections,同时探测免密 sudo);
  2. 通过system_profiler SPThunderboltDataType -jsonnetworksetup -listallhardwareports提取各节点的 Thunderbolt 连接关系(对应extract_connectivity),并用 UUID 反向索引构建连接矩阵;
  3. extract_rings(DFS 找环)或check_valid_mesh(全互联校验)确认拓扑;
  4. 通过IPConfigurator为每条 Thunderbolt 线缆分配192.168.x.y形式的独立子网地址并生成ifconfig/route命令(--auto-setup时自动执行,否则逐节点打印等待确认);
  5. 写出 hostfile(save_hostfile,可用--output-hostfile指定路径,否则打印到 stdout)。

JACCL 后端:Thunderbolt RDMA 低延迟通信

从 macOS 26.2 开始,Thunderbolt 支持 RDMA,可实现 Thunderbolt 5 Mac 之间的低延迟通信。MLX 提供的 JACCL 后端利用该能力,通信延迟比 RING 后端低一个数量级

名字由来:JACCL(读作 Jackal,豺狼)是Jack and Angelos' Collective Communication Library的缩写,既是向 NVIDIA NCCL 的致敬式双关,也纪念主导 Apple RDMA over Thunderbolt 研发的Jack Beasley

启用 RDMA

该功能尚未完全成熟,启用 RDMA 需要一定步骤,并且即使有 sudo 也无法远程完成,必须在 macOS 恢复模式(recovery)中操作:

  1. 将电脑启动到恢复模式;
  2. 通过 Utilities → Terminal 打开终端;
  3. 运行rdma_ctl enable
  4. 重启。

验证是否启用成功,可运行ibv_devices,以 M3 Ultra 为例输出类似:

~ % ibv_devices device node GUID ------ ---------------- rdma_en2 8096a9d9edbaac05 rdma_en3 8196a9d9edbaac05 rdma_en5 8396a9d9edbaac05 rdma_en4 8296a9d9edbaac05 rdma_en6 8496a9d9edbaac05 rdma_en7 8596a9d9edbaac05

定义 Mesh(全互联拓扑)

JACCL 后端只支持全互联(fully connected)拓扑:任意两台 Mac 之间都必须有 Thunderbolt 线缆直连。下图左侧是合法的 4 节点 mesh(任意节点两两相连),右侧则不合法(M3 Ultra 1 与 M3 Ultra 2 之间没有连接):

四台 M3 Ultra 的全互联 mesh(合法拓扑)。

非法 mesh:M3 Ultra 1 未与 M3 Ultra 2 相连。

与 RING 类似,使用 JACCL 最简单的方式是写一个供mlx.launch使用的 JSON hostfile。hostfile 需要包含:

  • 用于 ssh 启动脚本的主机名;
  • rank 0 的一个 IP,要求所有节点都能访问;
  • 连接每个节点到其他每个节点的 RDMA 设备列表。

下面这个 JSON 定义了上图中合法的 4 节点 mesh:

[ { "ssh": "m3-ultra-1", "ips": ["123.123.123.1"], "rdma": [null, "rdma_en5", "rdma_en4", "rdma_en3"] }, { "ssh": "m3-ultra-2", "ips": [], "rdma": ["rdma_en5", null, "rdma_en3", "rdma_en4"] }, { "ssh": "m3-ultra-3", "ips": [], "rdma": ["rdma_en4", "rdma_en3", null, "rdma_en5"] }, { "ssh": "m3-ultra-4", "ips": [], "rdma": ["rdma_en3", "rdma_en4", "rdma_en5", null] } ]

rdma是一个 N×N 矩阵:第i行第j列表示节点i用于连接节点j的 RDMA 设备名,对角线上是null(自己连自己没有意义)。尽管 Thunderbolt RDMA 通信不走 TCP/IP,仍然需要禁用 Thunderbolt bridge,并为每条 Thunderbolt 连接设置隔离的本地网络(这正是IPConfigurator做的事)。

mlx.launchlaunch_jaccl会对 hostfile 做校验:要求每个 host 的rdma列表长度与节点数一致、对角线必须为null、且每对节点都必须有 RDMA 设备(缺失时会报错并列出前三个缺失的对),随后把MLX_JACCL_COORDINATOR=ip:portMLX_IBV_DEVICES(rdma 矩阵 JSON)注入各进程。

上述所有配置也可改用mlx.distributed_config完成,该脚本会:

  • ssh 到每个节点;
  • 提取 Thunderbolt 连接关系;
  • 检查是否为合法 mesh;
  • 提供配置每个节点的命令(有 sudo 则直接执行);
  • 生成供mlx.launch使用的 hostfile。

端到端实战:JACCL 分布式推理

假设节点均可 ssh 访问且配置了免密 sudo,启动一个使用 JACCL 的分布式脚本非常直接。

首先,连接所有 Thunderbolt 线缆,然后用mlx.distributed_config可视化验证连接:

mlx.distributed_config --verbose \ --hosts m3-ultra-1,m3-ultra-2,m3-ultra-3,m3-ultra-4 \ --over thunderbolt --dot | dot -Tpng | open -f -a Preview

确认拓扑无误后,自动配置节点并把 hostfile 保存为m3-ultra-jaccl.json

mlx.distributed_config --verbose \ --hosts m3-ultra-1,m3-ultra-2,m3-ultra-3,m3-ultra-4 \ --over thunderbolt --backend jaccl \ --auto-setup --output m3-ultra-jaccl.json

之后即可运行分布式脚本,例如用 MLX LM 对超大模型做分布式推理:

mlx.launch --verbose --backend jaccl --hostfile m3-ultra-jaccl.json -- \ /path/to/remote/python -m mlx_lm chat --model mlx-community/DeepSeek-R1-0528-4bit

关于MLX_METAL_FAST_SYNCH:通过--env MLX_METAL_FAST_SYNCH=1设置该环境变量,可启用一种不同的、更快的 GPU 与 CPU 同步方式。它并非 JACCL 专用,任何需要 CPU 与 GPU 协作计算的场景都适用;对低延迟通信尤为重要,因为 RDMA 通信由 CPU 执行。但该机制不可靠,可能导致死锁并使 GPU 卡死,因此默认关闭,最好保持不设置

自定义 side channel(侧信道)

JACCL 初始化时,会在各 rank 之间通过一条侧信道(side channel)交换 RDMA 连接元数据。默认是一条简单的 TCP 星型 all-gather;你可以通过mlx.core.distributed.initall_gather_factory参数提供自定义的侧信道 all-gather(仅backend='jaccl'时可用)。

该 factory 会以(rank, size)为参数在每个 rank 上调用一次,必须返回一个签名为f(src: bytes, n_bytes: int) -> bytes的可调用对象。返回的 bytes 长度必须为size * n_bytes,内容是按 rank 顺序拼接的所有 rank 的输入。

def make_side_channel(rank, size): def all_gather(src: bytes, n_bytes: int) -> bytes: # 与所有 rank 交换 src,并返回 size * n_bytes 字节 ... return all_gather world = mx.distributed.init( backend="jaccl", all_gather_factory=make_side_channel, )

C++ 侧对应 mlx/distributed/distributed.cpp 的init(bool strict, const std::string& bk, AllGatherFactory factory)重载:非jaccl后端传入 factory 会抛出invalid_argument,且后端已初始化时 factory 会被忽略(返回缓存的 Group)。

NCCL 后端:CUDA 环境的首选

MLX 的 CUDA 版本内置了与 NCCL 通信的能力。NCCL 是高性能集体通信库,支持多 GPU 与多节点。在 CUDA 环境中,NCCL 是mlx.launch的默认后端,运行分布式任务只需:

mlx.launch -n 8 test.py # 适合交互式脚本 mlx.launch -n 8 python -m mlx_lm chat --model my-model

也可以用mlx.launch以同样方式 ssh 到远程节点启动:

mlx.launch --hosts my-cuda-node -n 8 test.py

实现上,mlx/distributed/nccl/nccl.cpp 通过 NCCL 的 C API 封装了通信原语,并把 MLX 的 dtype 映射为 NCCL 类型(int8_t→ncclCharfloat16_t→ncclHalfbfloat16_t→ncclBfloat16等),超时可通过MLX_NCCL_TIMEOUT环境变量调整(默认 300000 毫秒)。launch_nccl--verbose时会自动追加NCCL_DEBUG=INFO

很多场景下你可能不想用mlx.launch,而是由集群调度器直接拉起进程,此时需要手动设置 NCCL 后端初始化所需的环境变量(见下文「不使用 mlx.launch」)。

从 Mac 发起、在 Linux CUDA 节点执行

mlx.launch的 NCCL 后端默认针对 CUDA 环境;当从 Mac 启动、目标是带 CUDA 的 Linux 机器时,应显式指定--backend nccl--repeat-hosts, -n用于多节点多 GPU 作业,例如:

mlx.launch --backend nccl --hosts linux-1,linux-2 -n 8 -- ./my-job.sh

会尝试启动 16 个进程,每个节点 8 个,全部运行my-job.shlaunch_nccl会为每个进程设置NCCL_HOST_IPNCCL_PORT--nccl-port,默认 12345)、MLX_WORLD_SIZE,并按rank % repeat_hosts为进程分配CUDA_VISIBLE_DEVICES

MPI 后端:成熟生态与mpirun集成

只要机器上安装了 MPI,MLX 就具备“讲 MPI 语言”的能力。可以用mpirun启动,但下文示例使用mlx.launch --backend mpi,它会代为处理一些麻烦事,例如为mpirun可执行文件和libmpi.dyld共享库设置绝对路径。

最简单的用法(沿用本文开头的求和示例):

$ mlx.launch --backend mpi -n 2 test.py 1 array([2, 2, 2, ..., 2, 2, 2], dtype=float32) 0 array([2, 2, 2, ..., 2, 2, 2], dtype=float32)

上述命令在本机启动两个进程,可以看到两个进程的 stdout:进程把全 1 数组发给对方并求和后打印。用mlx.launch -n 4 ...则会打印 4。

launch_mpi的实现(见 launch.py)会用which mpirun找到可执行文件,通过otool/ldd探测libmpi的真实库名(兼容 Homebrew 与 pip 安装),自动注入DYLD_LIBRARY_PATHMLX_MPI_LIBNAME,并生成临时 hostfile(hostname slots=N格式)后以--output :raw模式调用mpirun,保证输出不被行缓冲。

安装 MPI

MPI 可以通过 Homebrew、pip、Anaconda 包管理器安装,或从源码编译。官方测试大多使用 Anaconda 安装的openmpi

$ conda install conda-forge::openmpi

用 Homebrew 或 pip 安装时,需要指定libmpi.dyld的位置以便 MLX 在运行时加载,只需通过DYLD_LIBRARY_PATH环境变量传给mpirun即可(mlx.launch会自动完成)。某些环境使用非标准库文件名,可通过MPI_LIBNAME环境变量指定(mlx.launch同样自动处理):

$ mpirun -np 2 -x DYLD_LIBRARY_PATH=/opt/homebrew/lib/ -x MPI_LIBNAME=libmpi.40.dylib python test.py $ # 或简写为 $ mlx.launch -n 2 test.py

配置远程主机

MPI 可以自动连接远程主机并建立网络通信,前提是远程主机可通过 ssh 访问。排查连通性问题的检查清单:

  • 从所有机器到所有机器执行ssh hostname都不要求密码或 host 确认;
  • 所有机器上都能找到mpirun
  • 确保 MPI 使用的hostname与所有机器~/.ssh/config中配置的一致。

调优 MPI All Reduce

提示:需要更快的 all reduce 时,可考虑使用 RING 后端(无论是 Thunderbolt 连接还是以太网)。

  • 通过--mca btl_tcp_links N配置每对主机间使用 N 条 TCP 连接以提升带宽;
  • 通过--mca btl_tcp_if_include <iface>强制 MPI 使用性能最优的网络接口(<iface>换成目标接口名)。

mlx.launch中可通过--mpi-arg透传这些参数给mpirun,例如:

mlx.launch --backend mpi --mpi-arg '--mca btl_tcp_if_include en0' --hostfile hosts.json my_script.py

使用 MPI 后端时注意三点(launch.pylaunch_mpi逻辑):hostfile 中的 IP 会被忽略;ssh 连通性要求更强,每个节点必须能连到其他所有节点;mpirun必须在每个节点上位于相同路径。

不使用mlx.launch:手动设置环境变量

没有任何一个分布式后端要求必须用mlx.launch启动mlx.launch本质上只是连接到各主机、按 rank 启动进程、设置必要的环境变量,然后委托给你的 MLX 脚本。很多场景下手动设置环境变量反而是最简单的方式——典型的例子是使用集群调度器:调度器在任务排定时主机尚未确定,由它来拉起所有进程。

下面列出每个后端所需的环境变量。

RING 后端

  • MLX_RANK:单个 0 起始的整数,定义本进程的 rank。
  • MLX_HOSTFILE:JSON 文件路径,内容为每个 rank 要监听的 IP 与端口列表,例如:
[ ["123.123.1.1:5000", "123.123.1.2:5000"], ["123.123.2.1:5000", "123.123.2.2:5000"], ["123.123.3.1:5000", "123.123.3.2:5000"], ["123.123.4.1:5000", "123.123.4.2:5000"] ]
  • MLX_RING_VERBOSE(可选):设为 1 时启用更多分布式后端日志。

JACCL 后端

  • MLX_RANK:单个 0 起始的整数,定义本进程的 rank。
  • MLX_JACCL_COORDINATOR:rank 0 监听的 IP 与端口,其他 rank 连接它以建立 RDMA 连接。
  • MLX_IBV_DEVICES:JSON 文件路径,内容为连接每对节点的 ibverbs 设备名矩阵,例如:
[ [null, "rdma_en5", "rdma_en4", "rdma_en3"], ["rdma_en5", null, "rdma_en3", "rdma_en4"], ["rdma_en4", "rdma_en3", null, "rdma_en5"], ["rdma_en3", "rdma_en4", "rdma_en5", null] ]

NCCL 后端

  • MLX_RANK:单个 0 起始的整数,定义本进程的 rank。
  • MLX_WORLD_SIZE:将要启动的进程总数。
  • NCCL_HOST_IPNCCL_PORT:所有主机都能访问、用于建立 NCCL 通信的 IP 与端口。
  • CUDA_VISIBLE_DEVICES:本进程对应的本地 GPU 索引。

当然,NCCL 使用的其他任何环境变量也都可以照常设置。

Tips and Tricks:分布式实践建议

以下技巧有助于更好利用 MLX 的分布式通信能力。

  • 先在本地小规模测试。mlx.launch -n2 -- my_script.py在单节点上先跑通小规模测试。

  • 批量合并通信。如数据并行训练示例所述,大量小通信会拖累性能。参考mlx.nn.average_gradients的做法,把许多小通信合并成一次大通信:该函数(python/mlx/nn/utils.py)默认把梯度按all_reduce_size=32MiB分组——先提取所有梯度的形状、大小、dtype,将 dtype 一致的梯度连续拼接后一次 all reduce,再按形状切分还原;dtype 混杂时会退化为逐张量 all reduce(all_reduce_size=0禁用分组)。单进程组(size()==1)时直接原样返回梯度,不做任何通信。

  • 可视化连接。mlx.distributed_config --hosts h1,h2,h3 --over thunderbolt --dot导出节点连接图(GraphViz DOT 格式),确保线缆连接正确(tb_connectivity_to_dot会输出带接口标签的图,例如en2/en2)。

  • 善用调试器。mlx.launch专为交互式使用设计:它把 stdin 广播给所有进程、汇总所有进程的 stdout,这让pdb调试变得非常轻松。

参考路径速查

  • 分布式使用指南(本文主体):docs/src/usage/distributed.rst
  • 启动工具专项文档:docs/src/usage/launching_distributed.rst
  • 环境变量说明:docs/src/usage/environment_variables.rst
  • 启动器实现:python/mlx/_distributed_utils/launch.py
  • 配置工具实现:python/mlx/_distributed_utils/config.py
  • hostfile 解析:python/mlx/_distributed_utils/common.py
  • 核心分布式实现:mlx/distributed/distributed.cpp、mlx/distributed/ops.h
  • RING 实现:mlx/distributed/ring/ring.cpp
  • NCCL 实现:mlx/distributed/nccl/nccl.cpp
  • 数据并行示例:docs/src/examples/data_parallelism.rst
  • 张量并行示例:docs/src/examples/tensor_parallelism.rst
  • 梯度平均工具:python/mlx/nn/utils.py </|DSML|parameter> </|DSML|invoke> </|DSML|tool_calls>

【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询