C语言MPI编程实现进程间数组数据均匀负载均衡方法
MPI相邻链式负载均衡实现思路
你已经通过MPI_Allreduce和MPI_Exscan拿到了全局总长度、单进程目标长度、当前进程前的全局前缀和,只需要基于全局索引区间做相邻进程的定向收发即可,全程仅需和rank-1、rank+1两个邻居通信,不需要跨多节点建立连接,通信开销最低。
前置修正(现有代码的遗漏点)
- 注意:MPI_Exscan 对0号进程的返回值无定义,必须手动将rank0的
exScan赋值为0 - 你当前的
new_size是向上取整的结果,总长度可能无法被进程数整除,必须单独修正最后一个进程的目标长度,避免数组越界:
if (myRank == 0) exScan = 0; if (myRank == nbTask - 1) { new_size = global_sum - myRank * new_size; }
核心逻辑
每个进程先计算两个连续索引区间:
- 本地原始数据覆盖的全局索引区间:
[local_start, local_end] = [exScan, exScan + size - 1] - 重分发后当前进程应该持有的全局索引区间:
[target_start, target_end] = [myRank * new_size, myRank * new_size + new_size - 1]
由于所有数组元素是按全局索引连续排布的,本地数据和目标区间的差异只会出现在数组头部和尾部:
- 头部如果有索引小于
target_start的元素,全部属于左邻居(rank-1),直接发给左邻居即可 - 头部如果缺索引小于
target_start的元素,缺失部分全部来自左邻居,直接从左邻居接收即可 - 尾部如果有索引大于
target_end的元素,全部属于右邻居(rank+1),直接发给右邻居即可 - 尾部如果缺索引大于
target_end的元素,缺失部分全部来自右邻居,直接从右邻居接收即可
全程使用MPI_Sendrecv完成收发,不需要手动协调收发顺序,从根源避免死锁。
参考实现片段
首先生成本地测试数组(可以给元素赋值为对应的全局索引,方便后续验证结果):
// 生成本地随机长度数组 int* local_arr = (int*)malloc(size * sizeof(int)); for (int i = 0; i < size; i++) { local_arr[i] = exScan + i; // 元素值等于全局索引,方便校验 } int local_start = exScan; int local_end = exScan + size - 1; int target_start = myRank * new_size; int target_end = target_start + new_size - 1; int* curr_arr = local_arr; int curr_size = size;
处理和左邻居(仅myRank>0时存在)的数据交换:
if (myRank > 0) { int diff = target_start - local_start; if (diff > 0) { // 头部多了diff个元素,发给左邻居 MPI_Sendrecv(curr_arr, diff, MPI_INT, myRank-1, 0, NULL, 0, MPI_INT, myRank-1, 1, MPI_COMM_WORLD, MPI_STATUS_IGNORE); // 移除本地数组头部已发送的部分 int* tmp = (int*)malloc((curr_size - diff) * sizeof(int)); memcpy(tmp, curr_arr + diff, (curr_size - diff) * sizeof(int)); free(curr_arr); curr_arr = tmp; curr_size -= diff; } else if (diff < 0) { // 头部缺-diff个元素,从左邻居接收 int recv_cnt = -diff; int* tmp = (int*)malloc((curr_size + recv_cnt) * sizeof(int)); MPI_Sendrecv(NULL, 0, MPI_INT, myRank-1, 1, tmp, recv_cnt, MPI_INT, myRank-1, 0, MPI_COMM_WORLD, MPI_STATUS_IGNORE); // 收到的元素插到数组头部 memcpy(tmp + recv_cnt, curr_arr, curr_size * sizeof(int)); free(curr_arr); curr_arr = tmp; curr_size += recv_cnt; } }
处理和右邻居(仅myRank<nbTask-1时存在)的数据交换,逻辑和左邻居对称:
if (myRank < nbTask - 1) { int diff = local_end - target_end; if (diff > 0) { // 尾部多了diff个元素,发给右邻居 MPI_Sendrecv(curr_arr + curr_size - diff, diff, MPI_INT, myRank+1, 0, NULL, 0, MPI_INT, myRank+1, 1, MPI_COMM_WORLD, MPI_STATUS_IGNORE); // 移除本地数组尾部已发送的部分 int* tmp = (int*)malloc((curr_size - diff) * sizeof(int)); memcpy(tmp, curr_arr, (curr_size - diff) * sizeof(int)); free(curr_arr); curr_arr = tmp; curr_size -= diff; } else if (diff < 0) { // 尾部缺-diff个元素,从右邻居接收 int recv_cnt = -diff; int* tmp = (int*)malloc((curr_size + recv_cnt) * sizeof(int)); memcpy(tmp, curr_arr, curr_size * sizeof(int)); free(curr_arr); MPI_Sendrecv(NULL, 0, MPI_INT, myRank+1, 1, tmp + curr_size, recv_cnt, MPI_INT, myRank+1, 0, MPI_COMM_WORLD, MPI_STATUS_IGNORE); curr_arr = tmp; curr_size += recv_cnt; } }
结果校验
传输完成后,每个进程的curr_size必然等于修正后的new_size,如果测试时给元素赋值为全局索引,此时curr_arr[0]应该等于target_start,curr_arr[curr_size-1]应该等于target_end,直接打印即可验证正确性。
内容的提问来源于stack exchange,提问作者bous
相关产品推荐
相关产品推荐

