MPI_Comm_split_type用法

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 完成了。

相关推荐
胖大和尚3 小时前
MPI_Probe
mpi
Logintern0921 天前
【每天学习一点点】MPI是什么?
学习·mpi·多机训练
HHFQ8 个月前
适用于Rocky Linux的MPI 管理程序
mpi·openmpi·mpich
2401_841495649 个月前
并行程序设计与实现
c++·python·算法·cuda·mpi·并行计算·openmp
戴草帽的大z10 个月前
rk3588上用rk_mpi_vi_test与ffmpeg实战
ffmpeg·rk3588·mpi·rk_mpi_vi
Eloudy2 年前
编译安装调试 scaLapack 和 openmpi 以及 lapack
mpi·lapack
noravinsc2 年前
paraview 跨节点并行渲染
mpi·paraview
长安er2 年前
高性能并行计算华为云实验一:MPI矩阵运算
华为云·并行·高性能·mpi·高性能并行计算·串行
canmoumou2 年前
【分布式通信】NPKit,NCCL的Profiling工具
人工智能·mpi·nccl