如何编写可返回块对角矩阵的C++ mex函数
C++ MEX 实现行块对角矩阵生成方案
实现逻辑
- 首先确定输入输出维度:假设输入为
M行N列矩阵,输出为M行(M*N)列矩阵,所有元素默认初始化为0 - 仅对输出矩阵的块对角区域赋值:输入第
i行的N个元素,对应输出第i行、列范围[i*N, (i+1)*N - 1]的区域 - 注意MATLAB矩阵为列优先存储,取值和赋值时需要按照列优先规则计算线性索引,避免位置错误
完整MEX代码
#include "mex.h" void mexFunction(int nlhs, mxArray *plhs[], int nrhs, const mxArray *prhs[]) { // 输入参数校验 if(nrhs != 1) { mexErrMsgIdAndTxt("block_diag_rows:nrhs", "仅需要1个输入参数"); } if(!mxIsDouble(prhs[0]) || mxIsComplex(prhs[0])) { mexErrMsgIdAndTxt("block_diag_rows:inputType", "输入必须为非复数双精度矩阵"); } if(nlhs > 1) { mexErrMsgIdAndTxt("block_diag_rows:nlhs", "最多返回1个输出参数"); } // 获取输入矩阵维度 mwSize M = mxGetM(prhs[0]); mwSize N = mxGetN(prhs[0]); // 创建输出矩阵 M x (M*N),默认初始化为0 plhs[0] = mxCreateDoubleMatrix(M, M*N, mxREAL); // 获取输入输出数据指针 double *in_ptr = mxGetPr(prhs[0]); double *out_ptr = mxGetPr(plhs[0]); // 逐行赋值块对角区域 for(mwSize i = 0; i < M; i++) { // 遍历输入每一行 mwSize block_start_col = i * N; // 当前块在输出中的起始列 for(mwSize j = 0; j < N; j++) { // 遍历当前行的所有元素 // 计算输入(i,j)的列优先索引 mwSize in_idx = j * M + i; // 计算输出(i, block_start_col + j)的列优先索引 mwSize out_idx = (block_start_col + j) * M + i; out_ptr[out_idx] = in_ptr[in_idx]; } } }
编译与使用
- 将上述代码保存为
block_diag_rows.cpp - 在MATLAB命令行执行编译命令:
mex block_diag_rows.cpp,首次使用需要先通过mex -setup配置好C++编译器 - 测试效果:
input_matrix = [1,2,3;4,5,6;7,8,9]; output_matrix = block_diag_rows(input_matrix)
运行后得到的output_matrix与你给出的示例完全一致。
内容的提问来源于stack exchange,提问作者Joshua Cochrane
相关产品推荐
相关产品推荐

