2026-08-04 更新
这个最初的 Tensor Parallel demo 已经继续扩展
当前项目已经从基础 TP 扩展到 Tensor、Sequence、Vocab、Pipeline 和 Context Parallelism,并支持 TP x PP x CP 的多维组合。
目前实现了:
- Tensor / Sequence / Vocab Parallelism;
- Pipeline Parallelism,包括 GPipe 和 1F1B;
- Context Parallelism,包括 All-Gather K/V 和 A2A 两种通信方式
- TP × CP × PP 多维组合;
- Dense 与并行模型的 forward、backward 和多步训练数值验证;
- 统一的显存与吞吐 benchmark。
项目同时整理了一套在线教程:
https://anker661.github.io/minimind-tron/
目前 TP / SP / VP 部分已整理完成,后续 PP, CP 以及多维并行部分将在整理好后发布。
感谢 MiniMind 提供了足够小且清晰的模型实现,让这些并行策略能够从原理、tensor layout、通信和 autograd 一步步展开。
jingyao你好!非常棒的项目,我在Minimind的基础上实现了一个用来学习Tensor Parallel的版本
https://github.com/ANKer661/minimind/tree/tp_demo
实现内容
当前实现参考了Nvidia 2019年的最初的Megatron LM论文的层内Tensor Parallel方式
- Q/K/V、gate/up projection 按输出维度切分。
- attention output、down projection 按输入维度切分,并通过 all-reduce 合并结果。
- 每个 TP rank 仅计算本地 attention heads 和 MLP intermediate 分片。
- 使用自定义
torch.autograd.Function 实现 TP 通信语义。
- 定义好Parallel版本的layer后,简单替换Minimind中的对应组件即可实现Tensor Parallel。
- 支持从普通 MiniMind 模型权重切片加载 TP 模型。
- 提供 Dense 与 TP 的 forward、backward 和多步 AdamW 数值对齐验证。
- 提供 Dense 与 TP 的显存及训练 step 耗时对比脚本。
为了避免影响现有代码,目前 TP 实现均放在独立文件中,没有修改原有模型和训练流程:
model/model_tp.py
scripts/demo_tensor_parallel.py
scripts/benchmark_tensor_parallel_memory.py
scripts/benchmark_tensor_parallel_worker.py
局限
当前实现目前主要用于理解Tensor Parallel的原理和实现和数值验证,存在以下限制:
- 仅支持单机多卡的 Tensor Parallel。
- 暂未与现有 DDP/训练脚本集成。
- 暂不支持 KV Cache 和带缓存的生成。
- 暂不支持 MoE。
- 暂不支持Linear之外的并行,例如 vocab parallel embedding、lm_head 和 parallel cross entropy。
- 暂不支持 TP checkpoint shard 的独立保存和加载。
- 目标是展示 TP 原理和数值正确性,并非替代 Megatron-LM、DeepSpeed 等成熟框架。
验证结果
数值验证
结果包括加载相同权重后,第一次forward的logits误差,第一次backward后所有被并行的layer的梯度误差,以及使用AdamW优化器进行更新后每x步的logits的累积误差。
torchrun --standalone --nproc_per_node=2 scripts/demo_tensor_parallel.py \
--tp_size 2 \
--hidden_size 768 \
--num_hidden_layers 8 \
--num_attention_heads 8 \
--num_key_value_heads 4 \
--vocab_size 6400 \
--seq_len 340 \
--batch_size 32 \
--dtype float32 \
--check_backward \
--optimizer_steps 100 \
--log_interval 10 \
--learning_rate 5e-4 \
--seed 42 \
--atol 1e-3
forward max logits diff: 6.198883e-06
forward loss diff: 0.000000e+00
backward max grad diff: 1.006993e-08
layer 0:
self_attn.q_proj.weight mean=8.708e-10 max=5.530e-09 rel_l2=2.877e-06
self_attn.k_proj.weight mean=1.237e-09 max=7.014e-09 rel_l2=2.889e-06
self_attn.v_proj.weight mean=1.690e-09 max=1.007e-08 rel_l2=2.366e-06
self_attn.o_proj.weight mean=1.176e-09 max=7.276e-09 rel_l2=2.304e-06
self_attn.q_norm.weight mean=1.133e-09 max=4.540e-09 rel_l2=2.104e-06
self_attn.k_norm.weight mean=1.248e-09 max=4.075e-09 rel_l2=2.375e-06
mlp.gate_proj.weight mean=4.554e-10 max=3.289e-09 rel_l2=2.535e-06
mlp.up_proj.weight mean=4.580e-10 max=3.609e-09 rel_l2=2.513e-06
mlp.down_proj.weight mean=4.606e-10 max=4.191e-09 rel_l2=2.516e-06
layer 1:
self_attn.q_proj.weight mean=1.862e-10 max=1.106e-09 rel_l2=3.106e-06
self_attn.k_proj.weight mean=2.643e-10 max=1.710e-09 rel_l2=3.109e-06
self_attn.v_proj.weight mean=8.508e-10 max=5.937e-09 rel_l2=1.940e-06
self_attn.o_proj.weight mean=6.060e-10 max=4.191e-09 rel_l2=1.955e-06
self_attn.q_norm.weight mean=3.068e-10 max=1.033e-09 rel_l2=3.072e-06
self_attn.k_norm.weight mean=2.920e-10 max=9.459e-10 rel_l2=2.852e-06
mlp.gate_proj.weight mean=2.613e-10 max=1.804e-09 rel_l2=2.604e-06
mlp.up_proj.weight mean=2.578e-10 max=2.154e-09 rel_l2=2.574e-06
mlp.down_proj.weight mean=2.583e-10 max=2.052e-09 rel_l2=2.570e-06
layer 2:
self_attn.q_proj.weight mean=1.112e-10 max=7.603e-10 rel_l2=3.181e-06
self_attn.k_proj.weight mean=1.570e-10 max=9.459e-10 rel_l2=3.190e-06
self_attn.v_proj.weight mean=6.215e-10 max=4.016e-09 rel_l2=1.835e-06
self_attn.o_proj.weight mean=4.358e-10 max=3.041e-09 rel_l2=1.858e-06
self_attn.q_norm.weight mean=1.763e-10 max=6.439e-10 rel_l2=3.296e-06
self_attn.k_norm.weight mean=1.733e-10 max=6.694e-10 rel_l2=3.466e-06
mlp.gate_proj.weight mean=1.927e-10 max=1.357e-09 rel_l2=2.603e-06
mlp.up_proj.weight mean=1.881e-10 max=1.382e-09 rel_l2=2.559e-06
mlp.down_proj.weight mean=1.882e-10 max=1.426e-09 rel_l2=2.566e-06
layer 3:
self_attn.q_proj.weight mean=7.808e-11 max=4.729e-10 rel_l2=3.138e-06
self_attn.k_proj.weight mean=1.115e-10 max=7.130e-10 rel_l2=3.158e-06
self_attn.v_proj.weight mean=4.587e-10 max=3.194e-09 rel_l2=1.724e-06
self_attn.o_proj.weight mean=3.223e-10 max=2.270e-09 rel_l2=1.722e-06
self_attn.q_norm.weight mean=1.201e-10 max=5.239e-10 rel_l2=3.376e-06
self_attn.k_norm.weight mean=1.122e-10 max=5.807e-10 rel_l2=3.396e-06
mlp.gate_proj.weight mean=1.553e-10 max=1.397e-09 rel_l2=2.583e-06
mlp.up_proj.weight mean=1.509e-10 max=1.259e-09 rel_l2=2.523e-06
mlp.down_proj.weight mean=1.510e-10 max=1.124e-09 rel_l2=2.534e-06
layer 4:
self_attn.q_proj.weight mean=6.025e-11 max=4.193e-10 rel_l2=3.089e-06
self_attn.k_proj.weight mean=8.500e-11 max=5.348e-10 rel_l2=3.057e-06
self_attn.v_proj.weight mean=3.460e-10 max=2.212e-09 rel_l2=1.609e-06
self_attn.o_proj.weight mean=2.412e-10 max=1.630e-09 rel_l2=1.584e-06
self_attn.q_norm.weight mean=9.615e-11 max=2.328e-10 rel_l2=2.924e-06
self_attn.k_norm.weight mean=9.981e-11 max=3.183e-10 rel_l2=2.993e-06
mlp.gate_proj.weight mean=1.308e-10 max=1.106e-09 rel_l2=2.558e-06
mlp.up_proj.weight mean=1.262e-10 max=1.091e-09 rel_l2=2.492e-06
mlp.down_proj.weight mean=1.264e-10 max=1.106e-09 rel_l2=2.512e-06
layer 5:
self_attn.q_proj.weight mean=4.951e-11 max=3.311e-10 rel_l2=3.021e-06
self_attn.k_proj.weight mean=7.024e-11 max=4.520e-10 rel_l2=3.016e-06
self_attn.v_proj.weight mean=2.676e-10 max=1.921e-09 rel_l2=1.464e-06
self_attn.o_proj.weight mean=1.860e-10 max=1.237e-09 rel_l2=1.439e-06
self_attn.q_norm.weight mean=7.855e-11 max=3.020e-10 rel_l2=3.367e-06
self_attn.k_norm.weight mean=8.454e-11 max=2.874e-10 rel_l2=3.909e-06
mlp.gate_proj.weight mean=1.139e-10 max=9.750e-10 rel_l2=2.524e-06
mlp.up_proj.weight mean=1.092e-10 max=9.732e-10 rel_l2=2.473e-06
mlp.down_proj.weight mean=1.101e-10 max=9.459e-10 rel_l2=2.483e-06
layer 6:
self_attn.q_proj.weight mean=4.187e-11 max=2.829e-10 rel_l2=2.945e-06
self_attn.k_proj.weight mean=5.946e-11 max=4.002e-10 rel_l2=2.967e-06
self_attn.v_proj.weight mean=2.114e-10 max=1.339e-09 rel_l2=1.360e-06
self_attn.o_proj.weight mean=1.424e-10 max=1.004e-09 rel_l2=1.293e-06
self_attn.q_norm.weight mean=6.910e-11 max=2.401e-10 rel_l2=3.077e-06
self_attn.k_norm.weight mean=7.143e-11 max=2.265e-10 rel_l2=3.271e-06
mlp.gate_proj.weight mean=1.010e-10 max=7.349e-10 rel_l2=2.498e-06
mlp.up_proj.weight mean=9.703e-11 max=8.731e-10 rel_l2=2.432e-06
mlp.down_proj.weight mean=9.790e-11 max=8.440e-10 rel_l2=2.461e-06
layer 7:
self_attn.q_proj.weight mean=3.678e-11 max=2.442e-10 rel_l2=2.880e-06
self_attn.k_proj.weight mean=5.276e-11 max=3.420e-10 rel_l2=2.879e-06
self_attn.v_proj.weight mean=1.647e-10 max=1.222e-09 rel_l2=1.192e-06
self_attn.o_proj.weight mean=1.086e-10 max=7.421e-10 rel_l2=1.112e-06
self_attn.q_norm.weight mean=5.902e-11 max=1.910e-10 rel_l2=3.262e-06
self_attn.k_norm.weight mean=5.784e-11 max=1.728e-10 rel_l2=3.097e-06
mlp.gate_proj.weight mean=8.968e-11 max=6.839e-10 rel_l2=2.431e-06
mlp.up_proj.weight mean=8.535e-11 max=6.548e-10 rel_l2=2.386e-06
mlp.down_proj.weight mean=8.784e-11 max=6.985e-10 rel_l2=2.449e-06
AdamW accumulated logits diff:
step 10: mean=1.496e-05 max=3.853e-04 rel_l2=3.307e-05
step 20: mean=1.597e-05 max=3.788e-04 rel_l2=3.328e-05
step 30: mean=1.682e-05 max=3.462e-04 rel_l2=3.379e-05
step 40: mean=1.645e-05 max=3.592e-04 rel_l2=3.249e-05
step 50: mean=1.578e-05 max=3.576e-04 rel_l2=3.092e-05
step 60: mean=1.516e-05 max=3.442e-04 rel_l2=2.959e-05
step 70: mean=1.468e-05 max=3.333e-04 rel_l2=2.859e-05
step 80: mean=1.432e-05 max=3.321e-04 rel_l2=2.786e-05
step 90: mean=1.405e-05 max=3.282e-04 rel_l2=2.730e-05
step 100: mean=1.384e-05 max=3.318e-04 rel_l2=2.686e-05
AdamW final parameter max abs diff: 9.860471e-05
显存和运行时间对比
python scripts/benchmark_tensor_parallel_memory.py \
--layers 2 4 8 16 32 48 64 96 128 \
--tp_size 2 \
--hidden_size 768 \
--num_attention_heads 8 \
--num_key_value_heads 4 \
--vocab_size 6400 \
--seq_len 340 \
--batch_size 4 \
--dtype float32 \
--seed 42 \
--learning_rate 5e-4 \
--warmup_iters 3 \
--benchmark_iters 10 \
--output_csv tp_memory_scaling.csv \
--output_plot tp_memory_scaling.png
layers= 2 dense: 588.60 MiB, 16.68 ms
layers= 2 TP: 402.58 MiB, 19.58 ms
layers= 4 dense: 998.38 MiB, 29.99 ms
layers= 4 TP: 635.31 MiB, 35.29 ms
layers= 8 dense: 1824.40 MiB, 56.89 ms
layers= 8 TP: 1093.79 MiB, 65.84 ms
layers= 16 dense: 3471.01 MiB, 111.00 ms
layers= 16 TP: 2011.12 MiB, 130.01 ms
layers= 32 dense: 6764.02 MiB, 220.03 ms
layers= 32 TP: 3841.29 MiB, 252.47 ms
layers= 48 dense: 10036.89 MiB, 335.47 ms
layers= 48 TP: 5677.38 MiB, 378.14 ms
layers= 64 dense: 13325.74 MiB, 433.46 ms
layers= 64 TP: 7514.72 MiB, 498.86 ms
layers= 96 dense: 19926.52 MiB, 646.44 ms
layers= 96 TP: 11183.27 MiB, 750.35 ms
layers=128 dense: OOM/failed
layers=128 TP: 14849.10 MiB, 1011.53 ms
saved CSV: tp_memory_scaling.csv
saved plot: tp_memory_scaling.png

