基于MPI实现的矩阵乘法代码仅支持4x4矩阵运行的问题排查
MPI矩阵乘法代码无法支持大尺寸矩阵的修复方案
原代码核心问题
- 硬编码矩阵尺寸:宏
#define N 4写死了所有固定栈数组的长度为4,输入大于4的尺寸时,数组写入直接越界破坏内存,程序必然崩溃或计算错误 - 矩阵维度未同步:仅0号进程读取了输入的行、列值,其余进程的这两个变量是未初始化的随机值,后续计算、通信的尺寸参数全错
- 通信分片逻辑写死:
MPI_Scatter/MPI_Gather/MPI_Bcast的传输计数全部用写死的N计算,完全不匹配实际输入的矩阵大小,且未校验进程数能否整除矩阵行数,通信时会出现缓冲区溢出或数据截断 - 计算逻辑硬编码边界:多层循环的终止条件全部绑定固定值
N,和实际矩阵尺寸脱钩,就算改大宏定义,计算结果也会出错 - 栈数组不适合大尺寸矩阵:固定长度的栈上数组空间上限极低,尺寸稍大就会触发栈溢出。
MPI下动态数组使用说明
MPI通信要求发送/接收的缓冲区是连续内存块,因此不要用int**二级指针做逐行malloc的非连续数组,直接用malloc分配一维连续内存即可,二维索引手动换算成一维偏移(比如矩阵第i行j列元素对应arr[i*n + j]),所有进程都根据广播得到的实际矩阵尺寸分配本地缓冲区即可,不需要特殊的MPI专属内存接口。
修复后可运行代码
#include <stdio.h> #include <stdlib.h> #include <time.h> #include "mpi.h" void print_results(char *prompt, int *a, int n); int main(int argc, char *argv[]) { int i, j, k, rank, size, sum = 0; int *a, *b, *c; int *aa, *cc; int n, rows_per_proc; double time1, time2, duration, global; MPI_Status status; MPI_Init(&argc, &argv); MPI_Comm_size(MPI_COMM_WORLD, &size); MPI_Comm_rank(MPI_COMM_WORLD, &rank); if(rank == 0){ printf("enter the matrix size n (n x n square matrix) ="); scanf("%d",&n); // 校验进程数可整除行数,保证平均分配 if (n % size != 0) { printf("Error: process count %d must divide matrix size %d\n", size, n); MPI_Abort(MPI_COMM_WORLD, 1); } // 动态分配0号进程上的三个矩阵 a = (int*)malloc(n * n * sizeof(int)); b = (int*)malloc(n * n * sizeof(int)); c = (int*)malloc(n * n * sizeof(int)); srand(time(NULL)); for(i=0;i<n;i++) { for(j=0;j<n;j++){ a[i*n + j] = rand() % 10; } } for(i=0;i<n;i++){ for(j=0;j<n;j++){ b[i*n + j] = rand() % 10; } } } // 先把矩阵尺寸广播给所有进程 MPI_Bcast(&n, 1, MPI_INT, 0, MPI_COMM_WORLD); // 校验非0进程上的进程数和尺寸匹配性 if (n % size != 0) { MPI_Finalize(); exit(1); } rows_per_proc = n / size; // 所有进程分配本地缓冲区 aa = (int*)malloc(rows_per_proc * n * sizeof(int)); cc = (int*)malloc(rows_per_proc * n * sizeof(int)); // 非0进程分配b矩阵的接收缓冲区 if (rank != 0) { b = (int*)malloc(n * n * sizeof(int)); } MPI_Barrier(MPI_COMM_WORLD); time1 = MPI_Wtime(); // 按行分片散射A矩阵 MPI_Scatter(a, rows_per_proc * n, MPI_INT, aa, rows_per_proc * n, MPI_INT, 0, MPI_COMM_WORLD); // 广播整个B矩阵到所有进程 MPI_Bcast(b, n*n, MPI_INT, 0, MPI_COMM_WORLD); MPI_Barrier(MPI_COMM_WORLD); // 矩阵乘法计算:每个进程算分到的rows_per_proc行C矩阵 for (i = 0; i < rows_per_proc; i++) { for (j = 0; j < n; j++) { sum = 0; for (k = 0; k < n; k++) { sum += aa[i*n + k] * b[k*n + j]; } cc[i*n + j] = sum; } } // 收集所有进程算完的C矩阵分片 MPI_Gather(cc, rows_per_proc * n, MPI_INT, c, rows_per_proc * n, MPI_INT, 0, MPI_COMM_WORLD); MPI_Barrier(MPI_COMM_WORLD); time2 = MPI_Wtime(); duration = time2 - time1; MPI_Reduce(&duration,&global,1,MPI_DOUBLE,MPI_MAX,0,MPI_COMM_WORLD); if(rank == 0) { printf("Global runtime is %f s\n",global); } printf("Runtime at rank %d is %f s\n", rank,duration); if (rank == 0) print_results("C = ", c, n); // 释放所有动态分配的内存 free(aa); free(cc); free(b); if (rank == 0) { free(a); free(c); } MPI_Finalize(); return 0; } void print_results(char *prompt, int *a, int n) { int i, j; printf ("\n\n%s\n", prompt); for (i = 0; i < n; i++) { for (j = 0; j < n; j++) { printf(" %d", a[i*n + j]); } printf ("\n"); } printf ("\n\n"); }
编译运行注意事项
- 编译时链接mpi库,命令参考
mpicc matmul.c -o matmul - 运行时指定进程数必须能整除你输入的矩阵边长,比如跑8x8矩阵用4个进程就执行
mpirun -np 4 ./matmul,输入n=8即可 - 如果要支持非整除的分片场景,把平均分配改成
MPI_Scatterv/MPI_Gatherv做变长分片即可,逻辑和当前版本一致。
内容的提问来源于stack exchange,提问作者h3avyc0der
相关产品推荐
相关产品推荐

