实战指南:用 AllToShardedLinear 与 ShardedToAllLinear 实现 Llama 多设备推理)
MLX 张量并行Tensor Parallelism实战指南用 AllToShardedLinear 与 ShardedToAllLinear 实现 Llama 多设备推理【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx导读本指南以 MLX 官方示例 tensor_parallelism.rst 为主线系统讲解如何在 Apple silicon 多设备环境下通过mlx.nn提供的分布式线性层AllToShardedLinear/ShardedToAllLinear与两个 shard 工具函数shard_linear/shard_inplace实现张量并行TP并最终落地到 Llama 风格 Transformer 的多卡推理。读完本文你将掌握 TP 层的数据切分与通信语义、两种 shard 工具的选择依据以及如何用mlx.launch -n 2一行命令在 2 台及以上设备上运行大模型推理脚本。一、背景为什么张量并行适合 MLX张量并行与数据并行Data Parallelism的最大区别在于切分对象数据并行切分 batch每个设备持有一份完整权重而张量并行切分的是模型权重本身把单个超大线性层的参数分布到多个设备上每个设备只负责计算一部分。因此当模型权重超过单台设备如 Apple silicon 的统一内存可容纳的规模时张量并行是更自然的选择。仓库中 examples/python/distributed_tensor_parallel.py 的模块 docstring 对这一思路做了精炼概括Unlike data parallelism this splits the model rather than the batch, so it is what you reach for when the weights are too big for one machine.与数据并行不同张量并行切分的是模型而不是 batch因此当权重对单机来说过大时它是你应当考虑的方案。MLX 在mlx.nn中为张量并行提供了四个开箱即用的分布式层和两个工具函数全部实现在 python/mlx/nn/layers/distributed.py并通过 python/mlx/nn/layers/init.py 导出为nn.AllToShardedLinear、nn.ShardedToAllLinear、nn.QuantizedAllToShardedLinear、nn.QuantizedShardedToAllLinear以及nn.layers.distributed.shard_linear、nn.layers.distributed.shard_inplace。二、分片线性层Sharded Layers2.1 AllToShardedLinear列方向切分输出保持分片nn.AllToShardedLinear接收每台设备都完整的公共输入把权重矩阵沿输出维度切分到分布式组mlx.core.distributed.Group的所有设备上产出分片输出。以文档中的示例参数为例input_dims2、output_dims2输入形状(4, 2)设备组大小为 2。此时权重按输出维度被切成两半每台设备持有形状为(1, 2)的权重切片收到完整的(4, 2)输入后各自计算出一个(4, 1)的部分输出。从源码 python/mlx/nn/layers/distributed.py 可以看到其构造逻辑初始化时直接用output_dims // N作为本地权重行数N为组内设备数并显式检查output_dims % N 0若不满足可整除条件会抛出ValueError# python/mlx/nn/layers/distributed.py (AllToShardedLinear.__init__) N self.group.size() if (output_dims % N) ! 0: raise ValueError( fCannot shard the output of size {output_dims} across {N} devices. ) self.weight mx.random.uniform( low-scale, highscale, shape(output_dims // N, input_dims), )关键点该层不会自动把各设备的输出 gather 回来。这是文档中特别强调的“有意设计”见第四节“Useful Design Choices”目的是让它的输出形状天然成为下一个 sharded-to-all 层的输入。2.2 ShardedToAllLinear行方向切分自动 all_sum 聚合nn.ShardedToAllLinear恰好与前者互补它期望输入已经沿特征维度被切分权重则沿输入维度切分到各设备并在计算完本地结果后通过mx.distributed.all_sum自动聚合组内所有设备最终拿到完全一致的结果。仍以input_dims2、output_dims2、输入(4, 2)、2 台设备为例权重沿输入维度一分为二每台设备持有(2, 1)的权重切片输入也被切成两份分别喂给对应设备每台设备先算出(4, 2)的本地结果再经过all_sum得到最终层输出。源码实现中前向路径是先做本地矩阵乘法再调用mx.distributed.all_sum做组内归约最后加上偏置# python/mlx/nn/layers/distributed.py (ShardedToAllLinear.__call__) def __call__(self, x: mx.array) - mx.array: x x self[weight].T x mx.distributed.all_sum(x, groupself.group) if bias in self: x x self[bias] return x注意该层不会自动帮你切分输入你必须在喂入之前构造好“partial”已切分的输入结构这也是有意设计见第四节。2.3 量化版本Quantized 两个变体nn.QuantizedAllToShardedLinearAllToShardedLinear的量化等价物权重用mx.quantize量化后参与mx.quantized_matmul计算nn.QuantizedShardedToAllLinearShardedToAllLinear的量化等价物同样在本地量化矩阵乘后执行all_sum聚合。与nn.QuantizedLinear类似这两类量化层的参数是冻结的frozen不会进入任何梯度计算。从源码可见构造末尾会调用self.freeze()并且unfreeze被覆写为“解冻内部子层但自身参数保持冻结”。量化层默认参数为group_size64、bits4、modeaffine均透传给mx.quantize/mx.quantized_matmul也可通过modemxfp8等非 affine 模式使用。三、分片工具函数shard_linear 与 shard_inplace3.1 shard_linear零侵入地把普通 Linear 变成分布式层nn.layers.distributed.shard_linear(module, sharding, *, segments1, groupNone)输入一个已有的nn.Linear或nn.QuantizedLinear层输出一个全新的分布式层AllToShardedLinear/ShardedToAllLinear或其量化变体sharding只能是all-to-sharded或sharded-to-all其他取值会触发ValueErrorsegments若权重是融合矩阵如融合的 QKV可传入段数默认1group默认使用全局分布式组原层不会被修改只是被读取后作为新层参数的数据来源。源码 python/mlx/nn/layers/distributed.py 中的shard_linear会根据sharding类型与module是否为Linear分派到对应的from_linear/from_quantized_linear工厂方法fns { (all-to-sharded, True): AllToShardedLinear.from_linear, (all-to-sharded, False): QuantizedAllToShardedLinear.from_quantized_linear, (sharded-to-all, True): ShardedToAllLinear.from_linear, (sharded-to-all, False): QuantizedShardedToAllLinear.from_quantized_linear, } return fnssharding, isinstance(module, Linear)以AllToShardedLinear.from_linear为例它从原层取出output_dims, input_dims先构造一个新层再调用内部_shard按_all_to_sharded谓词切分参数并update进新层classmethod def from_linear(cls, linear_layer, *, segments1, groupNone): group group or mx.distributed.init() output_dims, input_dims linear_layer.weight.shape sl cls(input_dims, output_dims, hasattr(linear_layer, bias), group) sl.update(_shard(linear_layer.parameters(), _all_to_sharded(segments), group)) return sl3.2 shard_inplace原地切分不引入通信nn.layers.distributed.shard_inplace(module, sharding, *, segments1, groupNone)原地修改已有层的参数字典把每个参数替换为当前 rank 应持有的切片不创建新层不添加任何分布式通信如果该层自身不支持或未启用分布式通信前向/反向中不会自动产生all_sum等操作需要你手动处理sharding除了字符串还可以是自定义可调用对象给定(path, weight)返回切分轴int或(axis, segments)元组这为融合权重如 QKV 合并矩阵segments3等复杂切分提供了灵活性。内部_shard实现python/mlx/nn/layers/distributed.py使用tree_map_with_path遍历参数树对每个mx.array先按segments分段、再按组大小N二次切分取出当前 rank 的那一段并沿原轴mx.concatenate后做mx.contiguousreturn mx.contiguous( mx.concatenate( [_split(part, N, axis)[r] for part in _split(weight, segments, axis)], axisaxis, ) )文档与源码的权衡建议当你的模型层是标准的nn.Linear/nn.QuantizedLinear时优先使用shard_linear——它一次搞定“参数切分 分布式通信”只有在你自定义的层已经内建了通信逻辑、只需要把参数切片放到位时才考虑shard_inplace。四、有意的设计选择Useful Design Choices为什么要让AllToShardedLinear不自动 gather、ShardedToAllLinear不自动 shard因为这两类层天然成对出现前者all-to-sharded的输出恰好就是后者sharded-to-all所需的输入。两个层串联时中间无需任何 gather/shard 转换步骤直接省掉一次通信往返降低通信开销。文档用一个两层的简单模型演示了这一衔接x ... # some (4, 2) model input: batch size 4, feature size 2 l1 nn.AllToShardedLinear(2, 2, biasFalse) # initialize the layer l1_out l1(x) # (4, 1) output l2 nn.ShardedToAllLinear(2, 2, biasFalse) l2_out l2(l1_out) # (4, 2) output在 2 设备场景下l1把输出维切成两份每台设备产出(4, 1)l1_out以分片形式直接喂给l2l2按输入维切分权重、本地算完后all_sum聚合两台设备都得到一致的(4, 2)。整个链条只有一次 all reduce 通信。相关实现细节反向传播中的梯度聚合值得注意的是AllToShardedLinear的“不 gather”只针对前向输出梯度仍然会自动跨组聚合。源码中用mx.custom_function定义了一个带 VJP 的sum_gradients包装lru_cache def sum_gradients(group): if group.size() 1: return lambda x: x mx.custom_function def f(x): return x f.vjp def f(x, dx, _): return mx.distributed.all_sum(dx, groupgroup) return f前向时sum_gradients(group)(x)是恒等函数而反向时把来自各分片的梯度做all_sum保证训练场景下梯度在分片上的正确归约。这也解释了AllToShardedLinear.__call__中第一行x sum_gradients(self.group)(x)的作用。一个可以运行的端到端验证仓库中的 examples/python/distributed_tensor_parallel.py 把上述理念落成可运行脚本随机初始化一个 MLPup: Linear(dims, hidden)down: Linear(hidden, dims)先计算完整模型的输出作为基准再用AllToShardedLinear.from_linear和ShardedToAllLinear.from_linear替换两个线性层最后比较分片模型与完整模型的输出最大绝对误差world mx.distributed.init() # Seeding the global rng gives every rank the same weights to shard. mx.random.seed(0) model MLP(dims, hidden) mx.eval(model.parameters()) x mx.random.normal((num_tokens, dims), keymx.random.key(1)) expected model(x) mx.eval(expected) model.up nn.AllToShardedLinear.from_linear(model.up, groupworld) model.down nn.ShardedToAllLinear.from_linear(model.down, groupworld) mx.eval(model.parameters()) y model(x) mx.eval(y) # every rank must evaluate: down ends in an all reduce运行方式python examples/python/distributed_tensor_parallel.py mlx.launch -n 2 python examples/python/distributed_tensor_parallel.py mlx.launch -n 4 python examples/python/distributed_tensor_parallel.py脚本会在 rank 0 上打印Max |sharded - full|用于核对分片模型与完整模型输出的一致性。这里mx.eval(y)是必须的——由于down层以 all reduce 收尾跳过求值的 rank 会让其他设备一直等待。该示例同时也被 python/tests/mlx_distributed_tests.py 的test_shard_linear测试覆盖测试先构造nn.Linear(1024, 1024)分别用shard_linear(lin, all-to-sharded)与shard_linear(lin, sharded-to-all)生成两个分片层再断言y2 slin2(x[part])与完整输出y lin(x)逐元素接近、y1 slin1(x)与y[part]接近随后还验证了量化版本与mxfp8模式mode正确传播、mxfp8 无biases参数以及一个四层Sequentialall-to-sharded / sharded-to-all 交替在nn.value_and_grad下与完整模型梯度一致。五、Llama 推理中的张量并行实战5.1 从单机推理到多设备推理本节把 TP 应用到 Llama Inference 示例 的 Transformer 上实现思路与单机版完全一致只是多了“初始化分布式组 对每层做 shard”两步最终目标是用mlx.launch -n 2 llama.py跨两台及以上设备推理。第一步是初始化分布式通信组并取得当前进程的 rankworld mx.distributed.init() rank world.rank()注意 MLX 分布式 API 的语义当组大小为 1 时mx.distributed下的操作全部退化为 no-op因此单设备运行时这段代码不会有任何副作用无需额外的if world.size() 1分支保护详见 docs/src/usage/distributed.rst。5.2 Transformer 块中两个天然的 TP 切入点Llama Transformer 块中有两处适合张量并行Attention 块与FFN 块。两者遵循同一模式——多个并行线性层作用于相同输入再接一个输出线性层Attention 块Q、K、V 三个投影沿输出维切分all-to-sharded输出投影wo沿输入维切分sharded-to-allFFN 块gatew1与 upw3投影改为 all-to-shardeddownw2投影改为 sharded-to-all。5.3 为什么中间算子不影响切分线性层之间的中间运算不会破坏 TP 范式因为它们属于两类安全操作逐元素操作RoPE、FFN 中的逐元素乘法对每个元素/位置独立运算保持既有分片模式无需跨设备通信作用于未切分维度的操作softmax、scaled dot-product attention它们沿序列长度或 head 维度计算而这些维度未被切分各设备可独立完成。Attention 中的Q K^T与scores V之所以在 Q/K/V 已分片时依然正确是因为矩阵乘法恰好沿着被切分的特征维进行结果仍保持正确的分片布局可供后续 sharded-to-all 层直接消费。5.4 代码改动Attention 与 FeedForward 的 shard 方法实现上使用shard_linear而非shard_inplace来同时获得参数切分与分布式通信避免在__call__中手写通信步骤。Attention 块的shard方法如下# ... in Attention class def shard(self, group: mx.distributed.Group): self.n_heads self.n_heads // group.size() self.n_kv_heads self.n_kv_heads // group.size() self.wq nn.layers.distributed.shard_linear(self.wq, all-to-sharded, groupgroup) self.wk nn.layers.distributed.shard_linear(self.wk, all-to-sharded, groupgroup) self.wv nn.layers.distributed.shard_linear(self.wv, all-to-sharded, groupgroup) self.wo nn.layers.distributed.shard_linear(self.wo, sharded-to-all, groupgroup)两个细节值得注意head 数量按组大小整除n_heads与n_kv_heads都要除以group.size()因为每个 head 的 Q/K/V 投影在 shard 后只由一台设备持有head 总数必须能被设备数整除否则无法均分wq/wk/wv 是 all-to-shardedwo 是 sharded-to-all与 Attention 的数据流完全对应——Q/K/V 投影共享同一输入、输出沿输出维分片wo消费分片输入、输出聚合回全量。FFN 块的改动同理# ... in FeedForward class def shard(self, group: mx.distributed.Group): self.w1 nn.layers.distributed.shard_linear(self.w1, all-to-sharded, groupgroup) self.w2 nn.layers.distributed.shard_linear(self.w2, sharded-to-all, groupgroup) self.w3 nn.layers.distributed.shard_linear(self.w3, all-to-sharded, groupgroup)5.5 在 load_model 中统一应用最后在load_model函数中对所有 Transformer 层批量调用 shard仅在多设备时生效# ... in load_model function if world.size() 1: # convert Linear layers in Transformer/FFN to appropriate Sharded Layers for layer in model.layers: layer.attention.shard(groupworld) layer.feed_forward.shard(groupworld)完成上述改动后原有的单机推理文件无需任何其他修改python llama.py依旧可用要跨设备运行只需mlx.launch -n 2 llama.pymlx.launch是 MLX 自带的分布式启动脚本实现见 python/mlx/_distributed_utils/launch.py-n/--repeat-hosts指定每个 host 上启动的进程数默认在127.0.0.1本机启动也可用--hosts ip1,ip2,ip3,ip4指定远程机器前提是脚本存在于所有 host 且可 ssh 免密访问。启动脚本会为每个进程注入MLX_RANK环境变量随后mx.distributed.init()即可据此构建通信组。后端的选取逻辑见 docs/src/usage/distributed.rstmlx.launch默认在 CUDA 可用时选nccl、否则选ring也可通过--backend或 hostfile 中的backend字段指定ring/mpi/nccl/jaccl/jaccl-ring。一个易踩的坑参与 TP 的所有设备必须以相同随机种子初始化模型mx.random.seed(0)否则各 rank 持有不同权重的切片all_sum聚合出来的结果毫无意义。这与 examples/python/distributed_tensor_parallel.py 中“Seeding the global rng gives every rank the same weights to shard”的做法一致。六、量化与张量并行组合的注意事项量化层QuantizedAllToShardedLinear/QuantizedShardedToAllLinear的参数冻结不会参与梯度计算天然适配推理场景QuantizedMatmul目前仅在非 CUDA 环境可用见 python/tests/mlx_distributed_tests.py 中# QuantizedMatmul is not supported on CUDA的注释与分支使用shard_linear对QuantizedLinear做分片时量化配置group_size、bits、mode会从原层透传到新层分片前提是维度可整除AllToShardedLinear要求output_dims % N 0ShardedToAllLinear要求input_dims % N 0不满足时构造会直接抛ValueError这也是选择设备数时的重要约束例如 7B 模型某层输出维为 4096则可整除 2、4、8。七、小结本文围绕 MLX 官方的张量并行示例完整覆盖了以下知识链路两种分布式线性层AllToShardedLinear列切分、输出分片、反向自动聚合梯度与ShardedToAllLinear行切分、前向all_sum聚合以及各自的量化变体两种工具函数shard_linear新建分布式层、自动带通信与shard_inplace原地切分、不引入通信的适用场景与参数语义设计理念两个层刻意不做自动 gather/shard是为了前后串联时省掉一次中间通信Llama 实战Attention 的 Q/K/V 与 FFN 的 gate/up 走 all-to-shardedAttention 的wo与 FFN 的down走 sharded-to-allhead 数按组大小整除最后在load_model中统一 shard并用mlx.launch -n 2 llama.py一键多设备运行。如需继续深入可进一步阅读仓库中的 llama-inference.rst单机版完整模型实现含 RoPE、KV cache 与权重转换脚本以及 distributed.rst分布式通信后端与mlx.launch的完整说明分布式层的底层实现可对照 distributed.py数值正确性验证可参考 mlx_distributed_tests.py。【免费下载链接】mlxMLX: An array framework for Apple silicon项目地址: https://gitcode.com/GitHub_Trending/ml/mlx创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考