Cublas实现矩阵乘法报DGEMM参数8非法 领先维度配置问题求助
CUBLAS DGEMM参数非法报错问题解决方法
问题根因
你遇到的DGEMM参数8非法报错,直接原因是设备内存分配大小写反,导致矩阵A的设备指针越界非法:
你的A矩阵共12个元素,B矩阵共9个元素,但mulWithCuda函数中给dev_a只分配了9个元素的空间,给dev_b分配了12个元素的空间,后续往dev_a拷贝12个double时发生越界,导致指针非法。
修复步骤
第一步:修正设备内存分配逻辑
将mulWithCuda中的内存分配代码修改为正确大小:
// 修正后的内存分配代码 cudaMalloc((void**)&dev_c, 12 * sizeof(double)); cudaMalloc((void**)&dev_a, 12 * sizeof(double)); // A矩阵共12个元素 cudaMalloc((void**)&dev_b, size * sizeof(double)); // B矩阵共9个元素
第二步:适配CUBLAS的列优先存储规则
CUBLAS默认所有矩阵按列优先存储,而你的代码中矩阵是行优先存储,需要调整DGEMM参数适配:
cublasDgemm的运算逻辑为 C = α * op(A) * op(B) + β * C,针对你4×3的A矩阵乘3×3的B矩阵得到4×3的C矩阵的需求,行优先存储下的正确调用参数为:
cublasDgemm(handle, CUBLAS_OP_T, CUBLAS_OP_T, 3, 4, 3, &alpha, dev_b, 3, dev_a, 3, &beta, dev_c, 3);
你也可以选择将矩阵改为列优先存储,此时调用参数为:
// 若A、B改为列优先存储的调用写法 cublasDgemm(handle, CUBLAS_OP_N, CUBLAS_OP_N, 4, 3, 3, &alpha, dev_a, 4, dev_b, 3, &beta, dev_c, 4);
第三步:添加错误检查(建议)
为所有CUDA、CUBLAS接口调用添加返回值校验,可以快速定位后续可能出现的其他问题,示例如下:
// CUBLAS调用错误检查示例 cublasStatus_t stat = cublasDgemm(/*参数*/); if (stat != CUBLAS_STATUS_SUCCESS) { printf("cublasDgemm failed, error code: %d\n", stat); }
修复完成后运行即可得到正确的矩阵乘法结果。
内容的提问来源于stack exchange,提问作者Dresult
相关产品推荐
相关产品推荐

