如何在Matlab mex函数中按行读取输入矩阵以得到正确对角展开结果
问题原因分析
- 输出矩阵维度定义错误:需求中输出矩阵行数始终等于输入矩阵的行数,列数为输入矩阵行数×列数,原代码在行列数不等时错误将输出行数设为输入列数,导致输出矩阵形状不符合预期。
- 输入矩阵元素读取索引错误:Matlab矩阵按列优先存储,第i行(0起始)第j列(0起始)的元素对应的一维索引为
i + j * rows,原代码索引写反,导致按列读取元素。 - 输出赋值逻辑冗余复杂:原代码使用count变量的寻址逻辑不符合行展开需求,无需复杂计算偏移。
修正后完整代码
#include <matrix.h> #include <mex.h> void mexFunction(int nlhs, mxArray *plhs[], int nrhs, const mxArray *prhs[]) { const mwSize *dims; double *a, *b; int rows, cols; // 获取输入矩阵维度 dims = mxGetDimensions(prhs[0]); rows = (int) dims[0]; cols = (int) dims[1]; // 创建输出矩阵:行数=输入行数,列数=输入行数×输入列数,默认初始化为0 plhs[0] = mxCreateDoubleMatrix(rows, rows * cols, mxREAL); // 获取输入输出数组的指针 a = mxGetPr(prhs[0]); b = mxGetPr(plhs[0]); // 逐行处理输入矩阵 for (int i = 0; i < rows; i++) { for(int j = 0; j < cols; j++){ // 读取输入第i行第j列的元素 double val = a[i + j * rows]; // 计算输出位置:第i行第 (i*cols + j) 列 int output_col = i * cols + j; int output_idx = i + output_col * rows; b[output_idx] = val; } } }
核心修改说明
- 统一输出矩阵维度:输出行数固定为输入行数rows,列数固定为rows*cols,符合示例输出形状,同时去掉了多余的输入矩阵复制逻辑,减少不必要的内存开销。
- 修正输入元素索引:将原错误的
a[j + rows * i]改为a[i + j * rows],正确读取输入矩阵对应行列的元素。 - 简化输出赋值逻辑:直接根据要写入的输出行列位置计算列优先下的一维索引,逻辑清晰不易出错,其余位置默认初始化为0无需额外处理。
内容的提问来源于stack exchange,提问作者Joshua Cochrane
相关产品推荐
相关产品推荐

