MPI_Comm_split_type 和你前面问的 MPI_Comm_split 很像,都是从一个已有的通信器中创建新的通信器。
但两者分组依据不同:
MPI_Comm_split:你自己通过color指定怎么分组。MPI_Comm_split_type:让 MPI 根据硬件/拓扑关系来分组。
在你现在的 MPI + 多节点 + 多 GPU 场景里,最常见的用途就是:
把同一台计算节点上的 MPI 进程放进同一个 communicator,从而得到
local_rank。
函数原型
cpp
int MPI_Comm_split_type(
MPI_Comm comm,
int split_type,
int key,
MPI_Info info,
MPI_Comm *newcomm
);
最常见的写法:
cpp
MPI_Comm localComm;
MPI_Comm_split_type(
MPI_COMM_WORLD,
MPI_COMM_TYPE_SHARED,
0,
MPI_INFO_NULL,
&localComm
);
其中最关键的是:
cpp
MPI_COMM_TYPE_SHARED
它要求 MPI:
把能够共享内存的进程划分到同一个 communicator。
在常见的 HPC 集群中,这通常就对应于同一计算节点上的进程。
假设:
bash
mpirun -np 8 ...
8 个进程分布在两个节点:
text
MPI_COMM_WORLD
Node A Node B
───────────────── ─────────────────
world rank 0 world rank 4
world rank 1 world rank 5
world rank 2 world rank 6
world rank 3 world rank 7
执行:
cpp
MPI_Comm_split_type(
MPI_COMM_WORLD,
MPI_COMM_TYPE_SHARED,
0,
MPI_INFO_NULL,
&localComm
);
MPI 会创建两个不同的通信器:
text
MPI_COMM_WORLD
0 1 2 3 4 5 6 7
│
MPI_Comm_split_type
│
MPI_COMM_TYPE_SHARED
│
┌────────┴────────┐
↓ ↓
Node A Node B
0 1 2 3 4 5 6 7 ← world rank
↓ ↓
0 1 2 3 0 1 2 3 ← local rank
注意这里发生了非常重要的事情:
text
world rank = 4
local rank = 0
完全没问题。
因为 rank 永远是相对于某个 communicator而言的。
获取 local_rank
因此实际代码通常是:
cpp
int worldRank;
MPI_Comm_rank(MPI_COMM_WORLD, &worldRank);
MPI_Comm localComm;
MPI_Comm_split_type(
MPI_COMM_WORLD,
MPI_COMM_TYPE_SHARED,
0,
MPI_INFO_NULL,
&localComm
);
int localRank;
int localSize;
MPI_Comm_rank(localComm, &localRank);
MPI_Comm_size(localComm, &localSize);
假设两个节点,每个节点 4 个进程:
text
worldRank localRank localSize
0 0 4
1 1 4
2 2 4
3 3 4
4 0 4
5 1 4
6 2 4
7 3 4
这就解释了你之前遇到过的一个现象:
为什么 8 个 MPI 进程、4 张 GPU 的情况下,第二个节点的
local_rank又从 0 开始?
因为第二个节点拥有自己的 localComm。
这对 GPU 绑定特别有用
假设每个节点都有 4 张 GPU:
text
Node A Node B
localRank 0 → GPU 0 localRank 0 → GPU 0
localRank 1 → GPU 1 localRank 1 → GPU 1
localRank 2 → GPU 2 localRank 2 → GPU 2
localRank 3 → GPU 3 localRank 3 → GPU 3
于是可以:
cpp
cudaSetDevice(localRank);
而不能简单:
cpp
cudaSetDevice(worldRank); // 多节点情况下通常错误
否则第二个节点:
text
worldRank 4 → cudaSetDevice(4) ❌
worldRank 5 → cudaSetDevice(5) ❌
worldRank 6 → cudaSetDevice(6) ❌
worldRank 7 → cudaSetDevice(7) ❌
如果节点只有 GPU 0~3,就会出现你之前见过的:
text
invalid device ordinal
一个完整 Demo
cpp
#include <mpi.h>
#include <iostream>
int main(int argc, char **argv)
{
MPI_Init(&argc, &argv);
int worldRank;
int worldSize;
MPI_Comm_rank(MPI_COMM_WORLD, &worldRank);
MPI_Comm_size(MPI_COMM_WORLD, &worldSize);
// 创建"节点内部"的 communicator
MPI_Comm localComm;
MPI_Comm_split_type(
MPI_COMM_WORLD,
MPI_COMM_TYPE_SHARED,
0,
MPI_INFO_NULL,
&localComm
);
// 获取节点内部 rank
int localRank;
int localSize;
MPI_Comm_rank(localComm, &localRank);
MPI_Comm_size(localComm, &localSize);
std::cout
<< "worldRank = " << worldRank
<< ", worldSize = " << worldSize
<< ", localRank = " << localRank
<< ", localSize = " << localSize
<< std::endl;
// 自己创建的 communicator 要释放
MPI_Comm_free(&localComm);
MPI_Finalize();
return 0;
}
如果是:
text
2 个节点
每节点 4 个 MPI 进程
典型结果就是:
text
Node A:
worldRank 0 → localRank 0
worldRank 1 → localRank 1
worldRank 2 → localRank 2
worldRank 3 → localRank 3
Node B:
worldRank 4 → localRank 0
worldRank 5 → localRank 1
worldRank 6 → localRank 2
worldRank 7 → localRank 3
所以对于你现在的代码,先把这条关系牢牢记住:
text
MPI_COMM_WORLD
│
│ MPI_Comm_split_type
│ MPI_COMM_TYPE_SHARED
↓
localComm
│
│ MPI_Comm_rank
↓
localRank
│
↓
节点内部进程编号
│
↓
可以进一步用于 GPU 绑定
而 MPI_Comm_split 与它的区别,可以简单记成:
text
MPI_Comm_split
↓
我告诉 MPI 怎么分组(color)
MPI_Comm_split_type
↓
我告诉 MPI 按什么"类型/拓扑关系"分组
↓
MPI_COMM_TYPE_SHARED
↓
通常得到同一节点内的进程组
对于 MPI + CUDA/HIP 多节点程序,MPI_Comm_split_type(..., MPI_COMM_TYPE_SHARED, ...) 是获取节点内 localRank 的非常标准的办法。
localComm是怎么区分不同的计算节点的
对,你这个疑问正好碰到了 MPI communicator 最容易误解的地方:
localComm只是你代码里的一个变量名。虽然每个进程都写的是localComm,但每个 MPI 进程有自己独立的地址空间,所以它们各自拥有自己的localComm变量。
并不存在一个全局的 localComm 变量里面同时装着"节点1通信器"和"节点2通信器"。
假设两个节点,每个节点 4 个进程:
text
Node A Node B
worldRank worldRank
0 1 2 3 4 5 6 7
所有 8 个进程都执行:
cpp
MPI_Comm localComm;
MPI_Comm_split_type(
MPI_COMM_WORLD,
MPI_COMM_TYPE_SHARED,
0,
MPI_INFO_NULL,
&localComm
);
执行完以后,可以概念上理解成:
text
进程 0:
localComm ──→ [0,1,2,3]
进程 1:
localComm ──→ [0,1,2,3]
进程 2:
localComm ──→ [0,1,2,3]
进程 3:
localComm ──→ [0,1,2,3]
进程 4:
localComm ──→ [4,5,6,7]
进程 5:
localComm ──→ [4,5,6,7]
进程 6:
localComm ──→ [4,5,6,7]
进程 7:
localComm ──→ [4,5,6,7]
所以从逻辑上确实产生了两个 communicator:
text
MPI_COMM_WORLD
0 1 2 3 4 5 6 7
│
↓
MPI_Comm_split_type
│
┌─────┴─────┐
↓ ↓
communicator A communicator B
0 1 2 3 4 5 6 7
但是代码里的:
cpp
MPI_Comm localComm;
不是"保存所有 communicator 的容器"。
它表示:
对于当前这个 MPI 进程,我属于哪个新 communicator。
因此,对于 worldRank == 2 的进程:
cpp
localComm
指向的是 Node A 那个通信器。
对于 worldRank == 6 的进程:
cpp
localComm
指向的是 Node B 那个通信器。
那 MPI 到底怎么知道谁在同一个节点?
关键就是:
cpp
MPI_COMM_TYPE_SHARED
执行:
cpp
MPI_Comm_split_type(
MPI_COMM_WORLD,
MPI_COMM_TYPE_SHARED,
...
);
MPI 实现根据运行时掌握的进程拓扑,把能够创建共享内存区域的进程归为一组。在典型集群部署里,这通常就是同一个计算节点上的进程。
所以不是你的:
cpp
localComm
变量在判断节点。
而是:
text
MPI runtime
│
│ 根据 MPI_COMM_TYPE_SHARED
↓
判断哪些进程可以共享内存
│
├──────────────┐
↓ ↓
0 1 2 3 4 5 6 7
│ │
↓ ↓
communicator A communicator B
最后 MPI 把当前进程所属的那个 communicator 的句柄写入当前进程自己的:
cpp
localComm
再看 localRank 就更清楚了
所有进程继续执行:
cpp
int localRank;
MPI_Comm_rank(
localComm,
&localRank
);
注意这里问的不是:
"我在
MPI_COMM_WORLD里面是多少?"
而是在问:
"我在我自己的
localComm里面是多少?"
所以:
text
MPI_COMM_WORLD
worldRank
0 1 2 3 4 5 6 7
│ │ │ │ │ │ │ │
│ Node A │ │ Node B │
│ │ │ │
└───────────┬───────────┘ └───────────┬───────────┘
↓ ↓
localComm A localComm B
0 1 2 3 4 5 6 7
│ │ │ │ │ │ │ │
↓ ↓ ↓ ↓ ↓ ↓ ↓ ↓
local 0 1 2 3 0 1 2 3
rank
所以你之前看到:
cpp
MPI_Comm_rank(localComm, &localRank);
得到:
text
worldRank localRank
0 0
1 1
2 2
3 3
4 0 ← 关键
5 1
6 2
7 3
本质原因就是:
worldRank 0 和 worldRank 4 虽然代码里的变量都叫 localComm,但它们所属的 communicator 根本不是同一个。
还有一个更底层的点:不要把 MPI_Comm 当成"进程列表本身"。它更像一个 MPI 内部通信上下文的句柄(handle):
cpp
MPI_Comm localComm;
可以粗略理解成:
text
当前进程的 localComm
↓
一个句柄
↓
MPI 内部管理的 communicator
↓
知道当前进程属于哪个组、组内 rank 等信息
这和 C++ 里"每个进程都有一个同名局部变量"是两回事。
所以你问的"同一个 localComm 变量怎么区分两个节点",答案就是:它根本不需要区分。每个 MPI 进程只持有自己所属的那个 localComm。节点的划分已经由 MPI_Comm_split_type 完成了。