目前实现了:
项目同时整理了一套在线教程:
https://anker661.github.io/minimind-tron/
目前 TP / SP / VP 部分已整理完成,后续 PP, CP 以及多维并行部分将在整理好后发布。
感谢 MiniMind 提供了足够小且清晰的模型实现,让这些并行策略能够从原理、tensor layout、通信和 autograd 一步步展开。
jingyao你好!非常棒的项目,我在Minimind的基础上实现了一个用来学习Tensor Parallel的版本
https://github.com/ANKer661/minimind/tree/tp_demo
实现内容
当前实现参考了Nvidia 2019年的最初的Megatron LM论文的层内Tensor Parallel方式
torch.autograd.Function实现 TP 通信语义。为了避免影响现有代码,目前 TP 实现均放在独立文件中,没有修改原有模型和训练流程:
model/model_tp.pyscripts/demo_tensor_parallel.pyscripts/benchmark_tensor_parallel_memory.pyscripts/benchmark_tensor_parallel_worker.py局限
当前实现目前主要用于理解Tensor Parallel的原理和实现和数值验证,存在以下限制:
验证结果
数值验证
结果包括加载相同权重后,第一次forward的logits误差,第一次backward后所有被并行的layer的梯度误差,以及使用AdamW优化器进行更新后每x步的logits的累积误差。
显存和运行时间对比