使用cublasSgemmStridedBatched遇CUBLAS_STATUS_INVALID_VALUE错误求助
问题分析与修复方案
错误码CUBLAS_STATUS_INVALID_VALUE表明传入的参数不符合API要求,结合行主序矩阵的特性和cuBLAS接口的列主序设计,问题出在转置参数、矩阵维度顺序、变量名大小写三个方面,具体修复如下:
1. 修正变量名大小写
C++是大小写敏感语言,代码中定义的变量是batch_count,但调用API时误写为batchCount,需统一为batch_count。
2. 调整转置参数与矩阵维度
cuBLAS默认以列主序处理矩阵,而你的张量是行主序存储,直接使用CUBLAS_OP_N会导致矩阵维度不匹配。要计算行主序下的A(20×9) * B(9×32) = C(20×32),等价于列主序下的B^T(32×9) * A^T(9×20) = C^T(32×20),因此需调整:
- 转置参数改为
CUBLAS_OP_T(对列主序存储的矩阵转置,还原为行主序逻辑矩阵) - 维度参数顺序调整为
N, M, K(对应列主序下结果矩阵的行数、列数、公共维度)
3. 修正Leading Dimension参数
Leading Dimension(lda/ldb/ldc)是列主序矩阵的行数,需对应转置后的矩阵维度:
lda:列主序B的行数=32(行主序B是9×32,列主序存储为32×9)ldb:列主序A的行数=9(行主序A是20×9,列主序存储为9×20)ldc:列主序C^T的行数=32(行主序C是20×32,列主序存储为32×20)
修复后的完整代码
int batch_count = 60000; int M = 20; int K = 9; int N = 32; cublasHandle_t handle; cublasCreate(&handle); float alpha = 1.0; float beta = 0.0; int strideA = 20 * 9; int strideB = 0; int strideC = 20 * 32; // A(60000 * 20 * 9) * B(9 * 32) = C(60000 * 20 * 32) cublasStatus_t ret = cublasSgemmStridedBatched( handle, CUBLAS_OP_T, // 转置列主序B,得到行主序B CUBLAS_OP_T, // 转置列主序A,得到行主序A N, // 结果矩阵C^T的行数(对应行主序C的列数) M, // 结果矩阵C^T的列数(对应行主序C的行数) K, // 矩阵乘法公共维度 &alpha, B.data<float>(), // GPU上的列主序B(32×9) N, // lda:列主序B的行数 strideB, A.data<float>(), // GPU上的列主序A(9×20) K, // ldb:列主序A的行数 strideA, &beta, C.data<float>(), // GPU上的列主序C^T(32×20) N, // ldc:列主序C^T的行数 strideC, batch_count); // 修正变量名大小写 cublasDestroy(handle); if(ret != CUBLAS_STATUS_SUCCESS){ printf("cublasSgemmStridedBatched failed %d line (%d)\n", ret, __LINE__); }
额外验证建议
- 先用小批量(如
batch_count=1)测试,对比CPU计算结果,确保矩阵乘法逻辑正确。 - 检查GPU内存分配:A需60000×20×9×4=43.2MB,B需9×32×4=1.152KB,C需60000×20×32×4=153.6MB,总内存约196MB,主流GPU均可满足。
内容的提问来源于stack exchange,提问作者user9875189
相关产品推荐
相关产品推荐

